From 7661cc44beea344ec5296ef8172742b030f39198 Mon Sep 17 00:00:00 2001 From: ssongliu Date: Thu, 27 Aug 2026 13:58:46 +0800 Subject: [PATCH] refactor: simplify firewall service structure --- agent/app/api/v2/docker_port_guard_test.go | 49 - agent/app/api/v2/firewall_rule_test.go | 70 - agent/app/dto/firewall.go | 8 +- agent/app/model/firewall.go | 110 + agent/app/model/firewall_test.go | 78 - agent/app/repo/firewall_rule_test.go | 120 - agent/app/repo/forwarding_rule_test.go | 47 - .../{firewall_service.go => firewall.go} | 569 +---- .../service/firewall_database_sync_test.go | 41 - agent/app/service/firewall_docker.go | 201 +- agent/app/service/firewall_docker_test.go | 572 ----- agent/app/service/firewall_service_test.go | 2038 ----------------- agent/app/service/firewall_setting.go | 26 +- agent/app/service/firewall_setting_test.go | 174 -- agent/app/service/firewall_sync.go | 397 +--- agent/app/service/firewall_sync_test.go | 683 ------ .../{forward_service.go => forward.go} | 17 +- agent/app/service/forwarding_contract_test.go | 686 ------ agent/init/migration/migrate_test.go | 37 - .../migration/migrations/firewall_test.go | 327 --- .../migrations/utils/firewall_transfer.go | 5 +- .../utils/firewall_transfer_test.go | 316 --- .../utils/host_firewall_transfer_test.go | 365 --- agent/utils/docker/docker.go | 3 + .../firewall/docker_guard/manager_test.go | 255 --- .../firewall/docker_guard/nftables_test.go | 210 -- agent/utils/firewall/docker_guard/policy.go | 131 ++ agent/utils/firewall/docker_guard/runtime.go | 83 + agent/utils/firewall/filter/check_flag.go | 170 ++ agent/utils/firewall/filter/check_test.go | 366 --- agent/utils/firewall/filter/identity_test.go | 195 -- agent/utils/firewall/filter/inventory_test.go | 256 --- agent/utils/firewall/filter/model.go | 29 + agent/utils/firewall/filter/model_test.go | 85 - agent/utils/firewall/filter/normalize_test.go | 270 --- .../providers/firewalld/adapter_test.go | 579 ----- .../filter/providers/iptables/adapter_test.go | 753 ------ .../filter/providers/nftables/adapter_test.go | 243 -- .../filter/providers/ufw/adapter_test.go | 782 ------- .../utils/firewall/filter/runtime/runtime.go | 307 +++ agent/utils/firewall/filter/safety_test.go | 139 -- agent/utils/firewall/forwarding/adapter.go | 52 - .../utils/firewall/forwarding/adapter_test.go | 20 - agent/utils/firewall/forwarding/forwarding.go | 203 ++ agent/utils/firewall/forwarding/manager.go | 83 - .../providers/adapter_contract_test.go | 329 --- .../firewall/forwarding/providers/iptables.go | 2 +- .../firewall/forwarding/providers/nftables.go | 2 +- .../forwarding/providers/nftables_test.go | 137 -- .../forwarding/providers/normalize.go | 82 - .../firewall/iptables_helper/inspect_test.go | 12 - .../iptables_helper/manager_restore_test.go | 214 -- agent/utils/firewall/lifecycle/client.go | 105 - agent/utils/firewall/lifecycle/client_test.go | 61 - agent/utils/firewall/lifecycle/lifecycle.go | 230 ++ .../utils/firewall/lifecycle/operator_test.go | 48 - .../utils/firewall/lifecycle/provider_test.go | 43 - .../lifecycle/providers/firewalld_test.go | 107 - agent/utils/firewall/lifecycle/runtime.go | 84 - .../utils/firewall/lifecycle/runtime_test.go | 41 - agent/utils/firewall/lifecycle/status.go | 50 - agent/utils/firewall/lifecycle/status_test.go | 55 - .../utils/firewall/nftables_helper/command.go | 66 - .../firewall/nftables_helper/command_test.go | 51 - .../firewall/nftables_helper/manager_test.go | 45 - .../{inspect.go => runtime.go} | 59 + agent/utils/firewall/port_whitelist.go | 91 + agent/utils/firewall/port_whitelist_test.go | 72 - agent/utils/firewall/sync/diff.go | 120 + agent/utils/firewall/sync/order.go | 141 ++ agent/utils/re/firewall.go | 1 + core/init/migration/migrate_test.go | 24 - .../migration/migrations/firewall_test.go | 114 - frontend/src/api/interface/firewall.ts | 1 + .../src/views/host/firewall/sync/index.vue | 13 +- 75 files changed, 1858 insertions(+), 12692 deletions(-) delete mode 100644 agent/app/api/v2/docker_port_guard_test.go delete mode 100644 agent/app/api/v2/firewall_rule_test.go delete mode 100644 agent/app/model/firewall_test.go delete mode 100644 agent/app/repo/firewall_rule_test.go delete mode 100644 agent/app/repo/forwarding_rule_test.go rename agent/app/service/{firewall_service.go => firewall.go} (82%) delete mode 100644 agent/app/service/firewall_database_sync_test.go delete mode 100644 agent/app/service/firewall_docker_test.go delete mode 100644 agent/app/service/firewall_service_test.go delete mode 100644 agent/app/service/firewall_setting_test.go delete mode 100644 agent/app/service/firewall_sync_test.go rename agent/app/service/{forward_service.go => forward.go} (97%) delete mode 100644 agent/app/service/forwarding_contract_test.go delete mode 100644 agent/init/migration/migrate_test.go delete mode 100644 agent/init/migration/migrations/firewall_test.go delete mode 100644 agent/init/migration/migrations/utils/firewall_transfer_test.go delete mode 100644 agent/init/migration/migrations/utils/host_firewall_transfer_test.go delete mode 100644 agent/utils/firewall/docker_guard/manager_test.go delete mode 100644 agent/utils/firewall/docker_guard/nftables_test.go create mode 100644 agent/utils/firewall/docker_guard/policy.go create mode 100644 agent/utils/firewall/docker_guard/runtime.go create mode 100644 agent/utils/firewall/filter/check_flag.go delete mode 100644 agent/utils/firewall/filter/check_test.go delete mode 100644 agent/utils/firewall/filter/identity_test.go delete mode 100644 agent/utils/firewall/filter/inventory_test.go delete mode 100644 agent/utils/firewall/filter/model_test.go delete mode 100644 agent/utils/firewall/filter/normalize_test.go delete mode 100644 agent/utils/firewall/filter/providers/firewalld/adapter_test.go delete mode 100644 agent/utils/firewall/filter/providers/iptables/adapter_test.go delete mode 100644 agent/utils/firewall/filter/providers/nftables/adapter_test.go delete mode 100644 agent/utils/firewall/filter/providers/ufw/adapter_test.go create mode 100644 agent/utils/firewall/filter/runtime/runtime.go delete mode 100644 agent/utils/firewall/filter/safety_test.go delete mode 100644 agent/utils/firewall/forwarding/adapter.go delete mode 100644 agent/utils/firewall/forwarding/adapter_test.go create mode 100644 agent/utils/firewall/forwarding/forwarding.go delete mode 100644 agent/utils/firewall/forwarding/manager.go delete mode 100644 agent/utils/firewall/forwarding/providers/adapter_contract_test.go delete mode 100644 agent/utils/firewall/forwarding/providers/nftables_test.go delete mode 100644 agent/utils/firewall/forwarding/providers/normalize.go delete mode 100644 agent/utils/firewall/iptables_helper/inspect_test.go delete mode 100644 agent/utils/firewall/iptables_helper/manager_restore_test.go delete mode 100644 agent/utils/firewall/lifecycle/client.go delete mode 100644 agent/utils/firewall/lifecycle/client_test.go create mode 100644 agent/utils/firewall/lifecycle/lifecycle.go delete mode 100644 agent/utils/firewall/lifecycle/operator_test.go delete mode 100644 agent/utils/firewall/lifecycle/provider_test.go delete mode 100644 agent/utils/firewall/lifecycle/providers/firewalld_test.go delete mode 100644 agent/utils/firewall/lifecycle/runtime.go delete mode 100644 agent/utils/firewall/lifecycle/runtime_test.go delete mode 100644 agent/utils/firewall/lifecycle/status.go delete mode 100644 agent/utils/firewall/lifecycle/status_test.go delete mode 100644 agent/utils/firewall/nftables_helper/command.go delete mode 100644 agent/utils/firewall/nftables_helper/command_test.go delete mode 100644 agent/utils/firewall/nftables_helper/manager_test.go rename agent/utils/firewall/nftables_helper/{inspect.go => runtime.go} (53%) delete mode 100644 agent/utils/firewall/port_whitelist_test.go create mode 100644 agent/utils/firewall/sync/diff.go create mode 100644 agent/utils/firewall/sync/order.go delete mode 100644 core/init/migration/migrate_test.go delete mode 100644 core/init/migration/migrations/firewall_test.go diff --git a/agent/app/api/v2/docker_port_guard_test.go b/agent/app/api/v2/docker_port_guard_test.go deleted file mode 100644 index cfe26c51c95f..000000000000 --- a/agent/app/api/v2/docker_port_guard_test.go +++ /dev/null @@ -1,49 +0,0 @@ -package v2 - -import ( - "encoding/json" - "fmt" - "net/http" - "net/http/httptest" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/app/service" - agenti18n "github.com/1Panel-dev/1Panel/agent/i18n" - "github.com/gin-gonic/gin" -) - -func TestHandleDockerPortGuardErrorReturnsStableBusinessCode(t *testing.T) { - agenti18n.Init() - gin.SetMode(gin.TestMode) - recorder := httptest.NewRecorder() - context, _ := gin.CreateTestContext(recorder) - handleDockerPortGuardError(context, fmt.Errorf("normalize policy: %w", service.ErrDockerGuardInvalid)) - - var response dto.Response - if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil { - t.Fatalf("decode response: %v", err) - } - if response.Code != http.StatusBadRequest || response.ErrorCode != "FW_DOCKER_GUARD_INVALID" { - t.Fatalf("unexpected Docker guard error response: %#v", response) - } -} - -func TestHandleDockerPortGuardErrorLocalizesDockerUnavailable(t *testing.T) { - agenti18n.Init() - gin.SetMode(gin.TestMode) - recorder := httptest.NewRecorder() - context, _ := gin.CreateTestContext(recorder) - handleDockerPortGuardError(context, fmt.Errorf("inspect Docker: %w", service.ErrDockerUnavailable)) - - var response dto.Response - if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil { - t.Fatalf("decode response: %v", err) - } - if response.Code != http.StatusServiceUnavailable || response.ErrorCode != "FW_DOCKER_UNAVAILABLE" { - t.Fatalf("unexpected Docker unavailable response: %#v", response) - } - if response.Message != agenti18n.Get("ErrDockerFailed") { - t.Fatalf("message = %q, want localized Docker failure", response.Message) - } -} diff --git a/agent/app/api/v2/firewall_rule_test.go b/agent/app/api/v2/firewall_rule_test.go deleted file mode 100644 index d3a68c6c2bb1..000000000000 --- a/agent/app/api/v2/firewall_rule_test.go +++ /dev/null @@ -1,70 +0,0 @@ -package v2 - -import ( - "encoding/json" - "fmt" - "net/http" - "net/http/httptest" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/app/repo" - agenti18n "github.com/1Panel-dev/1Panel/agent/i18n" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - "github.com/gin-gonic/gin" -) - -func TestHandleFirewallRuleErrorReturnsStableBusinessCode(t *testing.T) { - agenti18n.Init() - gin.SetMode(gin.TestMode) - tests := []struct { - name string - err error - code string - }{ - {name: "stale rule", err: filter.ErrRuleStale, code: "FW_RULE_STALE"}, - {name: "check required", err: filter.ErrRuleCheckRequired, code: "FW_RULE_CHECK_REQUIRED"}, - {name: "revision conflict", err: fmt.Errorf("persist rule: %w", repo.ErrFirewallRuleRevisionConflict), code: "FW_RULE_REVISION_CONFLICT"}, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - recorder := httptest.NewRecorder() - context, _ := gin.CreateTestContext(recorder) - handleFirewallRuleError(context, test.err) - var response dto.Response - if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil { - t.Fatalf("decode response: %v", err) - } - if response.Code != 409 || response.ErrorCode != test.code { - t.Fatalf("unexpected firewall error response: %#v", response) - } - }) - } -} - -func TestNormalizeFirewallRuleUUID(t *testing.T) { - gin.SetMode(gin.TestMode) - t.Run("trims valid UUID", func(t *testing.T) { - recorder := httptest.NewRecorder() - context, _ := gin.CreateTestContext(recorder) - value := " managed-rule " - if !normalizeFirewallRuleUUID(context, &value) || value != "managed-rule" { - t.Fatalf("normalized UUID = %q", value) - } - }) - t.Run("rejects blank UUID", func(t *testing.T) { - recorder := httptest.NewRecorder() - context, _ := gin.CreateTestContext(recorder) - value := " " - if normalizeFirewallRuleUUID(context, &value) { - t.Fatal("blank UUID was accepted") - } - var response dto.Response - if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil { - t.Fatalf("decode response: %v", err) - } - if recorder.Code != http.StatusOK || response.Code != http.StatusBadRequest { - t.Fatalf("transport status = %d, response = %#v", recorder.Code, response) - } - }) -} diff --git a/agent/app/dto/firewall.go b/agent/app/dto/firewall.go index f784af0954cf..9f60eaf22e54 100644 --- a/agent/app/dto/firewall.go +++ b/agent/app/dto/firewall.go @@ -1,6 +1,9 @@ package dto -import "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" +import ( + "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" + firewallsync "github.com/1Panel-dev/1Panel/agent/utils/firewall/sync" +) type FirewallSubsystemStatus struct { Name string `json:"name"` @@ -245,7 +248,8 @@ type FirewallRuleSyncItem struct { Rule *filter.FirewallRule `json:"rule,omitempty"` ForwardRule *ForwardRule `json:"forwardRule,omitempty"` DockerRule *DockerPortGuardEndpoint `json:"dockerRule,omitempty"` - Status string `json:"status"` + Status firewallsync.Status `json:"status"` + ReasonCode firewallsync.ReasonCode `json:"reasonCode,omitempty"` Reason string `json:"reason,omitempty"` } diff --git a/agent/app/model/firewall.go b/agent/app/model/firewall.go index 3abc45adb86b..ef21eeb9ebde 100644 --- a/agent/app/model/firewall.go +++ b/agent/app/model/firewall.go @@ -5,6 +5,7 @@ import ( "encoding/hex" "encoding/json" "fmt" + "sort" "strings" "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" @@ -67,6 +68,19 @@ func FirewallRuleOwner(sourceKind, sourceID string) string { return sourceKind + ":" + sourceID } +func FirewallRulesRevision(rules []FirewallRule) (string, error) { + ordered := append([]FirewallRule(nil), rules...) + sort.Slice(ordered, func(i, j int) bool { + return ordered[i].UUID < ordered[j].UUID + }) + payload, err := json.Marshal(ordered) + if err != nil { + return "", err + } + sum := sha256.Sum256(payload) + return hex.EncodeToString(sum[:]), nil +} + func FirewallRuleFromDomain(rule filter.FirewallRule) (FirewallRule, error) { normalized, err := filter.NormalizeRule(rule) if err != nil { @@ -115,3 +129,99 @@ func (rule FirewallRule) PolicyKey() string { sum := sha256.Sum256(payload) return hex.EncodeToString(sum[:]) } + +func (rule FirewallRule) RulesForProvider(provider filter.Provider) ([]filter.FirewallRule, error) { + if rule.CompatibilityError != "" { + return nil, fmt.Errorf("%w: %s", filter.ErrUnsupportedScope, rule.CompatibilityError) + } + connectionStates := make([]string, 0) + if rule.ConnectionStates != "" { + connectionStates = strings.Split(rule.ConnectionStates, ",") + } + base := filter.FirewallRule{ + Protocol: rule.Protocol, SourceAddress: rule.SourceAddress, SourcePort: rule.SourcePort, + DestinationAddress: rule.DestinationAddress, DestinationPort: rule.DestinationPort, + Interface: rule.Interface, ConnectionStates: connectionStates, + Action: filter.Action(rule.Action), Description: rule.Description, + } + if provider == filter.ProviderFirewalld { + base.Priority = rule.Priority + } + families := []filter.Family{filter.Family(rule.Family)} + if provider != filter.ProviderFirewalld && families[0] == filter.FamilyInet { + hasIPv4, hasIPv6 := ruleAddressFamilies(base) + switch { + case hasIPv4 && hasIPv6: + return nil, fmt.Errorf("%w: inet policy contains both IPv4 and IPv6 addresses", filter.ErrUnsupportedScope) + case hasIPv6 || strings.EqualFold(base.Protocol, "icmpv6"): + families = []filter.Family{filter.FamilyIPv6} + case hasIPv4: + families = []filter.Family{filter.FamilyIPv4} + default: + families = []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} + } + } + result := make([]filter.FirewallRule, 0, len(families)) + for _, family := range families { + compiled := base + compiled.Scope = filter.Scope{Provider: provider, Family: family, Direction: filter.DirectionInput} + switch provider { + case filter.ProviderIptables, filter.ProviderNftables: + compiled.Scope.Table, compiled.Scope.Chain = "filter", filter.IptablesInputChain + case filter.ProviderFirewalld: + compiled.Scope.Zone = filter.FirewalldInputZone + case filter.ProviderUFW: + compiled.Scope.Chain = filter.UFWInputChain + default: + return nil, fmt.Errorf("%w: unsupported firewall provider %q", filter.ErrProviderUnavailable, provider) + } + expanded, err := filter.ExpandAtomicRules(compiled) + if err != nil { + return nil, err + } + result = append(result, expanded...) + } + return result, nil +} + +func SortFirewallRules(rules []FirewallRule, provider filter.Provider) { + sort.SliceStable(rules, func(i, j int) bool { + left, right := rules[i], rules[j] + if provider == filter.ProviderFirewalld { + switch { + case left.Priority == nil && right.Priority != nil: + return false + case left.Priority != nil && right.Priority == nil: + return true + case left.Priority != nil && right.Priority != nil && *left.Priority != *right.Priority: + return *left.Priority < *right.Priority + } + } else { + switch { + case left.Sequence == nil && right.Sequence != nil: + return false + case left.Sequence != nil && right.Sequence == nil: + return true + case left.Sequence != nil && right.Sequence != nil && *left.Sequence != *right.Sequence: + return *left.Sequence < *right.Sequence + } + } + return left.UUID < right.UUID + }) +} + +func ruleAddressFamilies(rule filter.FirewallRule) (bool, bool) { + hasIPv4, hasIPv6 := false, false + for _, address := range []string{rule.SourceAddress, rule.DestinationAddress} { + address = strings.TrimSpace(address) + if address == "" { + continue + } + if strings.Contains(address, ":") { + hasIPv6 = true + } else { + hasIPv4 = true + } + } + return hasIPv4, hasIPv6 +} diff --git a/agent/app/model/firewall_test.go b/agent/app/model/firewall_test.go deleted file mode 100644 index b22004ec1c62..000000000000 --- a/agent/app/model/firewall_test.go +++ /dev/null @@ -1,78 +0,0 @@ -package model - -import ( - "errors" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -func TestFirewallRuleFromDomainUsesProviderNeutralIdentity(t *testing.T) { - rule := filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - }, - Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, - } - iptables, err := FirewallRuleFromDomain(rule) - if err != nil { - t.Fatal(err) - } - rule.Scope.Provider = filter.ProviderNftables - nftables, err := FirewallRuleFromDomain(rule) - if err != nil { - t.Fatal(err) - } - if iptables.PolicyKey() != nftables.PolicyKey() { - t.Fatalf("provider leaked into desired-rule identity: iptables=%#v nftables=%#v", iptables, nftables) - } -} - -func TestFirewallRuleFromDomainRejectsProviderNativeOnlyRule(t *testing.T) { - _, err := FirewallRuleFromDomain(filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, - Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput, - }, - NativeKind: filter.NativeKindZoneService, - Protocol: "all", - Action: filter.ActionAccept, - }) - if !errors.Is(err, filter.ErrUnsupportedScope) { - t.Fatalf("provider-native rule error = %v, want unsupported scope", err) - } -} - -func TestFirewallRuleFromDomainPersistsOnlyFirewalldPriority(t *testing.T) { - priority := -100 - rule := filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, - Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput, - }, - NativeKind: filter.NativeKindRichRule, Protocol: "tcp", DestinationPort: "443", - Action: filter.ActionAccept, Priority: &priority, - } - firewalld, err := FirewallRuleFromDomain(rule) - if err != nil { - t.Fatal(err) - } - if firewalld.Priority == nil || *firewalld.Priority != priority || firewalld.Sequence != nil { - t.Fatalf("firewalld placement was not persisted: %#v", firewalld) - } - - rule.Scope = filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - rule.NativeKind = filter.NativeKindRule - rule.Priority = nil - iptables, err := FirewallRuleFromDomain(rule) - if err != nil { - t.Fatal(err) - } - if iptables.Priority != nil || iptables.Sequence != nil { - t.Fatalf("positional backend persisted provider priority: %#v", iptables) - } -} diff --git a/agent/app/repo/firewall_rule_test.go b/agent/app/repo/firewall_rule_test.go deleted file mode 100644 index b069dafe9ba8..000000000000 --- a/agent/app/repo/firewall_rule_test.go +++ /dev/null @@ -1,120 +0,0 @@ -package repo - -import ( - "context" - "errors" - "fmt" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/glebarez/sqlite" - "github.com/google/uuid" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -func TestFirewallRuleRepoRevision(t *testing.T) { - db := newFirewallRepoTestDB(t) - repository := NewFirewallRuleRepo(db) - ctx := context.Background() - rule := newFirewallRuleModel() - - if err := repository.Create(ctx, &rule); err != nil { - t.Fatalf("create rule: %v", err) - } - if rule.UUID == "" || rule.Revision != 1 { - t.Fatalf("rule defaults were not applied: %#v", rule) - } - - if err := repository.UpdateWithRevision(ctx, rule.UUID, 2, map[string]interface{}{"description": "updated"}); !errors.Is(err, ErrFirewallRuleRevisionConflict) { - t.Fatalf("expected revision conflict, got %v", err) - } - if err := repository.UpdateWithRevision(ctx, rule.UUID, 1, map[string]interface{}{"description": "updated"}); err != nil { - t.Fatalf("update rule: %v", err) - } - - updated, err := repository.GetByUUID(ctx, rule.UUID) - if err != nil { - t.Fatalf("get updated rule: %v", err) - } - if updated.Revision != 2 || updated.Description != "updated" { - t.Fatalf("unexpected updated rule: %#v", updated) - } - - if err := repository.DeleteWithRevision(ctx, rule.UUID, updated.Revision-1); !errors.Is(err, ErrFirewallRuleRevisionConflict) { - t.Fatalf("expected revision conflict, got %v", err) - } - if err := repository.DeleteWithRevision(ctx, rule.UUID, updated.Revision); err != nil { - t.Fatalf("delete rule: %v", err) - } - if _, err := repository.GetByUUID(ctx, rule.UUID); !errors.Is(err, gorm.ErrRecordNotFound) { - t.Fatalf("expected deleted rule to be absent, got %v", err) - } - var deletedCount int64 - if err := db.Model(&model.FirewallRule{}).Where("uuid = ?", rule.UUID).Count(&deletedCount).Error; err != nil || deletedCount != 0 { - t.Fatalf("rule was not hard deleted: count=%d err=%v", deletedCount, err) - } -} - -func TestFirewallRepositoriesRejectIncompleteRecords(t *testing.T) { - db := newFirewallRepoTestDB(t) - ctx := context.Background() - if err := NewFirewallRuleRepo(db).Create(ctx, &model.FirewallRule{}); !errors.Is(err, ErrFirewallPersistenceInvalid) { - t.Fatalf("expected invalid rule error, got %v", err) - } -} - -func TestFirewallRepositoriesUseContextTransaction(t *testing.T) { - db := newFirewallRepoTestDB(t) - ruleRepo := NewFirewallRuleRepo(db) - wantRollback := errors.New("rollback") - - err := db.Transaction(func(tx *gorm.DB) error { - ctx := context.WithValue(context.Background(), constant.DB, tx) - rule := newFirewallRuleModel() - if err := ruleRepo.Create(ctx, &rule); err != nil { - return err - } - return wantRollback - }) - if !errors.Is(err, wantRollback) { - t.Fatalf("expected rollback error, got %v", err) - } - - var ruleCount int64 - if err := db.Model(&model.FirewallRule{}).Count(&ruleCount).Error; err != nil { - t.Fatalf("count rules: %v", err) - } - if ruleCount != 0 { - t.Fatalf("transaction did not roll back: rules=%d", ruleCount) - } -} - -func newFirewallRepoTestDB(t *testing.T) *gorm.DB { - t.Helper() - dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString()) - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) - if err != nil { - t.Fatalf("open sqlite: %v", err) - } - sqlDB, err := db.DB() - if err != nil { - t.Fatalf("load sql db: %v", err) - } - sqlDB.SetMaxOpenConns(1) - t.Cleanup(func() { _ = sqlDB.Close() }) - if err := db.AutoMigrate(&model.FirewallRule{}); err != nil { - t.Fatalf("migrate models: %v", err) - } - return db -} - -func newFirewallRuleModel() model.FirewallRule { - return model.FirewallRule{ - Family: "ipv4", - Protocol: "tcp", - DestinationPort: "22", - Action: "accept", - } -} diff --git a/agent/app/repo/forwarding_rule_test.go b/agent/app/repo/forwarding_rule_test.go deleted file mode 100644 index 29fe7274a6f0..000000000000 --- a/agent/app/repo/forwarding_rule_test.go +++ /dev/null @@ -1,47 +0,0 @@ -package repo - -import ( - "context" - "fmt" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/global" - "github.com/glebarez/sqlite" - "github.com/google/uuid" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -func TestForwardingRuleRepoReplaceAll(t *testing.T) { - previousDB := global.DB - dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString()) - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) - if err != nil { - t.Fatal(err) - } - if err := db.AutoMigrate(&model.ForwardingRule{}); err != nil { - t.Fatal(err) - } - global.DB = db - t.Cleanup(func() { global.DB = previousDB }) - - repository := NewIForwardingRuleRepo() - first := []model.ForwardingRule{{Family: "ipv4", Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80"}} - if err := repository.ReplaceAll(context.Background(), first); err != nil { - t.Fatal(err) - } - rules, err := repository.List(context.Background()) - if err != nil || len(rules) != 1 || rules[0].Port != "8080" { - t.Fatalf("rules = %#v, err = %v", rules, err) - } - - second := []model.ForwardingRule{{Family: "ipv6", Protocol: "udp", Port: "5353", TargetIP: "::1", TargetPort: "53", Interface: "eth0"}} - if err := repository.ReplaceAll(context.Background(), second); err != nil { - t.Fatal(err) - } - rules, err = repository.List(context.Background()) - if err != nil || len(rules) != 1 || rules[0].Family != "ipv6" || rules[0].Port != "5353" { - t.Fatalf("rules = %#v, err = %v", rules, err) - } -} diff --git a/agent/app/service/firewall_service.go b/agent/app/service/firewall.go similarity index 82% rename from agent/app/service/firewall_service.go rename to agent/app/service/firewall.go index e15335132cf9..4b76bdb00291 100644 --- a/agent/app/service/firewall_service.go +++ b/agent/app/service/firewall.go @@ -2,11 +2,6 @@ package service import ( "context" - "crypto/hmac" - "crypto/sha256" - "encoding/base64" - "encoding/hex" - "encoding/json" "errors" "fmt" "os" @@ -22,10 +17,7 @@ import ( "github.com/1Panel-dev/1Panel/agent/global" "github.com/1Panel-dev/1Panel/agent/utils/firewall" "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - filterfirewalld "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/firewalld" - filteriptables "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/iptables" - filternftables "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/nftables" - filterufw "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/ufw" + filterruntime "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/runtime" "github.com/1Panel-dev/1Panel/agent/utils/firewall/iptables_helper" "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" "github.com/1Panel-dev/1Panel/agent/utils/firewall/nftables_helper" @@ -35,7 +27,7 @@ import ( type FirewallService struct { rules repo.IFirewallRuleRepo - adapters firewallRuleRuntimeRegistry + adapters firewallRuleRuntimeResolver forwardingSync firewallDatabaseSyncAdapter dockerSync firewallDatabaseSyncAdapter selectedProvider func(context.Context) (filter.Provider, error) @@ -48,23 +40,28 @@ type FirewallService struct { installedProviders func() []string } +type firewallRuleRuntimeResolver interface { + Resolve(filter.Provider) (*firewallRuleRuntime, error) + Providers() []filter.Provider +} + var firewallRuleMutationMu sync.Mutex type IFirewallService interface { LoadBaseInfo(chainGroup string) (dto.FirewallSubsystemStatus, error) OperateFirewall(request dto.FirewallLifecycleOperation) error OperateFilterChain(request dto.FilterChainOperation) error + Reset(context.Context, dto.FirewallRuleReset) (dto.FirewallRuleResetResponse, error) Inventory(context.Context, dto.FirewallRuleInventory) (dto.FirewallRuleInventoryResponse, error) LoadFirewallNativeDetail(context.Context, dto.FirewallNativeDetail) (string, error) Check(context.Context, string, dto.FirewallRuleCheck) (dto.FirewallRuleCheckResponse, error) Create(context.Context, dto.FirewallRuleCreate) (dto.FirewallRuleCreateResponse, error) - PreviewRuleSync(context.Context, string, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncPreview, error) - SyncRules(context.Context, string, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error) - CurrentRuleSyncTask() (dto.FirewallRuleSyncTask, error) Delete(context.Context, dto.FirewallRuleDelete) (dto.FirewallRuleDeleteResponse, error) Update(context.Context, string, dto.FirewallRuleUpdate) error Reorder(context.Context, string, dto.FirewallRuleReorder) error - Reset(context.Context, dto.FirewallRuleReset) (dto.FirewallRuleResetResponse, error) + PreviewRuleSync(context.Context, string, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncPreview, error) + SyncRules(context.Context, string, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error) + CurrentRuleSyncTask() (dto.FirewallRuleSyncTask, error) } func NewIFirewallService() IFirewallService { @@ -425,11 +422,7 @@ func (s *FirewallService) LoadFirewallNativeDetail(ctx context.Context, request if err != nil { return "", err } - informer, ok := runtime.adapter.(filter.NativeDetailReader) - if !ok { - return "", fmt.Errorf("%w: native details for %s", filter.ErrAdapterUnavailable, provider) - } - return informer.NativeDetail(ctx, request.Name, request.Permanent) + return runtime.NativeDetail(ctx, request.Name, request.Permanent) } func (s *FirewallService) checkUpdate( @@ -531,7 +524,7 @@ func (s *FirewallService) Check( if desiredErr != nil { return dto.FirewallRuleCheckResponse{}, desiredErr } - managedRevision, revisionErr := firewallManagedRevision(stored) + managedRevision, revisionErr := model.FirewallRulesRevision(stored) if revisionErr != nil { return dto.FirewallRuleCheckResponse{}, revisionErr } @@ -742,7 +735,7 @@ func (s *FirewallService) prepareCreate( if listErr != nil { return nil, index, listErr } - managedRevision, revisionErr := firewallManagedRevision(stored) + managedRevision, revisionErr := model.FirewallRulesRevision(stored) if revisionErr != nil { return nil, index, revisionErr } @@ -773,7 +766,7 @@ func nativeCreateBatchEnd(prepared []preparedFirewallRuleCreate, start int) int return start } first := prepared[start] - if first.runtime == nil || first.runtime.adapter == nil || !supportsNativeRuleBatch(first.runtime.adapter.Provider()) || + if first.runtime == nil || !supportsNativeRuleBatch(first.runtime.Provider()) || first.authorization.Operation != filter.ChangeCreate || first.request.Rule.OrderIndex != nil { return start + 1 } @@ -1160,14 +1153,14 @@ func (s *FirewallService) createRule( domainRule := request.Rule appendRule := false if authorization.Operation == filter.ChangeCreate && domainRule.Scope.Provider == filter.ProviderUFW && domainRule.OrderIndex == nil { - appendPosition, err := ufwAppendPosition(ctx, runtime, snapshot, domainRule) + appendPosition, err := runtime.AppendPosition(ctx, snapshot, domainRule) if err != nil { return err } domainRule.OrderIndex = &appendPosition appendRule = true } else if authorization.Operation == filter.ChangeCreate && domainRule.OrderIndex != nil { - maxPosition, err := maxPositionForRule(ctx, runtime, snapshot, domainRule) + maxPosition, err := runtime.MaxPosition(ctx, snapshot, domainRule) if err != nil { return err } @@ -1286,7 +1279,7 @@ func (s *FirewallService) deleteRule(ctx context.Context, ruleUUID string) error } restoreAtEnd := false if desired.Rule.Scope.Provider == filter.ProviderUFW && observed.Locator.Position != nil { - maxPosition, maxErr := maxPositionForRule(ctx, runtime, snapshot, desired.Rule) + maxPosition, maxErr := runtime.MaxPosition(ctx, snapshot, desired.Rule) if maxErr != nil { return rollback(maxErr) } @@ -1378,7 +1371,7 @@ func (s *FirewallService) reorderRule(ctx context.Context, clientIP, ruleUUID st if targetPosition == nil || *targetPosition < 1 { return fmt.Errorf("%w: target position is required", filter.ErrInvalidRule) } - if err := validatePositionTarget(ctx, runtime, snapshot, before.Rule, *targetPosition); err != nil { + if err := runtime.ValidatePosition(ctx, snapshot, before.Rule, *targetPosition); err != nil { return err } after.OrderIndex = targetPosition @@ -1408,101 +1401,6 @@ func (s *FirewallService) reorderRule(ctx context.Context, clientIP, ruleUUID st }) } -func validatePositionTarget( - ctx context.Context, - runtime *firewallRuleRuntime, - snapshot filter.Snapshot, - rule filter.FirewallRule, - targetPosition int64, -) error { - if rule.Scope.Provider == filter.ProviderUFW { - minPosition, maxPosition := snapshotPositionBounds(snapshot) - if targetPosition < minPosition || targetPosition > maxPosition { - return fmt.Errorf( - "%w: target position %d is outside the %s range %d-%d", - filter.ErrInvalidRule, targetPosition, rule.Scope.Family, minPosition, maxPosition, - ) - } - return nil - } - maxPosition, err := maxPositionForRule(ctx, runtime, snapshot, rule) - if err != nil { - return err - } - if targetPosition > maxPosition { - return fmt.Errorf("%w: target position %d is out of range 1-%d", filter.ErrInvalidRule, targetPosition, maxPosition) - } - return nil -} - -func ufwAppendPosition( - ctx context.Context, - runtime *firewallRuleRuntime, - snapshot filter.Snapshot, - rule filter.FirewallRule, -) (int64, error) { - if rule.Scope.Family == filter.FamilyIPv4 { - return snapshotMaxPosition(snapshot) + 1, nil - } - maxPosition, err := maxPositionForRule(ctx, runtime, snapshot, rule) - if err != nil { - return 0, err - } - return maxPosition + 1, nil -} - -func snapshotPositionBounds(snapshot filter.Snapshot) (int64, int64) { - minPosition, maxPosition := int64(0), int64(0) - for _, observed := range snapshot.Rules { - if observed.Locator.Position == nil { - continue - } - position := int64(*observed.Locator.Position) - if minPosition == 0 || position < minPosition { - minPosition = position - } - if position > maxPosition { - maxPosition = position - } - } - return minPosition, maxPosition -} - -func maxPositionForRule( - ctx context.Context, - runtime *firewallRuleRuntime, - snapshot filter.Snapshot, - rule filter.FirewallRule, -) (int64, error) { - maxPosition := snapshotMaxPosition(snapshot) - if rule.Scope.Provider == filter.ProviderUFW { - relatedScope := rule.Scope - if relatedScope.Family == filter.FamilyIPv4 { - relatedScope.Family = filter.FamilyIPv6 - } else { - relatedScope.Family = filter.FamilyIPv4 - } - relatedSnapshot, err := runtime.ObserveMutation(ctx, relatedScope) - if err != nil { - return 0, err - } - if relatedMax := snapshotMaxPosition(relatedSnapshot); relatedMax > maxPosition { - maxPosition = relatedMax - } - } - return maxPosition, nil -} - -func snapshotMaxPosition(snapshot filter.Snapshot) int64 { - var maxPosition int64 - for _, observed := range snapshot.Rules { - if observed.Locator.Position != nil && int64(*observed.Locator.Position) > maxPosition { - maxPosition = int64(*observed.Locator.Position) - } - } - return maxPosition -} - type managedMutationRequest struct { Stored model.FirewallRule Before filter.FirewallRule @@ -1566,7 +1464,7 @@ func (s *FirewallService) prepareManagedUpdate( if after.OrderIndex == nil { after.OrderIndex = ¤tPosition } else if *after.OrderIndex != currentPosition { - if err := validatePositionTarget(ctx, runtime, snapshot, before.Rule, *after.OrderIndex); err != nil { + if err := runtime.ValidatePosition(ctx, snapshot, before.Rule, *after.OrderIndex); err != nil { return preparedManagedUpdate{}, err } } @@ -1630,9 +1528,10 @@ func (s *FirewallService) selectedProviderForStoredRule( if s.selectedProvider != nil { return s.selectedProvider(ctx) } - if len(s.adapters) == 1 { - for provider := range s.adapters { - return provider, nil + if s.adapters != nil { + providers := s.adapters.Providers() + if len(providers) == 1 { + return providers[0], nil } } return "", fmt.Errorf("%w: selected provider is unavailable", filter.ErrProviderUnavailable) @@ -1649,7 +1548,7 @@ func (s *FirewallService) executeManagedMutation(ctx context.Context, request ma } appendRule, restoreAtEnd := false, false if after.Scope.Provider == filter.ProviderUFW && request.AdapterOperation == filter.ChangeUpdate { - maxPosition, maxErr := maxPositionForRule(ctx, request.Runtime, request.Snapshot, after) + maxPosition, maxErr := request.Runtime.MaxPosition(ctx, request.Snapshot, after) if maxErr != nil { return maxErr } @@ -1910,71 +1809,39 @@ func (s *FirewallService) adoptExternalSystemPort(ctx context.Context, port dto. } func systemPortRule(provider filter.Provider, port dto.FirewallSystemPort) filter.FirewallRule { - scope := filter.Scope{Provider: provider, Direction: filter.DirectionInput} - family := filter.Family(strings.ToLower(strings.TrimSpace(port.Family))) - switch provider { - case filter.ProviderIptables, filter.ProviderNftables: - if family != filter.FamilyIPv6 { - family = filter.FamilyIPv4 - } - scope.Family = family - scope.Table = "filter" - case filter.ProviderFirewalld: - if family != filter.FamilyIPv4 && family != filter.FamilyIPv6 { - family = filter.FamilyInet - } - scope.Family = family - scope.Zone = filter.FirewalldInputZone - case filter.ProviderUFW: - if family != filter.FamilyIPv6 { - family = filter.FamilyIPv4 - } - scope.Family = family - } - return filter.FirewallRule{ - Scope: scope, Protocol: port.Protocol, DestinationPort: port.Port, - Action: filter.ActionAccept, Description: "1Panel managed accepted port", - } + return firewall.RuleForSystemPort(provider, firewall.SystemPort(port)) } func normalizeSystemPorts(ports []dto.FirewallSystemPort) (map[string]dto.FirewallSystemPort, error) { - result := make(map[string]dto.FirewallSystemPort, len(ports)) + domainPorts := make([]firewall.SystemPort, 0, len(ports)) for _, port := range ports { - normalized, err := filter.NormalizeRule(systemPortRule(filter.ProviderIptables, port)) - if err != nil { - return nil, err - } - family := strings.ToLower(strings.TrimSpace(port.Family)) - if family != "" { - family = string(normalized.Scope.Family) - } - item := dto.FirewallSystemPort{ - Family: family, Port: normalized.DestinationPort, Protocol: normalized.Protocol, - } - result[systemPortKey(item)] = item + domainPorts = append(domainPorts, firewall.SystemPort(port)) + } + normalized, err := firewall.NormalizeSystemPorts(domainPorts) + if err != nil { + return nil, err + } + result := make(map[string]dto.FirewallSystemPort, len(normalized)) + for key, port := range normalized { + result[key] = dto.FirewallSystemPort(port) } return result, nil } func systemPortKey(port dto.FirewallSystemPort) string { - key := legacySystemPortKey(port) - if family := strings.ToLower(strings.TrimSpace(port.Family)); family != "" { - return family + "/" + key - } - return key + return firewall.SystemPortKey(firewall.SystemPort(port)) } func legacySystemPortKey(port dto.FirewallSystemPort) string { - return strings.ToLower(strings.TrimSpace(port.Protocol)) + "/" + strings.TrimSpace(port.Port) + return firewall.LegacySystemPortKey(firewall.SystemPort(port)) } func sortedSystemPortKeys(ports map[string]dto.FirewallSystemPort) []string { - keys := make([]string, 0, len(ports)) - for key := range ports { - keys = append(keys, key) + domainPorts := make(map[string]firewall.SystemPort, len(ports)) + for key, port := range ports { + domainPorts[key] = firewall.SystemPort(port) } - sort.Strings(keys) - return keys + return firewall.SortedSystemPortKeys(domainPorts) } func firewallRuleModelForCreate(rule filter.FirewallRule, request dto.FirewallRuleCreateItem, origin string) (model.FirewallRule, error) { @@ -2162,7 +2029,7 @@ func mergeFirewallInventory( } func desiredFirewallRuleFromModel(stored model.FirewallRule) (filter.DesiredRule, error) { - rules, err := firewallPolicyRulesForProvider(stored, filter.ProviderIptables) + rules, err := stored.RulesForProvider(filter.ProviderIptables) if err != nil { return filter.DesiredRule{}, err } @@ -2188,7 +2055,7 @@ func (s *FirewallService) compileStoredFirewallRules( stored model.FirewallRule, target filter.Provider, ) ([]filter.DesiredRule, error) { - rules, err := firewallPolicyRulesForProvider(stored, target) + rules, err := stored.RulesForProvider(target) if err != nil { return nil, err } @@ -2196,41 +2063,7 @@ func (s *FirewallService) compileStoredFirewallRules( if err != nil { return nil, err } - result := make([]filter.DesiredRule, 0, len(rules)) - scopeOrdinals := make(map[string]int) - for _, rule := range rules { - rule, err = runtime.Prepare(rule) - if err != nil { - return nil, err - } - if err = runtime.CheckRule(ctx, rule); err != nil { - return nil, err - } - ruleKey, keyErr := filter.RuleKey(rule) - if keyErr != nil { - return nil, keyErr - } - scopeKey := rule.Scope.Key() - ordinal := scopeOrdinals[scopeKey] - scopeOrdinals[scopeKey] = ordinal + 1 - rule.UUID = compiledFirewallRuleUUID(stored.UUID, ruleKey, ordinal) - result = append(result, filter.DesiredRule{ - UUID: stored.UUID, Rule: rule, RuleKey: ruleKey, Origin: filter.RuleOrigin(stored.Origin), - Marker: "1panel-rule:" + rule.UUID, - }) - } - return result, nil -} - -func compiledFirewallRuleUUID(policyUUID, ruleKey string, scopeOrdinal int) string { - if scopeOrdinal == 0 { - return policyUUID - } - const suffixLength = 12 - if len(ruleKey) > suffixLength { - ruleKey = ruleKey[:suffixLength] - } - return fmt.Sprintf("%s-%d-%s", policyUUID, scopeOrdinal+1, ruleKey) + return runtime.CompileDesired(ctx, stored.UUID, filter.RuleOrigin(stored.Origin), rules) } func (s *FirewallService) desiredFirewallRulesForScope( @@ -2238,7 +2071,7 @@ func (s *FirewallService) desiredFirewallRulesForScope( stored []model.FirewallRule, scope filter.Scope, ) ([]filter.DesiredRule, error) { - sortFirewallPolicies(stored, scope.Provider) + model.SortFirewallRules(stored, scope.Provider) desired := make([]filter.DesiredRule, 0, len(stored)) for _, record := range stored { compiled, err := s.compileStoredFirewallRules(ctx, record, scope.Provider) @@ -2266,146 +2099,26 @@ func firewallRuleSelectedProvider(context.Context) (filter.Provider, error) { return selectedRuleProvider() } -func rollbackFirewallPlan(ctx context.Context, runtime *firewallRuleRuntime, plan filter.BackendPlan, cause error) error { - if runtime == nil { - return cause - } - if err := runtime.Rollback(ctx, plan); err != nil { - return errors.Join(cause, fmt.Errorf("rollback applied firewall plan: %w", err)) - } - return cause -} - -type firewallSnapshotPolicy func(context.Context, filter.Snapshot) (filter.Snapshot, error) - -type firewallRuleRuntime struct { - adapter filter.Adapter - policy firewallSnapshotPolicy -} - -type firewallRuleRuntimeRegistry map[filter.Provider]*firewallRuleRuntime +type firewallSnapshotPolicy = filterruntime.SnapshotPolicy +type firewallRuleRuntime = filterruntime.Engine +type firewallRuleRuntimeRegistry = filterruntime.Registry func newFirewallRuleRuntimeRegistry(policy firewallSnapshotPolicy) firewallRuleRuntimeRegistry { - return firewallRuleRuntimeRegistry{ - filter.ProviderIptables: newFirewallRuleRuntime(filteriptables.NewAdapter(), policy), - filter.ProviderNftables: newFirewallRuleRuntime(filternftables.NewAdapter(), policy), - filter.ProviderFirewalld: newFirewallRuleRuntime(filterfirewalld.NewAdapter(), policy), - filter.ProviderUFW: newFirewallRuleRuntime(filterufw.NewAdapter(), policy), - } + return filterruntime.NewRegistry(policy) } func newFirewallRuleRuntime(adapter filter.Adapter, policy firewallSnapshotPolicy) *firewallRuleRuntime { - return &firewallRuleRuntime{adapter: adapter, policy: policy} + return filterruntime.New(adapter, policy) } -func (r firewallRuleRuntimeRegistry) Resolve(provider filter.Provider) (*firewallRuleRuntime, error) { - runtime, exists := r[provider] - if !exists || runtime == nil || runtime.adapter == nil { - return nil, fmt.Errorf("%w: %s", filter.ErrAdapterUnavailable, provider) - } - return runtime, nil -} - -func (r *firewallRuleRuntime) Observe(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) { - snapshot, err := r.adapter.Observe(ctx, scope) - if err != nil { - return filter.Snapshot{}, err - } - if r.policy == nil { - return snapshot, nil - } - return r.policy(ctx, snapshot) -} - -func (r *firewallRuleRuntime) ObserveScopes(ctx context.Context, scopes []filter.Scope) ([]filter.Snapshot, error) { - observer, ok := r.adapter.(filter.MultiScopeObserver) - if !ok { - return nil, fmt.Errorf("%w: %s multi-scope inventory", filter.ErrAdapterUnavailable, r.adapter.Provider()) - } - snapshots, err := observer.ObserveScopes(ctx, scopes) - if err != nil { - return nil, err - } - if r.policy == nil { - return snapshots, nil - } - for index := range snapshots { - snapshots[index], err = r.policy(ctx, snapshots[index]) - if err != nil { - return nil, err - } - } - return snapshots, nil -} - -func (r *firewallRuleRuntime) ObserveMutation(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) { - snapshot, err := r.Observe(ctx, scope) - if err != nil { - return filter.Snapshot{}, err - } - for _, notice := range snapshot.Notices { - if notice.Code == filter.ScopeNoticeManagedScopeInactive || notice.Code == filter.ScopeNoticeManagedScopeMissing { - return filter.Snapshot{}, fmt.Errorf("%w: managed firewall scope is unavailable", filter.ErrProviderUnavailable) - } - } - return snapshot, nil -} - -func (r *firewallRuleRuntime) Prepare(rule filter.FirewallRule) (filter.FirewallRule, error) { - preparer, ok := r.adapter.(filter.RulePreparer) - if !ok { - return rule, nil - } - return preparer.PrepareRule(rule) -} - -func (r *firewallRuleRuntime) CheckRule(ctx context.Context, rule filter.FirewallRule) error { - checker, ok := r.adapter.(filter.RuleChecker) - if !ok { - return nil - } - return checker.CheckRule(ctx, rule) -} - -func (r *firewallRuleRuntime) Capabilities(ctx context.Context) (filter.Capabilities, error) { - return r.adapter.Capabilities(ctx) -} - -func (r *firewallRuleRuntime) Execute(ctx context.Context, snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, filter.VerifyResult, error) { - plan, err := r.adapter.Compile(snapshot, changes) - if err != nil { - return filter.BackendPlan{}, filter.VerifyResult{}, err - } - result, err := r.adapter.Apply(ctx, plan) - if err != nil { - return plan, filter.VerifyResult{}, err - } - if result.Verification != nil { - if !result.Verification.Matched { - if rollbackErr := r.Rollback(ctx, plan); rollbackErr != nil { - return plan, *result.Verification, errors.Join(filter.ErrVerificationFailed, rollbackErr) - } - } - return plan, *result.Verification, nil - } - verification, err := r.adapter.Verify(ctx, plan) - if err != nil { - return plan, verification, rollbackFirewallPlan(ctx, r, plan, err) - } - if !verification.Matched { - if rollbackErr := r.Rollback(ctx, plan); rollbackErr != nil { - return plan, verification, errors.Join(filter.ErrVerificationFailed, rollbackErr) - } +func rollbackFirewallPlan(ctx context.Context, runtime *firewallRuleRuntime, plan filter.BackendPlan, cause error) error { + if runtime == nil { + return cause } - return plan, verification, nil -} - -func (r *firewallRuleRuntime) Rollback(ctx context.Context, plan filter.BackendPlan) error { - rollbacker, ok := r.adapter.(filter.PlanRollbacker) - if !ok { - return fmt.Errorf("%w: provider %s does not support applied-plan rollback", filter.ErrAdapterUnavailable, r.adapter.Provider()) + if err := runtime.Rollback(ctx, plan); err != nil { + return errors.Join(cause, fmt.Errorf("rollback applied firewall plan: %w", err)) } - return rollbacker.Rollback(ctx, plan) + return cause } func selectedRuleProvider() (filter.Provider, error) { @@ -2416,28 +2129,7 @@ func selectedRuleProvider() (filter.Provider, error) { return filter.Provider(provider), nil } -type firewallRuleCheckClaims struct { - Version int `json:"version"` - Provider filter.Provider `json:"provider"` - ScopeKey string `json:"scopeKey"` - RuleDigest string `json:"ruleDigest"` - SnapshotRevision string `json:"snapshotRevision"` - ManagedRevision string `json:"managedRevision"` - Decision filter.CheckDecision `json:"decision"` - Classification filter.CheckClassification `json:"classification"` - AllowedActions []filter.CheckAction `json:"allowedActions"` - AdoptionCandidates []firewallAdoptionCandidate `json:"adoptionCandidates,omitempty"` -} - -type firewallAdoptionCandidate struct { - InstanceKey string `json:"instanceKey"` - Locator filter.Locator `json:"locator"` -} - -type firewallRuleCreateAuthorization struct { - Operation filter.ChangeOperation - Locator *filter.Locator -} +type firewallRuleCreateAuthorization = filter.CreateAuthorization func refreshCreateAuthorization( snapshot filter.Snapshot, @@ -2486,36 +2178,7 @@ func refreshCreateAuthorization( } func signFirewallRuleCheck(result filter.RuleCheckResult, snapshot filter.Snapshot, managedRevision string) (string, error) { - ruleDigest, err := firewallRuleDigest(result.RequestedRule) - if err != nil { - return "", err - } - claims := firewallRuleCheckClaims{ - Version: constant.FirewallRuleCheckVersion, - Provider: result.RequestedRule.Scope.Provider, - ScopeKey: result.RequestedRule.Scope.Key(), - RuleDigest: ruleDigest, - SnapshotRevision: snapshot.Revision, - ManagedRevision: managedRevision, - Decision: result.Decision, - Classification: result.Classification, - AllowedActions: result.AllowedActions, - } - if result.Classification == filter.CheckClassificationExactExternal { - claims.AdoptionCandidates = make([]firewallAdoptionCandidate, 0, len(result.Candidates)) - for _, candidate := range result.Candidates { - claims.AdoptionCandidates = append(claims.AdoptionCandidates, firewallAdoptionCandidate{ - InstanceKey: candidate.InstanceKey, - Locator: candidate.Locator, - }) - } - } - payload, err := json.Marshal(claims) - if err != nil { - return "", err - } - signature := firewallRuleCheckSignature(payload) - return base64.RawURLEncoding.EncodeToString(payload) + "." + base64.RawURLEncoding.EncodeToString(signature), nil + return firewallCheckFlagCodec().Sign(result, snapshot, managedRevision) } func authorizeFirewallRuleCreate( @@ -2526,94 +2189,12 @@ func authorizeFirewallRuleCreate( snapshot filter.Snapshot, managedRevision string, ) (firewallRuleCreateAuthorization, error) { - claims, err := parseFirewallRuleCheck(checkFlag) - if err != nil { - return firewallRuleCreateAuthorization{}, err - } - ruleDigest, err := firewallRuleDigest(rule) - if err != nil { - return firewallRuleCreateAuthorization{}, err - } - if claims.Version != constant.FirewallRuleCheckVersion || - claims.Provider != rule.Scope.Provider || - claims.ScopeKey != rule.Scope.Key() || - claims.RuleDigest != ruleDigest || - claims.SnapshotRevision != snapshot.Revision || - claims.ManagedRevision != managedRevision { - return firewallRuleCreateAuthorization{}, fmt.Errorf("%w: firewall or managed rules changed", filter.ErrRuleCheckRequired) - } - if claims.Decision != filter.CheckDecisionReady && claims.Decision != filter.CheckDecisionConfirmationRequired { - return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation - } - if !containsFirewallCheckAction(claims.AllowedActions, action) { - return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation - } - - switch action { - case filter.CheckActionCreate, filter.CheckActionCreateAnyway: - if strings.TrimSpace(adoptInstanceKey) != "" { - return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation - } - return firewallRuleCreateAuthorization{Operation: filter.ChangeCreate}, nil - case filter.CheckActionAdopt, filter.CheckActionSelectAdopt: - for _, candidate := range claims.AdoptionCandidates { - if candidate.InstanceKey == adoptInstanceKey && adoptInstanceKey != "" { - locator := candidate.Locator - return firewallRuleCreateAuthorization{Operation: filter.ChangeAdopt, Locator: &locator}, nil - } - } - return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation - default: - return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation - } -} - -func parseFirewallRuleCheck(checkFlag string) (firewallRuleCheckClaims, error) { - parts := strings.Split(strings.TrimSpace(checkFlag), ".") - if len(parts) != 2 || parts[0] == "" || parts[1] == "" { - return firewallRuleCheckClaims{}, filter.ErrRuleCheckRequired - } - payload, err := base64.RawURLEncoding.DecodeString(parts[0]) - if err != nil { - return firewallRuleCheckClaims{}, filter.ErrRuleCheckRequired - } - signature, err := base64.RawURLEncoding.DecodeString(parts[1]) - if err != nil || !hmac.Equal(signature, firewallRuleCheckSignature(payload)) { - return firewallRuleCheckClaims{}, filter.ErrRuleCheckRequired - } - var claims firewallRuleCheckClaims - if err := json.Unmarshal(payload, &claims); err != nil { - return firewallRuleCheckClaims{}, filter.ErrRuleCheckRequired - } - return claims, nil + return firewallCheckFlagCodec().Authorize(checkFlag, action, adoptInstanceKey, rule, snapshot, managedRevision) } -func firewallRuleCheckSignature(payload []byte) []byte { - mac := hmac.New(sha256.New, []byte(global.CONF.Base.EncryptKey+"\x00firewall-rule-check-v1")) - _, _ = mac.Write(payload) - return mac.Sum(nil) -} - -func firewallRuleDigest(rule filter.FirewallRule) (string, error) { - payload, err := json.Marshal(rule) - if err != nil { - return "", err - } - sum := sha256.Sum256(payload) - return hex.EncodeToString(sum[:]), nil -} - -func firewallManagedRevision(rules []model.FirewallRule) (string, error) { - ordered := append([]model.FirewallRule(nil), rules...) - sort.Slice(ordered, func(i, j int) bool { - return ordered[i].UUID < ordered[j].UUID - }) - payload, err := json.Marshal(ordered) - if err != nil { - return "", err - } - sum := sha256.Sum256(payload) - return hex.EncodeToString(sum[:]), nil +func firewallCheckFlagCodec() *filter.CheckFlagCodec { + secret := []byte(global.CONF.Base.EncryptKey + "\x00firewall-rule-check-v1") + return filter.NewCheckFlagCodec(secret, constant.FirewallRuleCheckVersion) } func containsFirewallCheckAction(actions []filter.CheckAction, expected filter.CheckAction) bool { @@ -2678,13 +2259,7 @@ func OperateFirewallPort(oldPorts, newPorts []int) error { } func containsFirewallPort(ports []firewall.PortWhitelist, target firewall.PortWhitelist) bool { - for _, item := range ports { - familyMatches := item.Family == "" || target.Family == "" || item.Family == target.Family - if familyMatches && item.Port == target.Port && item.Protocol == target.Protocol { - return true - } - } - return false + return firewall.ContainsPort(ports, target) } func LoadPanelPort() string { @@ -2949,21 +2524,7 @@ func systemPorts(ports []firewall.PortWhitelist) []dto.FirewallSystemPort { } func excludeFirewallPorts(ports, excluded []firewall.PortWhitelist) []firewall.PortWhitelist { - result := make([]firewall.PortWhitelist, 0, len(ports)) - for _, port := range ports { - exists := false - for _, item := range excluded { - familyMatches := item.Family == "" || port.Family == "" || item.Family == port.Family - if familyMatches && item.Port == port.Port && item.Protocol == port.Protocol { - exists = true - break - } - } - if !exists { - result = append(result, port) - } - } - return result + return firewall.ExcludePorts(ports, excluded) } const ( diff --git a/agent/app/service/firewall_database_sync_test.go b/agent/app/service/firewall_database_sync_test.go deleted file mode 100644 index 40efe7ed014d..000000000000 --- a/agent/app/service/firewall_database_sync_test.go +++ /dev/null @@ -1,41 +0,0 @@ -package service - -import ( - "errors" - "fmt" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -func TestDatabaseSyncPlanMatchesOnceAndSummarizesStates(t *testing.T) { - desired := []databaseSyncDesired[int]{ - {value: 1, item: dto.FirewallRuleSyncItem{SourceUUID: "existing"}}, - {value: 1, item: dto.FirewallRuleSyncItem{SourceUUID: "duplicate"}}, - {value: 2, item: dto.FirewallRuleSyncItem{SourceUUID: "blocked"}, err: errors.New("invalid policy")}, - } - plan := buildDatabaseSyncPlan( - "test", filter.ProviderNftables, desired, []int{1, 3}, - func(value int) string { return fmt.Sprint(value) }, - func(value int) dto.FirewallRuleSyncItem { - return dto.FirewallRuleSyncItem{SourceUUID: "actual"} - }, - ) - - preview := plan.preview() - if preview.Total != 3 || preview.Ready != 1 || preview.Existing != 1 || - preview.Blocked != 1 || preview.Removed != 1 { - t.Fatalf("unexpected preview: %#v", preview) - } - completed := plan.completedResult() - if completed.Total != 3 || completed.Succeeded != 1 || completed.Skipped != 1 || - completed.Failed != 1 || completed.Removed != 1 || len(completed.Errors) != 1 { - t.Fatalf("unexpected completed result: %#v", completed) - } - failed := plan.failedResult(errors.New("reconcile failed")) - if failed.Succeeded != 0 || failed.Skipped != 1 || failed.Failed != 2 || failed.Removed != 0 || - len(failed.Errors) != 2 { - t.Fatalf("unexpected failed result: %#v", failed) - } -} diff --git a/agent/app/service/firewall_docker.go b/agent/app/service/firewall_docker.go index 16807e809d63..72b06d51cfa8 100644 --- a/agent/app/service/firewall_docker.go +++ b/agent/app/service/firewall_docker.go @@ -7,7 +7,6 @@ import ( "fmt" "net/netip" "sort" - "strconv" "strings" "sync" @@ -33,16 +32,7 @@ const ( dockerGuardComposeCreatedBy = "createdBy" ) -type dockerGuardRuntime interface { - Initialize([]docker_guard.Policy) error - Bind() error - Reconcile([]docker_guard.Policy) error - Unbind() error - Cleanup() error - Initialized(string) (bool, error) - Status(string) docker_guard.FamilyStatus - ListPolicies() ([]docker_guard.Policy, error) -} +type dockerGuardRuntime = docker_guard.Runtime type DockerPortGuardService struct { policies repo.IDockerPortGuardRepo @@ -52,20 +42,12 @@ type DockerPortGuardService struct { version func(string) string } -type normalizedDockerGuardPolicy struct { - Family string - HostIP string - HostPort uint16 - Protocol string - Mode string -} - var ( dockerPortGuardServiceMu sync.Mutex dockerPortGuardSyncMu sync.RWMutex dockerPortGuardSyncErr error - ErrDockerGuardInvalid = errors.New("invalid Docker port guard request") - ErrDockerUnavailable = errors.New("Docker is unavailable") + ErrDockerGuardInvalid = docker_guard.ErrInvalidPolicy + ErrDockerUnavailable = docker.ErrUnavailable ) type IDockerPortGuardService interface { @@ -156,7 +138,7 @@ func matchDockerGuardPolicies( if !ok { continue } - endpoints[i].PolicyUUID, endpoints[i].Mode, endpoints[i].Sources = policy.UUID, policy.Mode, decodeGuardSources(policy.Sources) + endpoints[i].PolicyUUID, endpoints[i].Mode, endpoints[i].Sources = policy.UUID, policy.Mode, docker_guard.DecodeSources(policy.Sources) endpoints[i].Description = policy.Description endpoints[i].Effective = (policy.Family == docker_guard.FamilyIPv4 && base.IPv4.Effective) || (policy.Family == docker_guard.FamilyIPv6 && base.IPv6.Effective) delete(byEndpoint, key) @@ -165,7 +147,7 @@ func matchDockerGuardPolicies( for _, policy := range byEndpoint { orphanPolicies = append(orphanPolicies, dto.DockerPortGuardEndpoint{ Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, - PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources), Description: policy.Description, + PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources), Description: policy.Description, }) } return endpoints, orphanPolicies @@ -224,7 +206,7 @@ func (s *DockerPortGuardService) Operate(ctx context.Context, request dto.Docker func (s *DockerPortGuardService) DeletePolicies(ctx context.Context, request dto.DockerPortGuardPolicyBatchDelete) error { dockerPortGuardServiceMu.Lock() defer dockerPortGuardServiceMu.Unlock() - uuids, err := normalizeDockerGuardPolicyUUIDs(request.UUIDs) + uuids, err := docker_guard.NormalizePolicyUUIDs(request.UUIDs) if err != nil { return err } @@ -234,33 +216,16 @@ func (s *DockerPortGuardService) DeletePolicies(ctx context.Context, request dto return s.reconcileLocked(ctx) } -func normalizeDockerGuardPolicyUUIDs(values []string) ([]string, error) { - uuids := make([]string, 0, len(values)) - seen := make(map[string]struct{}, len(values)) - for _, policyUUID := range values { - policyUUID = strings.TrimSpace(policyUUID) - if policyUUID == "" { - return nil, fmt.Errorf("%w: policy UUID cannot be empty", ErrDockerGuardInvalid) - } - if _, exists := seen[policyUUID]; exists { - continue - } - seen[policyUUID] = struct{}{} - uuids = append(uuids, policyUUID) - } - if len(uuids) == 0 { - return nil, fmt.Errorf("%w: policy UUIDs cannot be empty", ErrDockerGuardInvalid) - } - return uuids, nil -} - func (s *DockerPortGuardService) UpsertPolicies(ctx context.Context, request dto.DockerPortGuardPolicyBatch) error { dockerPortGuardServiceMu.Lock() defer dockerPortGuardServiceMu.Unlock() policies := make([]model.DockerPortGuardPolicy, 0, len(request.Endpoints)) seen := make(map[string]struct{}, len(request.Endpoints)) for _, endpoint := range request.Endpoints { - normalized, sources, err := normalizeGuardPolicy(endpoint.Family, endpoint.HostIP, endpoint.HostPort, endpoint.Protocol, request.Mode, request.Sources) + normalized, err := docker_guard.NormalizePolicy(docker_guard.Policy{ + Family: endpoint.Family, HostIP: endpoint.HostIP, HostPort: endpoint.HostPort, + Protocol: endpoint.Protocol, Mode: request.Mode, Sources: request.Sources, + }) if err != nil { return err } @@ -269,7 +234,7 @@ func (s *DockerPortGuardService) UpsertPolicies(ctx context.Context, request dto continue } seen[key] = struct{}{} - encoded, _ := json.Marshal(sources) + encoded, _ := json.Marshal(normalized.Sources) policies = append(policies, model.DockerPortGuardPolicy{ UUID: uuid.NewString(), Family: normalized.Family, HostIP: normalized.HostIP, HostPort: normalized.HostPort, Protocol: normalized.Protocol, Mode: normalized.Mode, @@ -348,92 +313,14 @@ func dockerGuardPoliciesFromModels(policies []model.DockerPortGuardPolicy) []doc func dockerGuardPolicyFromModel(policy model.DockerPortGuardPolicy) docker_guard.Policy { return docker_guard.Policy{ UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, - Protocol: policy.Protocol, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources), - } -} - -func dockerGuardPolicySyncKey(policy docker_guard.Policy) string { - mode := policy.Mode - if mode == docker_guard.ModeAllow && len(policy.Sources) == 0 { - mode = docker_guard.ModeAll - } - sources := append([]string(nil), policy.Sources...) - sort.Strings(sources) - return strings.Join([]string{ - policy.UUID, policy.Family, canonicalGuardHost(policy.HostIP), strconv.Itoa(int(policy.HostPort)), - policy.Protocol, mode, strings.Join(sources, ","), - }, "\x00") -} - -func canonicalGuardHost(value string) string { - if address, err := netip.ParseAddr(value); err == nil { - return address.String() - } - return value -} - -func verifyDockerGuardRuleSync(runtime dockerGuardRuntime, desired []docker_guard.Policy) error { - actual, err := runtime.ListPolicies() - if err != nil { - return fmt.Errorf("verify synchronized Docker firewall policies: %w", err) + Protocol: policy.Protocol, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources), } - if !databaseSyncStatesEqual(actual, desired, dockerGuardPolicySyncKey) { - return fmt.Errorf("verify synchronized Docker firewall policies: target policies do not match the database") - } - return nil -} - -func reconcileDockerGuardSyncTarget(backend string, policies []docker_guard.Policy, runtime dockerGuardRuntime) error { - families := make(map[string]struct{}, len(policies)) - needsInitialize, needsBind := false, false - for _, policy := range policies { - families[policy.Family] = struct{}{} - } - if len(families) == 0 { - initialized := false - for _, family := range []string{docker_guard.FamilyIPv4, docker_guard.FamilyIPv6} { - status := runtime.Status(family) - if status.Reason == docker_guard.ReasonInspectFailed { - return fmt.Errorf("inspect Docker firewall target %s for %s failed", backend, family) - } - initialized = initialized || status.Initialized - } - if initialized { - return runtime.Reconcile(nil) - } - return nil - } - for family := range families { - status := runtime.Status(family) - needsInitialize = needsInitialize || !status.Initialized - needsBind = needsBind || !status.Bound || !status.Effective - } - var err error - if needsInitialize { - err = runtime.Initialize(policies) - } else { - if needsBind { - err = runtime.Bind() - } - if err == nil { - err = runtime.Reconcile(policies) - } - } - if err != nil { - return err - } - for family := range families { - if !runtime.Status(family).Effective { - return fmt.Errorf("Docker firewall target %s is not effective for %s", backend, family) - } - } - return nil } func dockerGuardRuleSyncDTO(policy model.DockerPortGuardPolicy) *dto.DockerPortGuardEndpoint { return &dto.DockerPortGuardEndpoint{ Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, - PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources), Description: policy.Description, + PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources), Description: policy.Description, } } @@ -551,7 +438,7 @@ func (s *DockerPortGuardService) runtimePolicies(ctx context.Context) ([]docker_ } policies := make([]docker_guard.Policy, 0, len(stored)) for _, policy := range stored { - policies = append(policies, docker_guard.Policy{UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources)}) + policies = append(policies, docker_guard.Policy{UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources)}) } return policies, nil } @@ -577,10 +464,7 @@ func (s *DockerPortGuardService) guardRuntime(backend string) dockerGuardRuntime if s.runtime != nil { return s.runtime } - if backend == constant.FirewallProviderNftables { - return docker_guard.NewNftablesManager() - } - return docker_guard.NewManager() + return docker_guard.NewRuntime(backend) } func (s *DockerPortGuardService) runtimeForDocker(ctx context.Context) (dockerGuardRuntime, string, error) { @@ -668,65 +552,12 @@ func discoverDockerEndpoints(ctx context.Context, cli *client.Client) ([]dto.Doc return endpoints, nil } -func normalizeGuardPolicy(family, hostIP string, hostPort uint16, protocol, mode string, sources []string) (normalizedDockerGuardPolicy, []string, error) { - family, hostIP, protocol, mode = strings.ToLower(strings.TrimSpace(family)), strings.TrimSpace(hostIP), strings.ToLower(strings.TrimSpace(protocol)), strings.ToLower(strings.TrimSpace(mode)) - if hostPort == 0 || (protocol != "tcp" && protocol != "udp") || (family != docker_guard.FamilyIPv4 && family != docker_guard.FamilyIPv6) || (mode != docker_guard.ModeAll && mode != docker_guard.ModeSources && mode != docker_guard.ModeAllow) { - return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: invalid policy fields", ErrDockerGuardInvalid) - } - addr, err := netip.ParseAddr(hostIP) - if err != nil || (family == docker_guard.FamilyIPv4) != addr.Is4() { - return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: host IP does not match address family", ErrDockerGuardInvalid) - } - normalizedSources := make([]string, 0, len(sources)) - seen := map[string]struct{}{} - for _, source := range sources { - source = strings.TrimSpace(source) - if source == "" { - continue - } - prefix, err := netip.ParsePrefix(source) - if err != nil { - if sourceAddr, addrErr := netip.ParseAddr(source); addrErr == nil { - bits := 128 - if sourceAddr.Is4() { - bits = 32 - } - prefix = netip.PrefixFrom(sourceAddr, bits) - } else { - return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: invalid source address %q", ErrDockerGuardInvalid, source) - } - } - if (family == docker_guard.FamilyIPv4) != prefix.Addr().Is4() { - return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: source %q does not match address family", ErrDockerGuardInvalid, source) - } - canonical := prefix.Masked().String() - if _, ok := seen[canonical]; !ok { - seen[canonical] = struct{}{} - normalizedSources = append(normalizedSources, canonical) - } - } - if mode == docker_guard.ModeSources && len(normalizedSources) == 0 { - return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: deny_sources requires at least one source", ErrDockerGuardInvalid) - } - if mode == docker_guard.ModeAll { - normalizedSources = []string{} - } - sort.Strings(normalizedSources) - return normalizedDockerGuardPolicy{Family: family, HostIP: hostIP, HostPort: hostPort, Protocol: protocol, Mode: mode}, normalizedSources, nil -} - -func decodeGuardSources(value string) []string { - result := []string{} - _ = json.Unmarshal([]byte(value), &result) - return result -} - func dockerGuardPolicyEndpoints(policies []model.DockerPortGuardPolicy) []dto.DockerPortGuardEndpoint { endpoints := make([]dto.DockerPortGuardEndpoint, 0, len(policies)) for _, policy := range policies { endpoints = append(endpoints, dto.DockerPortGuardEndpoint{ Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, - PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources), + PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources), Description: policy.Description, }) } diff --git a/agent/app/service/firewall_docker_test.go b/agent/app/service/firewall_docker_test.go deleted file mode 100644 index 157064c6ff65..000000000000 --- a/agent/app/service/firewall_docker_test.go +++ /dev/null @@ -1,572 +0,0 @@ -package service - -import ( - "context" - "errors" - "fmt" - "reflect" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/global" - agenti18n "github.com/1Panel-dev/1Panel/agent/i18n" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - "github.com/docker/docker/client" - "github.com/glebarez/sqlite" - "github.com/google/uuid" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -type persistentDockerGuardRuntime struct { - initialized bool - initialize int - reconcile int - bind int - unbind int - policies []docker_guard.Policy - statuses map[string]docker_guard.FamilyStatus -} - -func (r *persistentDockerGuardRuntime) Initialize(policies []docker_guard.Policy) error { - r.initialize++ - r.initialized = true - r.policies = append([]docker_guard.Policy(nil), policies...) - r.markFamiliesEffective(policies) - return nil -} - -func (r *persistentDockerGuardRuntime) Bind() error { - r.bind++ - r.initialized = true - for family, status := range r.statuses { - status.Initialized, status.Bound, status.Effective = true, true, true - r.statuses[family] = status - } - return nil -} - -func (r *persistentDockerGuardRuntime) Reconcile(policies []docker_guard.Policy) error { - r.reconcile++ - r.policies = append([]docker_guard.Policy(nil), policies...) - return nil -} - -func (r *persistentDockerGuardRuntime) markFamiliesEffective(policies []docker_guard.Policy) { - if r.statuses == nil { - return - } - for _, policy := range policies { - r.statuses[policy.Family] = docker_guard.FamilyStatus{Initialized: true, Bound: true, Effective: true} - } -} - -func (r *persistentDockerGuardRuntime) Unbind() error { - r.unbind++ - return nil -} - -func (r *persistentDockerGuardRuntime) Cleanup() error { return nil } - -func (r *persistentDockerGuardRuntime) Initialized(string) (bool, error) { - return r.initialized, nil -} - -func (r *persistentDockerGuardRuntime) Status(family string) docker_guard.FamilyStatus { - if r.statuses != nil { - return r.statuses[family] - } - return docker_guard.FamilyStatus{Initialized: r.initialized, Bound: r.initialized, Effective: r.initialized} -} - -func (r *persistentDockerGuardRuntime) ListPolicies() ([]docker_guard.Policy, error) { - return append([]docker_guard.Policy(nil), r.policies...), nil -} - -type persistentDockerGuardPolicies struct{ items []model.DockerPortGuardPolicy } - -func (r *persistentDockerGuardPolicies) List(context.Context) ([]model.DockerPortGuardPolicy, error) { - return append([]model.DockerPortGuardPolicy(nil), r.items...), nil -} - -func (r *persistentDockerGuardPolicies) DeleteBatch(context.Context, []string) error { return nil } - -func (r *persistentDockerGuardPolicies) UpsertBatch(context.Context, []model.DockerPortGuardPolicy) error { - return nil -} - -func TestDockerGuardOverviewLocalizesUnavailableDocker(t *testing.T) { - agenti18n.Init() - service := &DockerPortGuardService{ - policies: &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{{ - UUID: "orphan-policy", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, - Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]", - }}}, - runtime: &persistentDockerGuardRuntime{}, - version: func(string) string { return "1.8.10" }, - client: func() (*client.Client, error) { - return nil, errors.New("Cannot connect to the Docker daemon at unix:///var/run/docker.sock") - }, - } - overview, err := service.LoadOverview(context.Background()) - if err != nil { - t.Fatalf("load overview: %v", err) - } - if overview.Base.Message != agenti18n.Get("ErrDockerFailed") { - t.Fatalf("message = %q, want localized Docker failure", overview.Base.Message) - } - if overview.Base.Version != "1.8.10" { - t.Fatalf("version = %q, want 1.8.10", overview.Base.Version) - } - if len(overview.Containers) != 0 || len(overview.OrphanPolicies) != 1 { - t.Fatalf("overview = %#v, want persisted policy returned separately from Docker containers", overview) - } - if overview.OrphanPolicies[0].PolicyUUID != "orphan-policy" { - t.Fatalf("orphan policy = %#v", overview.OrphanPolicies[0]) - } -} - -func TestDockerGuardRuntimeStatusAggregatesAvailableFamilies(t *testing.T) { - runtime := &persistentDockerGuardRuntime{statuses: map[string]docker_guard.FamilyStatus{ - docker_guard.FamilyIPv6: {Initialized: true, Bound: true, Effective: true}, - }} - base := (&DockerPortGuardService{}).runtimeStatus(runtime, constant.FirewallProviderNftables) - if !base.Initialized || !base.Bound || base.IPv4.Initialized || !base.IPv6.Initialized { - t.Fatalf("unexpected aggregate Docker guard status: %#v", base) - } -} - -func TestDockerGuardRuntimeStatusReportsMissingBackend(t *testing.T) { - runtime := &persistentDockerGuardRuntime{statuses: map[string]docker_guard.FamilyStatus{ - docker_guard.FamilyIPv4: {Reason: docker_guard.ReasonCommandMissing}, - docker_guard.FamilyIPv6: {Reason: docker_guard.ReasonCommandMissing}, - }} - base := (&DockerPortGuardService{}).runtimeStatus(runtime, constant.FirewallProviderIptables) - if base.IsExist { - t.Fatalf("missing Docker firewall backend reported as installed: %#v", base) - } -} - -func TestMatchDockerGuardPoliciesReturnsUnmatchedDatabaseRules(t *testing.T) { - policies := []model.DockerPortGuardPolicy{ - {UUID: "matched", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]"}, - {UUID: "orphan", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 9090, Protocol: "tcp", Mode: docker_guard.ModeSources, Sources: `["203.0.113.0/24"]`}, - } - endpoints := []dto.DockerPortGuardEndpoint{{ - Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, Protocol: "tcp", ContainerID: "container-1", - }} - matched, orphanPolicies := matchDockerGuardPolicies(dto.DockerPortGuardBase{}, policies, endpoints) - if len(matched) != 1 || matched[0].PolicyUUID != "matched" { - t.Fatalf("matched endpoints = %#v", matched) - } - if len(orphanPolicies) != 1 || orphanPolicies[0].PolicyUUID != "orphan" || orphanPolicies[0].HostPort != 9090 { - t.Fatalf("orphan policies = %#v", orphanPolicies) - } -} - -func TestDockerGuardRuleSyncInitializesTargetWithPersistedPolicies(t *testing.T) { - setupDockerGuardSettingsDB(t) - selectDockerGuardBackend(t, constant.FirewallProviderNftables) - target := &persistentDockerGuardRuntime{} - policies := &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{{ - UUID: "policy-1", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, - Protocol: "tcp", Mode: docker_guard.ModeSources, Sources: `["203.0.113.0/24"]`, - }}} - service := &DockerPortGuardService{ - policies: policies, - runtimeForBackend: func(string) dockerGuardRuntime { return target }, - } - request := dto.FirewallRuleSyncRequest{Subsystem: "docker", TargetProvider: filter.ProviderNftables} - preview, err := service.previewRuleSync(context.Background(), request) - if err != nil { - t.Fatal(err) - } - if preview.Total != 1 || preview.Ready != 1 || preview.TargetProvider != filter.ProviderNftables || preview.Items[0].DockerRule == nil { - t.Fatalf("unexpected preview: %#v", preview) - } - result, err := service.syncRules(context.Background(), request) - if err != nil { - t.Fatal(err) - } - if result.Succeeded != 1 || result.Failed != 0 || target.initialize != 1 || len(target.policies) != 1 { - t.Fatalf("unexpected sync result=%#v target=%#v", result, target) - } - retry, err := service.previewRuleSync(context.Background(), request) - if err != nil { - t.Fatal(err) - } - if retry.Ready != 0 || retry.Existing != 1 || retry.Removed != 0 { - t.Fatalf("synchronized Docker policy was not recognized: %#v", retry) - } - retryResult, err := service.syncRules(context.Background(), request) - if err != nil { - t.Fatal(err) - } - if retryResult.Succeeded != 0 || retryResult.Skipped != 1 || retryResult.Removed != 0 { - t.Fatalf("existing Docker policy was not counted as skipped: %#v", retryResult) - } -} - -func TestDockerGuardRuleSyncReconcilesInitializedEffectiveTarget(t *testing.T) { - setupDockerGuardSettingsDB(t) - selectDockerGuardBackend(t, constant.FirewallProviderNftables) - policy := model.DockerPortGuardPolicy{ - UUID: "policy-1", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, - Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]", - } - target := &persistentDockerGuardRuntime{initialized: true} - service := &DockerPortGuardService{ - policies: &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{policy}}, - runtimeForBackend: func(string) dockerGuardRuntime { return target }, - } - result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "docker", TargetProvider: filter.ProviderNftables, - }) - if err != nil { - t.Fatal(err) - } - if result.Succeeded != 1 || result.Skipped != 0 || target.reconcile != 1 || len(target.policies) != 1 { - t.Fatalf("initialized target skipped runtime reconciliation: result=%#v target=%#v", result, target) - } -} - -func TestDockerGuardRuleSyncRebindsInitializedIneffectiveTarget(t *testing.T) { - setupDockerGuardSettingsDB(t) - selectDockerGuardBackend(t, constant.FirewallProviderNftables) - policy := model.DockerPortGuardPolicy{ - UUID: "policy-1", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, - Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]", - } - target := &persistentDockerGuardRuntime{ - initialized: true, - statuses: map[string]docker_guard.FamilyStatus{ - docker_guard.FamilyIPv4: {Initialized: true}, - }, - } - service := &DockerPortGuardService{ - policies: &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{policy}}, - runtimeForBackend: func(string) dockerGuardRuntime { return target }, - } - result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "docker", TargetProvider: filter.ProviderNftables, - }) - if err != nil { - t.Fatal(err) - } - if result.Succeeded != 1 || target.bind != 1 || target.reconcile != 1 || - !target.Status(docker_guard.FamilyIPv4).Effective { - t.Fatalf("initialized target was not rebound and reconciled: result=%#v target=%#v", result, target) - } -} - -func TestDockerGuardRuleSyncClearsInitializedTargetWhenDatabaseIsEmpty(t *testing.T) { - setupDockerGuardSettingsDB(t) - selectDockerGuardBackend(t, constant.FirewallProviderNftables) - target := &persistentDockerGuardRuntime{ - initialized: true, - policies: []docker_guard.Policy{{ - UUID: "stale-policy", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, - Protocol: "tcp", Mode: docker_guard.ModeAll, - }}, - } - service := &DockerPortGuardService{ - policies: &persistentDockerGuardPolicies{}, - runtimeForBackend: func(string) dockerGuardRuntime { return target }, - } - preview, err := service.previewRuleSync(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "docker", TargetProvider: filter.ProviderNftables, - }) - if err != nil { - t.Fatal(err) - } - if preview.Total != 0 || preview.Removed != 1 || len(preview.Items) != 1 || preview.Items[0].Status != "remove" { - t.Fatalf("stale runtime policy was not included in preview: %#v", preview) - } - result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "docker", TargetProvider: filter.ProviderNftables, - }) - if err != nil { - t.Fatal(err) - } - if result.Total != 0 || result.Removed != 1 || target.initialize != 0 || target.reconcile != 1 || len(target.policies) != 0 { - t.Fatalf("empty database did not clear initialized target: result=%#v target=%#v", result, target) - } -} - -func TestDockerGuardRuleSyncLeavesUninitializedTargetEmpty(t *testing.T) { - setupDockerGuardSettingsDB(t) - selectDockerGuardBackend(t, constant.FirewallProviderNftables) - target := &persistentDockerGuardRuntime{} - service := &DockerPortGuardService{ - policies: &persistentDockerGuardPolicies{}, - runtimeForBackend: func(string) dockerGuardRuntime { return target }, - } - result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "docker", TargetProvider: filter.ProviderNftables, - }) - if err != nil { - t.Fatal(err) - } - if result.Total != 0 || target.initialize != 0 || target.reconcile != 0 { - t.Fatalf("empty database initialized an unused target: result=%#v target=%#v", result, target) - } -} - -func TestDockerGuardRuleSyncRejectsUnselectedTarget(t *testing.T) { - setupDockerGuardSettingsDB(t) - selectDockerGuardBackend(t, constant.FirewallProviderIptables) - target := &persistentDockerGuardRuntime{} - service := &DockerPortGuardService{ - policies: &persistentDockerGuardPolicies{}, - runtimeForBackend: func(string) dockerGuardRuntime { return target }, - } - request := dto.FirewallRuleSyncRequest{Subsystem: "docker", TargetProvider: filter.ProviderNftables} - if _, err := service.previewRuleSync(context.Background(), request); !errors.Is(err, filter.ErrProviderUnavailable) { - t.Fatalf("preview error = %v, want provider unavailable", err) - } - if _, err := service.syncRules(context.Background(), request); !errors.Is(err, filter.ErrProviderUnavailable) { - t.Fatalf("sync error = %v, want provider unavailable", err) - } - if target.initialize != 0 || target.bind != 0 || target.reconcile != 0 { - t.Fatalf("unselected target was modified: %#v", target) - } -} - -func selectDockerGuardBackend(t *testing.T, backend string) { - t.Helper() - if err := settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, backend); err != nil { - t.Fatal(err) - } -} - -func setupDockerGuardSettingsDB(t *testing.T) { - t.Helper() - previousDB := global.DB - dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString()) - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) - if err != nil { - t.Fatalf("open settings database: %v", err) - } - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatalf("migrate settings database: %v", err) - } - global.DB = db - t.Cleanup(func() { global.DB = previousDB }) -} - -func TestNormalizeDockerPortGuardPolicy(t *testing.T) { - policy, sources, err := normalizeGuardPolicy("ipv4", "0.0.0.0", 8080, "TCP", "deny_sources", []string{"203.0.113.10", "192.0.2.7/24", "203.0.113.10/32"}) - if err != nil { - t.Fatal(err) - } - if policy.Protocol != "tcp" { - t.Fatalf("protocol was not normalized: %#v", policy) - } - want := []string{"192.0.2.0/24", "203.0.113.10/32"} - if !reflect.DeepEqual(sources, want) { - t.Fatalf("sources = %#v, want %#v", sources, want) - } -} - -func TestNormalizeDockerPortGuardPolicyRejectsInvalidSources(t *testing.T) { - for _, test := range []struct { - name string - family string - hostIP string - mode string - sources []string - }{ - {name: "empty deny sources", family: "ipv4", hostIP: "0.0.0.0", mode: "deny_sources"}, - {name: "mixed source family", family: "ipv4", hostIP: "0.0.0.0", mode: "deny_sources", sources: []string{"2001:db8::/64"}}, - {name: "mixed host family", family: "ipv6", hostIP: "0.0.0.0", mode: "deny_all"}, - } { - t.Run(test.name, func(t *testing.T) { - if _, _, err := normalizeGuardPolicy(test.family, test.hostIP, 80, "tcp", test.mode, test.sources); !errors.Is(err, ErrDockerGuardInvalid) { - t.Fatalf("expected typed validation error, got %v", err) - } - }) - } -} - -func TestNormalizeDockerPortGuardPolicyAllowsEmptyAllowList(t *testing.T) { - policy, sources, err := normalizeGuardPolicy("ipv4", "0.0.0.0", 5432, "tcp", "allow_sources", nil) - if err != nil { - t.Fatal(err) - } - if policy.Mode != "allow_sources" || len(sources) != 0 { - t.Fatalf("policy = %#v, sources = %#v", policy, sources) - } -} - -func TestDockerGuardEndpointKeyIncludesAddressFamilyAndProtocol(t *testing.T) { - first := guardEndpointKey("ipv4", "0.0.0.0", 53, "udp") - if first == guardEndpointKey("ipv6", "::", 53, "udp") || first == guardEndpointKey("ipv4", "0.0.0.0", 53, "tcp") { - t.Fatal("endpoint identity collapsed distinct endpoint dimensions") - } -} - -func TestNormalizeDockerGuardPolicyUUIDs(t *testing.T) { - got, err := normalizeDockerGuardPolicyUUIDs([]string{" first ", "second", "first"}) - if err != nil { - t.Fatal(err) - } - if want := []string{"first", "second"}; !reflect.DeepEqual(got, want) { - t.Fatalf("UUIDs = %#v, want %#v", got, want) - } - if _, err := normalizeDockerGuardPolicyUUIDs([]string{""}); err == nil { - t.Fatal("expected empty UUID to be rejected") - } -} - -func TestGroupDockerGuardContainersMergesCompatiblePorts(t *testing.T) { - endpoints := []dto.DockerPortGuardEndpoint{ - {Family: "ipv4", HostIP: "0.0.0.0", HostPort: 8001, Protocol: "tcp", ContainerID: "container-1", ContainerName: "demo", ContainerPort: 81, Mode: "deny_all", PolicyUUID: "policy-2", Effective: true, Sources: []string{}}, - {Family: "ipv4", HostIP: "0.0.0.0", HostPort: 8000, Protocol: "tcp", ContainerID: "container-1", ContainerName: "demo", ContainerPort: 80, Mode: "deny_all", PolicyUUID: "policy-1", Effective: true, Sources: []string{}}, - {Family: "ipv4", HostIP: "0.0.0.0", HostPort: 8002, Protocol: "tcp", ContainerID: "container-1", ContainerName: "demo", ContainerPort: 82, Mode: "allow_sources", PolicyUUID: "policy-3", Effective: true, Sources: []string{"192.0.2.1/32"}}, - } - containers := groupDockerGuardContainers(endpoints) - if len(containers) != 1 || len(containers[0].PortGroups) != 2 { - t.Fatalf("containers = %#v, want one container with two port groups", containers) - } - foundRange := false - for _, group := range containers[0].PortGroups { - foundRange = foundRange || group.Label == "0.0.0.0:8000-8001/tcp" - } - if !foundRange { - t.Fatalf("port groups = %#v, expected merged range", containers[0].PortGroups) - } - for _, group := range containers[0].PortGroups { - if group.Label == "0.0.0.0:8000-8001/tcp" && len(group.Endpoints) != 2 { - t.Fatalf("merged group endpoints = %#v, want 2 endpoints", group.Endpoints) - } - } -} - -func TestMarkDockerGuardReconcileFailureOnlyAffectsFailedFamily(t *testing.T) { - base := dto.DockerPortGuardBase{ - IPv4: dto.DockerPortGuardFamilyStatus{State: docker_guard.StatusEffective, Initialized: true, Bound: true, Effective: true}, - IPv6: dto.DockerPortGuardFamilyStatus{State: docker_guard.StatusEffective, Initialized: true, Bound: true, Effective: true}, - } - markDockerGuardReconcileFailure(&base, &docker_guard.FamilyError{Family: docker_guard.FamilyIPv6, Err: errors.New("restore failed")}) - if !base.IPv4.Effective || base.IPv4.State != docker_guard.StatusEffective { - t.Fatalf("IPv4 status changed unexpectedly: %#v", base.IPv4) - } - if base.IPv6.Effective || base.IPv6.State != docker_guard.StatusNotEffective || base.IPv6.Reason != docker_guard.ReasonInspectFailed { - t.Fatalf("IPv6 status = %#v, want not effective", base.IPv6) - } -} - -func TestMarkDockerGuardIPv4FailureAlsoMarksUnattemptedIPv6(t *testing.T) { - base := dto.DockerPortGuardBase{ - IPv4: dto.DockerPortGuardFamilyStatus{State: docker_guard.StatusEffective, Initialized: true, Bound: true, Effective: true}, - IPv6: dto.DockerPortGuardFamilyStatus{State: docker_guard.StatusEffective, Initialized: true, Bound: true, Effective: true}, - } - markDockerGuardReconcileFailure(&base, &docker_guard.FamilyError{Family: docker_guard.FamilyIPv4, Err: errors.New("restore failed")}) - if base.IPv4.Effective || base.IPv6.Effective { - t.Fatalf("statuses = IPv4 %#v, IPv6 %#v; both must be not effective", base.IPv4, base.IPv6) - } -} - -func TestDockerGuardReconcileErrorState(t *testing.T) { - t.Cleanup(func() { recordDockerPortGuardReconcileError(nil) }) - want := errors.New("restore failed") - recordDockerPortGuardReconcileError(want) - if got := lastDockerPortGuardReconcileError(); !errors.Is(got, want) { - t.Fatalf("last reconcile error = %v, want %v", got, want) - } - recordDockerPortGuardReconcileError(nil) - if got := lastDockerPortGuardReconcileError(); got != nil { - t.Fatalf("last reconcile error = %v, want nil", got) - } -} - -func TestDockerFirewallDisplayName(t *testing.T) { - for backend, want := range map[string]string{ - "iptables": "iptables-docker", - "nftables": "nftables-docker", - "": "iptables-docker", - } { - if got := dockerFirewallDisplayName(backend); got != want { - t.Fatalf("dockerFirewallDisplayName(%q) = %q, want %q", backend, got, want) - } - } -} - -func TestDockerGuardRuntimeMatchesDockerBackend(t *testing.T) { - service := &DockerPortGuardService{} - if _, ok := service.guardRuntime("iptables").(*docker_guard.Manager); !ok { - t.Fatal("iptables backend did not select the iptables Docker guard runtime") - } - if _, ok := service.guardRuntime("nftables").(*docker_guard.NftablesManager); !ok { - t.Fatal("nftables backend did not select the nftables Docker guard runtime") - } -} - -func TestDockerGuardInitializeAndUnbindPersistStatus(t *testing.T) { - setupDockerGuardSettingsDB(t) - runtime := &persistentDockerGuardRuntime{} - service := &DockerPortGuardService{ - policies: &persistentDockerGuardPolicies{}, - runtime: runtime, - } - if err := service.Operate(context.Background(), dto.DockerPortGuardOperation{Operation: "initialize"}); err != nil { - t.Fatalf("initialize Docker port guard: %v", err) - } - status, err := settingRepo.GetValueByKey(constant.FirewallDockerPortGuardStatusKey) - if err != nil || status != constant.StatusEnable { - t.Fatalf("persisted status = %q, %v; want %q", status, err, constant.StatusEnable) - } - if runtime.initialize != 1 { - t.Fatalf("initialize calls = %d, want 1", runtime.initialize) - } - - if err := service.Operate(context.Background(), dto.DockerPortGuardOperation{Operation: "unbind"}); err != nil { - t.Fatalf("unbind Docker port guard: %v", err) - } - status, err = settingRepo.GetValueByKey(constant.FirewallDockerPortGuardStatusKey) - if err != nil || status != constant.StatusDisable { - t.Fatalf("persisted status = %q, %v; want %q", status, err, constant.StatusDisable) - } -} - -func TestDockerGuardReconcileRestoresPersistedInitialization(t *testing.T) { - setupDockerGuardSettingsDB(t) - if err := settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusEnable); err != nil { - t.Fatal(err) - } - runtime := &persistentDockerGuardRuntime{} - service := &DockerPortGuardService{ - policies: &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{{ - UUID: "policy-1", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", - HostPort: 8080, Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]", - }}}, - runtime: runtime, - } - if err := service.Reconcile(context.Background()); err != nil { - t.Fatalf("restore Docker port guard: %v", err) - } - if runtime.initialize != 1 || !runtime.initialized { - t.Fatalf("runtime was not initialized: %#v", runtime) - } - if len(runtime.policies) != 1 || runtime.policies[0].UUID != "policy-1" { - t.Fatalf("restored policies = %#v", runtime.policies) - } -} - -func TestDockerGuardReconcileLeavesDisabledGuardUninitialized(t *testing.T) { - setupDockerGuardSettingsDB(t) - if err := settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusDisable); err != nil { - t.Fatal(err) - } - runtime := &persistentDockerGuardRuntime{} - service := &DockerPortGuardService{policies: &persistentDockerGuardPolicies{}, runtime: runtime} - if err := service.Reconcile(context.Background()); err != nil { - t.Fatalf("reconcile disabled Docker port guard: %v", err) - } - if runtime.initialize != 0 || runtime.initialized { - t.Fatalf("disabled guard was initialized: %#v", runtime) - } -} diff --git a/agent/app/service/firewall_service_test.go b/agent/app/service/firewall_service_test.go deleted file mode 100644 index 831bcfe7e7f7..000000000000 --- a/agent/app/service/firewall_service_test.go +++ /dev/null @@ -1,2038 +0,0 @@ -package service - -import ( - "context" - "errors" - "fmt" - "strings" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/app/repo" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/utils/firewall" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - filterfirewalld "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/firewalld" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" - "github.com/glebarez/sqlite" - "github.com/go-playground/validator/v10" - "github.com/google/uuid" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -func TestFirewallBaseInfoDistinguishesMissingAndConflictingProviders(t *testing.T) { - wantErr := errors.New("firewalld and ufw conflict") - for _, test := range []struct { - name string - installed []string - exists bool - message string - }{ - {name: "missing", installed: nil}, - {name: "conflict", installed: []string{constant.FirewallProviderFirewalld, constant.FirewallProviderUFW}, exists: true, message: wantErr.Error()}, - } { - t.Run(test.name, func(t *testing.T) { - service := &FirewallService{ - baseClient: func() (lifecycle.Client, error) { return nil, wantErr }, - installedProviders: func() []string { - return append([]string(nil), test.installed...) - }, - } - base, err := service.LoadBaseInfo("base") - if err != nil { - t.Fatal(err) - } - if base.IsExist != test.exists || base.Message != test.message { - t.Fatalf("unexpected firewall failure status: %#v", base) - } - }) - } -} - -func TestFirewallResetCleansDirectBackendAndKeepsStoredRules(t *testing.T) { - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - stored, err := model.FirewallRuleFromDomain(executorTestAddressRule("172.16.10.111")) - if err != nil { - t.Fatal(err) - } - stored.UUID = "reset-direct-rule" - stored.Origin = constant.FirewallRuleOriginCreated - stored.Owner = constant.FirewallRuleSourceUser - if err := ruleRepo.Create(context.Background(), &stored); err != nil { - t.Fatal(err) - } - cleaned := "" - service := &FirewallService{ - rules: ruleRepo, - selectedProvider: func(context.Context) (filter.Provider, error) { - return filter.ProviderIptables, nil - }, - cleanupBackend: func(provider string) error { - cleaned = provider - return nil - }, - } - result, err := service.Reset(context.Background(), dto.FirewallRuleReset{}) - if err != nil { - t.Fatal(err) - } - if cleaned != string(filter.ProviderIptables) || result.Removed != 1 { - t.Fatalf("unexpected reset result: cleaned=%q result=%#v", cleaned, result) - } - remaining, err := ruleRepo.List(context.Background()) - if err != nil { - t.Fatal(err) - } - if len(remaining) != 1 || remaining[0].UUID != stored.UUID { - t.Fatalf("reset removed provider-neutral database policy: %#v", remaining) - } -} - -func TestFirewallResetCleansInactiveDirectBackendWithoutChangingSelectedBackendState(t *testing.T) { - service := &FirewallService{ - rules: repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)), - selectedProvider: func(context.Context) (filter.Provider, error) { - return filter.ProviderNftables, nil - }, - cleanupBackend: func(string) error { - t.Fatal("inactive source cleanup used the selected-backend cleanup path") - return nil - }, - cleanupInactiveBackend: func(provider string) error { - if provider != string(filter.ProviderIptables) { - t.Fatalf("inactive cleanup provider = %q", provider) - } - return nil - }, - } - result, err := service.Reset(context.Background(), dto.FirewallRuleReset{Provider: filter.ProviderIptables}) - if err != nil { - t.Fatal(err) - } - if !result.Disabled { - t.Fatalf("inactive source cleanup result = %#v", result) - } -} - -func TestFirewallResetRestoresUFWDefaultsAndKeepsStoredRules(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - protectedRule := filter.FirewallRule{ - Scope: scope, NativeKind: filter.NativeKindUFWRule, Protocol: "tcp", - DestinationPort: "22", Action: filter.ActionAccept, - } - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - stored, err := model.FirewallRuleFromDomain(protectedRule) - if err != nil { - t.Fatal(err) - } - stored.UUID = "protected-reset-rule" - stored.Origin = constant.FirewallRuleOriginCreated - stored.Owner = constant.FirewallRuleSourceUser - if err := ruleRepo.Create(context.Background(), &stored); err != nil { - t.Fatal(err) - } - resetProvider := "" - service := &FirewallService{ - rules: ruleRepo, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil }, - resetBackend: func(provider string) error { - if provider != string(filter.ProviderUFW) { - t.Fatalf("reset unexpected provider %q", provider) - } - resetProvider = provider - return nil - }, - } - result, err := service.Reset(context.Background(), dto.FirewallRuleReset{Provider: filter.ProviderUFW}) - if err != nil { - t.Fatal(err) - } - if resetProvider != string(filter.ProviderUFW) || result.Removed != 1 || !result.Disabled { - t.Fatalf("unexpected reset result: provider=%q result=%#v", resetProvider, result) - } - remaining, err := ruleRepo.List(context.Background()) - if err != nil || len(remaining) != 1 || remaining[0].UUID != stored.UUID { - t.Fatalf("reset removed provider-neutral UFW policy: %#v err=%v", remaining, err) - } -} - -func TestFirewallRuleServiceCheckCreateInventoryWorkflow(t *testing.T) { - rule := executorTestAddressRule("172.16.10.111") - external := executorObservedRule(rule, "", 1) - adapter := newFakeFilterAdapter(t, rule.Scope, []filter.ObservedRule{external}) - adapter.snapshot.Notices = []filter.ScopeNotice{{Code: filter.ScopeNoticeManagedScopeInactive}} - db := newFirewallRuleTestDB(t) - service := &FirewallService{ - rules: repo.NewFirewallRuleRepo(db), - adapters: firewallRuleRuntimeRegistry{filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil)}, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - ctx := context.Background() - - before, err := service.Inventory(ctx, dto.FirewallRuleInventory{Scope: rule.Scope}) - if err != nil { - t.Fatalf("inventory before adoption: %v", err) - } - if len(before.Items) != 1 || before.Items[0].State != filter.InventoryStateExternal { - t.Fatalf("expected external inventory: %#v", before) - } - if len(before.Notices) != 1 || before.Notices[0].Code != filter.ScopeNoticeManagedScopeInactive { - t.Fatalf("scope notices were not returned: %#v", before.Notices) - } - adapter.snapshot.Notices = nil - check, err := service.checkRule(ctx, "", dto.FirewallRuleCheckItem{Rule: rule}) - if err != nil { - t.Fatalf("plan adoption: %v", err) - } - if check.Classification != filter.CheckClassificationExactExternal { - t.Fatalf("unexpected check result: %#v", check) - } - if check.CheckFlag == "" { - t.Fatal("check did not return a creation flag") - } - err = service.createFirewallRuleItem(ctx, dto.FirewallRuleCreateItem{ - CheckFlag: check.CheckFlag, Action: filter.CheckActionAdopt, - AdoptInstanceKey: check.Candidates[0].InstanceKey, Rule: check.RequestedRule, - }) - if err != nil { - t.Fatalf("commit adoption: %v", err) - } - after, err := service.Inventory(ctx, dto.FirewallRuleInventory{Scope: rule.Scope}) - if err != nil { - t.Fatalf("inventory after adoption: %v", err) - } - if len(after.Items) != 1 || after.Items[0].State != filter.InventoryStateAdopted || after.Items[0].Desired == nil { - t.Fatalf("adopted ownership missing from inventory: %#v", after) - } - deleted, err := service.Delete(ctx, dto.FirewallRuleDelete{UUIDs: []string{after.Items[0].Desired.UUID}}) - if err != nil || deleted.Failed > 0 { - t.Fatalf("delete adopted rule: result=%#v err=%v", deleted, err) - } - empty, err := service.Inventory(ctx, dto.FirewallRuleInventory{Scope: rule.Scope}) - if err != nil || len(empty.Items) != 0 { - t.Fatalf("deleted rule remained in inventory: inventory=%#v err=%v", empty, err) - } -} - -func TestFirewallRuleServiceCombinedUFWInventory(t *testing.T) { - ipv4Scope := filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - ipv6Scope := ipv4Scope - ipv6Scope.Family = filter.FamilyIPv6 - ipv4Rule := filter.FirewallRule{ - Scope: ipv4Scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: "8080", Action: filter.ActionAccept, - } - ipv6Rule := filter.FirewallRule{ - Scope: ipv6Scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "udp", SourceAddress: "2001:db8::/64", DestinationPort: "5353", Action: filter.ActionDrop, - } - adapter := newFakeFilterAdapter(t, ipv4Scope, []filter.ObservedRule{executorObservedRule(ipv4Rule, "", 1)}) - ipv6Snapshot, err := filter.NewSnapshot(ipv6Scope, []filter.ObservedRule{executorObservedRule(ipv6Rule, "", 2)}) - if err != nil { - t.Fatalf("create IPv6 snapshot: %v", err) - } - adapter.multiSnapshots = []filter.Snapshot{adapter.snapshot, ipv6Snapshot} - adapter.multiSnapshots[0].Notices = []filter.ScopeNotice{{Code: filter.ScopeNoticeManagedScopeInactive}} - adapter.multiSnapshots[1].Notices = []filter.ScopeNotice{{Code: filter.ScopeNoticeManagedScopeInactive}} - service := &FirewallService{ - rules: repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)), - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderUFW: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderUFW, nil }, - } - - result, err := service.Inventory(context.Background(), dto.FirewallRuleInventory{Scope: filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyInet, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - }}) - if err != nil { - t.Fatalf("load combined UFW inventory: %v", err) - } - if adapter.observeScopesCount != 1 || adapter.observeCount != 0 { - t.Fatalf("unexpected UFW observation counts: multi=%d single=%d", adapter.observeScopesCount, adapter.observeCount) - } - if len(result.Items) != 2 || result.Items[0].Rule.Scope.Family != filter.FamilyIPv4 || - result.Items[1].Rule.Scope.Family != filter.FamilyIPv6 { - t.Fatalf("unexpected combined UFW inventory: %#v", result.Items) - } - if len(result.Notices) != 1 || result.Notices[0].Code != filter.ScopeNoticeManagedScopeInactive { - t.Fatalf("duplicate or missing combined UFW notices: %#v", result.Notices) - } -} - -func TestFirewallRuleServiceCheckBlocksConfiguredProtectedPort(t *testing.T) { - rule := executorTestRule("8443") - rule.Action = filter.ActionDrop - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - service, _ := newTestFirewallExecutor(t, adapter) - service.selectedProvider = func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil } - service.protectedPorts = func() ([]firewall.PortWhitelist, error) { - return []firewall.PortWhitelist{{Family: "ipv4", Port: "8443", Protocol: "tcp"}}, nil - } - - result, err := service.checkRule(context.Background(), "203.0.113.9", dto.FirewallRuleCheckItem{Rule: rule}) - if err != nil { - t.Fatalf("check protected port: %v", err) - } - if result.Decision != filter.CheckDecisionBlocked || result.Reason != "current_management_connection" { - t.Fatalf("configured protected port was not blocked: %#v", result) - } -} - -func TestFirewallRuleServiceLoadsNativeDetailOnDemand(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, - Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput, - } - adapter := newFakeFilterAdapter(t, scope, nil) - adapter.nativeDetail = "ssh\n ports: 22/tcp\n protocols:\n source-ports:\n helpers:\n destination:" - service := &FirewallService{ - adapters: firewallRuleRuntimeRegistry{filter.ProviderFirewalld: newFirewallRuleRuntime(adapter, nil)}, - selectedProvider: func(context.Context) (filter.Provider, error) { - return filter.ProviderFirewalld, nil - }, - } - - info, err := service.LoadFirewallNativeDetail(context.Background(), dto.FirewallNativeDetail{ - Provider: filter.ProviderFirewalld, NativeKind: filter.NativeKindZoneService, - Name: "ssh", Permanent: true, - }) - if err != nil { - t.Fatalf("load service info: %v", err) - } - if info != adapter.nativeDetail || adapter.nativeDetailName != "ssh" || !adapter.nativeDetailPermanent { - t.Fatalf("service info request was not passed through: info=%q adapter=%#v", info, adapter) - } -} - -func TestFirewallRuleServiceCreateRequiresCheck(t *testing.T) { - rule := executorTestRule("8080") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - service, ruleRepo := newTestFirewallExecutor(t, adapter) - service.selectedProvider = func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil } - - result, err := service.Create(context.Background(), dto.FirewallRuleCreate{Items: []dto.FirewallRuleCreateItem{{ - Rule: rule, Action: filter.CheckActionCreate, - }}}) - if err != nil || result.Failed != 1 || len(result.Errors) != 1 || - result.Errors[0].Error != filter.ErrRuleCheckRequired.Error() { - t.Fatalf("expected check-required failure, result=%#v err=%v", result, err) - } - rules, _ := ruleRepo.List(context.Background()) - if adapter.applyCount != 0 || len(rules) != 0 { - t.Fatalf("unchecked rule was created: applyCount=%d rules=%#v", adapter.applyCount, rules) - } -} - -func TestFirewallRuleServiceHonorsExplicitUFWCreatePosition(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - existingRule := filter.FirewallRule{ - Scope: scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept, - } - existing := executorObservedRule(existingRule, "", 4) - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{existing}) - service, _ := newTestFirewallExecutor(t, adapter) - order := int64(2) - rule := filter.FirewallRule{ - Scope: scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: "8080", Action: filter.ActionAccept, OrderIndex: &order, - } - - if err := createExecutorRule(service, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, Action: filter.CheckActionCreate, SourceKind: constant.FirewallRuleSourceUser, - }); err != nil { - t.Fatalf("create UFW rule at explicit position: %v", err) - } - if adapter.lastChange.Append || adapter.lastChange.After == nil || adapter.lastChange.After.OrderIndex == nil || - *adapter.lastChange.After.OrderIndex != order { - t.Fatalf("UFW explicit create position was not preserved: %#v", adapter.lastChange) - } -} - -func TestFirewallRuleServiceDefaultsMissingUFWCreatePositionToAppend(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - existingRule := filter.FirewallRule{ - Scope: scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept, - } - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{executorObservedRule(existingRule, "", 4)}) - service, _ := newTestFirewallExecutor(t, adapter) - rule := filter.FirewallRule{ - Scope: scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: "8080", Action: filter.ActionAccept, - } - - if err := createExecutorRule(service, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, Action: filter.CheckActionCreate, SourceKind: constant.FirewallRuleSourceUser, - }); err != nil { - t.Fatalf("append UFW rule: %v", err) - } - if !adapter.lastChange.Append || adapter.lastChange.After == nil || adapter.lastChange.After.OrderIndex == nil || - *adapter.lastChange.After.OrderIndex != 5 { - t.Fatalf("missing UFW position did not default to append: %#v", adapter.lastChange) - } -} - -func TestFirewallRuleServiceDefaultsMissingUFWRequestScope(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - adapter := newFakeFilterAdapter(t, scope, nil) - service, _ := newTestFirewallExecutor(t, adapter) - service.selectedProvider = func(context.Context) (filter.Provider, error) { return filter.ProviderUFW, nil } - - checked, err := service.Check(context.Background(), "", dto.FirewallRuleCheck{Items: []dto.FirewallRuleCheckItem{{ - Rule: filter.FirewallRule{Protocol: "tcp", DestinationPort: "55101", Action: filter.ActionAccept}, - }}}) - if err != nil { - t.Fatalf("check UFW rule with omitted scope: %v", err) - } - if len(checked.Items) != 1 || checked.Items[0].RequestedRule.Scope.Normalize() != scope.Normalize() { - t.Fatalf("unexpected defaulted UFW scope: %#v", checked.Items) - } -} - -func TestValidateUFWPositionWithinFamilyBounds(t *testing.T) { - ipv4Scope := filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - ipv4Rule := filter.FirewallRule{ - Scope: ipv4Scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept, - } - ipv4Snapshot, err := filter.NewSnapshot(ipv4Scope, []filter.ObservedRule{ - executorObservedRule(ipv4Rule, "", 1), - executorObservedRule(ipv4Rule, "", 4), - }) - if err != nil { - t.Fatalf("create IPv4 snapshot: %v", err) - } - if err = validatePositionTarget(context.Background(), nil, ipv4Snapshot, ipv4Rule, 4); err != nil { - t.Fatalf("valid IPv4 position was rejected: %v", err) - } - if err = validatePositionTarget(context.Background(), nil, ipv4Snapshot, ipv4Rule, 5); !errors.Is(err, filter.ErrInvalidRule) { - t.Fatalf("IPv6 position was accepted for IPv4 rule: %v", err) - } - - ipv6Scope := ipv4Scope - ipv6Scope.Family = filter.FamilyIPv6 - ipv6Rule := ipv4Rule - ipv6Rule.Scope = ipv6Scope - ipv6Snapshot, err := filter.NewSnapshot(ipv6Scope, []filter.ObservedRule{ - executorObservedRule(ipv6Rule, "", 5), - executorObservedRule(ipv6Rule, "", 8), - }) - if err != nil { - t.Fatalf("create IPv6 snapshot: %v", err) - } - if err = validatePositionTarget(context.Background(), nil, ipv6Snapshot, ipv6Rule, 4); !errors.Is(err, filter.ErrInvalidRule) { - t.Fatalf("IPv4 position was accepted for IPv6 rule: %v", err) - } -} - -func TestUFWAppendPositionUsesFamilyBoundary(t *testing.T) { - ipv4Scope := filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - ipv4Rule := filter.FirewallRule{ - Scope: ipv4Scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept, - } - ipv4Snapshot, err := filter.NewSnapshot(ipv4Scope, []filter.ObservedRule{ - executorObservedRule(ipv4Rule, "", 1), - executorObservedRule(ipv4Rule, "", 4), - }) - if err != nil { - t.Fatalf("create IPv4 snapshot: %v", err) - } - if position, positionErr := ufwAppendPosition(context.Background(), nil, ipv4Snapshot, ipv4Rule); positionErr != nil || position != 5 { - t.Fatalf("unexpected IPv4 append position: position=%d err=%v", position, positionErr) - } - - ipv6Scope := ipv4Scope - ipv6Scope.Family = filter.FamilyIPv6 - ipv6Rule := ipv4Rule - ipv6Rule.Scope = ipv6Scope - ipv6Snapshot, err := filter.NewSnapshot(ipv6Scope, []filter.ObservedRule{ - executorObservedRule(ipv6Rule, "", 5), - executorObservedRule(ipv6Rule, "", 8), - }) - if err != nil { - t.Fatalf("create IPv6 snapshot: %v", err) - } - ipv4Adapter := newFakeFilterAdapter(t, ipv4Scope, ipv4Snapshot.Rules) - runtime := newFirewallRuleRuntime(ipv4Adapter, nil) - if position, positionErr := ufwAppendPosition(context.Background(), runtime, ipv6Snapshot, ipv6Rule); positionErr != nil || position != 9 { - t.Fatalf("unexpected IPv6 append position: position=%d err=%v", position, positionErr) - } -} - -func TestUFWGlobalEndPositionIncludesOtherFamily(t *testing.T) { - ipv4Scope := filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - ipv4Rule := filter.FirewallRule{ - Scope: ipv4Scope, NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: "4422,8088", Action: filter.ActionAccept, - } - ipv4Snapshot, err := filter.NewSnapshot(ipv4Scope, []filter.ObservedRule{ - executorObservedRule(ipv4Rule, "1panel-rule:managed", 4), - }) - if err != nil { - t.Fatalf("create IPv4 snapshot: %v", err) - } - - ipv6Scope := ipv4Scope - ipv6Scope.Family = filter.FamilyIPv6 - ipv6Rule := ipv4Rule - ipv6Rule.Scope = ipv6Scope - ipv6Snapshot, err := filter.NewSnapshot(ipv6Scope, []filter.ObservedRule{ - executorObservedRule(ipv6Rule, "", 5), - executorObservedRule(ipv6Rule, "", 8), - }) - if err != nil { - t.Fatalf("create IPv6 snapshot: %v", err) - } - runtime := newFirewallRuleRuntime(newFakeFilterAdapter(t, ipv6Scope, ipv6Snapshot.Rules), nil) - maxPosition, err := maxPositionForRule(context.Background(), runtime, ipv4Snapshot, ipv4Rule) - if err != nil { - t.Fatalf("load UFW global maximum position: %v", err) - } - if maxPosition != 8 { - t.Fatalf("IPv4 family boundary was mistaken for the global end: %d", maxPosition) - } - if int64(*ipv4Snapshot.Rules[0].Locator.Position) == maxPosition { - t.Fatal("last IPv4 rule was incorrectly classified as the global last UFW rule") - } -} - -func TestFirewallRuleServiceBatchCheckAndCreateSameScope(t *testing.T) { - rules := []filter.FirewallRule{executorTestRule("8080"), executorTestRule("8081")} - adapter := newFakeFilterAdapter(t, rules[0].Scope, nil) - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil)}, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - ctx := context.Background() - - checked, err := service.Check(ctx, "", dto.FirewallRuleCheck{Items: firewallRuleCheckItems(rules)}) - if err != nil || len(checked.Items) != len(rules) { - t.Fatalf("batch check: result=%#v err=%v", checked, err) - } - if adapter.observeCount != 1 { - t.Fatalf("same-scope batch check observed the chain %d times", adapter.observeCount) - } - request := dto.FirewallRuleCreate{Items: make([]dto.FirewallRuleCreateItem, 0, len(rules))} - for _, item := range checked.Items { - request.Items = append(request.Items, dto.FirewallRuleCreateItem{ - Rule: item.RequestedRule, CheckFlag: item.CheckFlag, Action: filter.CheckActionCreate, - SourceKind: constant.FirewallRuleSourceUser, - }) - } - created, err := service.Create(ctx, request) - if err != nil || created.Succeeded != 2 || created.Failed != 0 { - t.Fatalf("batch create: result=%#v err=%v", created, err) - } - stored, _ := ruleRepo.List(ctx) - if len(stored) != 2 || len(adapter.snapshot.Rules) != 2 || adapter.applyCount != 1 { - t.Fatalf("same-scope batch did not commit all rules: stored=%#v snapshot=%#v applies=%d", stored, adapter.snapshot, adapter.applyCount) - } - if adapter.observeCount != 3 { - t.Fatalf("same-scope batch create used repeated snapshots: observes=%d", adapter.observeCount) - } -} - -func TestFirewallRuleServiceBatchDeleteSameIptablesScope(t *testing.T) { - rules := []filter.FirewallRule{executorTestRule("8080"), executorTestRule("8081")} - adapter := newFakeFilterAdapter(t, rules[0].Scope, nil) - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil)}, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - ctx := context.Background() - checked, err := service.Check(ctx, "", dto.FirewallRuleCheck{Items: firewallRuleCheckItems(rules)}) - if err != nil { - t.Fatalf("batch check: %v", err) - } - createRequest := dto.FirewallRuleCreate{Items: make([]dto.FirewallRuleCreateItem, 0, len(rules))} - for _, item := range checked.Items { - createRequest.Items = append(createRequest.Items, dto.FirewallRuleCreateItem{ - Rule: item.RequestedRule, CheckFlag: item.CheckFlag, Action: filter.CheckActionCreate, - }) - } - created, err := service.Create(ctx, createRequest) - if err != nil || created.Succeeded != 2 { - t.Fatalf("batch create: result=%#v err=%v", created, err) - } - stored, err := ruleRepo.List(ctx) - if err != nil || len(stored) != 2 { - t.Fatalf("list created rules: %#v err=%v", stored, err) - } - deleted, err := service.Delete(ctx, dto.FirewallRuleDelete{UUIDs: []string{stored[0].UUID, stored[1].UUID}}) - if err != nil || deleted.Succeeded != 2 || deleted.Failed != 0 { - t.Fatalf("batch delete: result=%#v err=%v", deleted, err) - } - remaining, listErr := ruleRepo.List(ctx) - if listErr != nil || len(remaining) != 0 || len(adapter.snapshot.Rules) != 0 || adapter.applyCount != 2 { - t.Fatalf( - "same-scope delete did not use one backend apply: remaining=%#v snapshot=%#v applies=%d err=%v", - remaining, adapter.snapshot, adapter.applyCount, listErr, - ) - } -} - -func TestFirewallRuleServiceBatchCreateAndDeleteSameNftablesScope(t *testing.T) { - rules := []filter.FirewallRule{executorTestRule("8080"), executorTestRule("8081")} - for index := range rules { - rules[index].Scope.Provider = filter.ProviderNftables - } - adapter := newFakeFilterAdapter(t, rules[0].Scope, nil) - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil)}, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil }, - } - ctx := context.Background() - checked, err := service.Check(ctx, "", dto.FirewallRuleCheck{Items: firewallRuleCheckItems(rules)}) - if err != nil { - t.Fatalf("batch check: %v", err) - } - createRequest := dto.FirewallRuleCreate{Items: make([]dto.FirewallRuleCreateItem, 0, len(rules))} - for _, item := range checked.Items { - createRequest.Items = append(createRequest.Items, dto.FirewallRuleCreateItem{ - Rule: item.RequestedRule, CheckFlag: item.CheckFlag, Action: filter.CheckActionCreate, - }) - } - created, err := service.Create(ctx, createRequest) - if err != nil || created.Succeeded != 2 || adapter.applyCount != 1 { - t.Fatalf("nftables batch create: result=%#v applies=%d err=%v", created, adapter.applyCount, err) - } - stored, err := ruleRepo.List(ctx) - if err != nil || len(stored) != 2 { - t.Fatalf("list nftables rules: %#v err=%v", stored, err) - } - deleted, err := service.Delete(ctx, dto.FirewallRuleDelete{UUIDs: []string{stored[0].UUID, stored[1].UUID}}) - if err != nil || deleted.Succeeded != 2 || deleted.Failed != 0 || adapter.applyCount != 2 { - t.Fatalf("nftables batch delete: result=%#v applies=%d err=%v", deleted, adapter.applyCount, err) - } - remaining, err := ruleRepo.List(ctx) - if err != nil || len(remaining) != 0 || len(adapter.snapshot.Rules) != 0 { - t.Fatalf("nftables batch delete did not clear rules: stored=%#v snapshot=%#v err=%v", remaining, adapter.snapshot, err) - } -} - -func TestFirewallRuleServiceBatchDeleteRollsBackOnMetadataFailure(t *testing.T) { - rules := []filter.FirewallRule{executorTestRule("8080"), executorTestRule("8081")} - adapter := newFakeFilterAdapter(t, rules[0].Scope, nil) - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - failingRepo := &failingFirewallRuleRepo{IFirewallRuleRepo: ruleRepo, deleteErr: errors.New("metadata delete failed")} - service := &FirewallService{ - rules: failingRepo, - adapters: firewallRuleRuntimeRegistry{filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil)}, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - ctx := context.Background() - checked, err := service.Check(ctx, "", dto.FirewallRuleCheck{Items: firewallRuleCheckItems(rules)}) - if err != nil { - t.Fatalf("batch check: %v", err) - } - createRequest := dto.FirewallRuleCreate{Items: make([]dto.FirewallRuleCreateItem, 0, len(rules))} - for _, item := range checked.Items { - createRequest.Items = append(createRequest.Items, dto.FirewallRuleCreateItem{ - Rule: item.RequestedRule, CheckFlag: item.CheckFlag, Action: filter.CheckActionCreate, - }) - } - created, err := service.Create(ctx, createRequest) - if err != nil || created.Succeeded != 2 { - t.Fatalf("batch create: result=%#v err=%v", created, err) - } - stored, err := ruleRepo.List(ctx) - if err != nil || len(stored) != 2 { - t.Fatalf("list created rules: %#v err=%v", stored, err) - } - deleted, err := service.Delete(ctx, dto.FirewallRuleDelete{UUIDs: []string{stored[0].UUID, stored[1].UUID}}) - if err != nil || deleted.Succeeded != 0 || deleted.Failed != 2 { - t.Fatalf("batch delete failure: result=%#v err=%v", deleted, err) - } - remaining, listErr := ruleRepo.List(ctx) - if listErr != nil || len(remaining) != 2 || len(adapter.snapshot.Rules) != 2 || adapter.applyCount != 2 || adapter.rollbackCount != 1 { - t.Fatalf( - "failed delete batch was not rolled back: remaining=%#v snapshot=%#v applies=%d rollbacks=%d err=%v", - remaining, adapter.snapshot, adapter.applyCount, adapter.rollbackCount, listErr, - ) - } -} - -func TestFirewallRuleServiceBatchCreateRejectsDuplicateManagedRule(t *testing.T) { - rule := executorTestRule("8080") - trailingRule := executorTestRule("8081") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil)}, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - ctx := context.Background() - - checked, err := service.Check(ctx, "", dto.FirewallRuleCheck{ - Items: firewallRuleCheckItems([]filter.FirewallRule{rule, rule, trailingRule}), - }) - if err != nil || len(checked.Items) != 3 { - t.Fatalf("batch check duplicate rules: result=%#v err=%v", checked, err) - } - request := dto.FirewallRuleCreate{Items: make([]dto.FirewallRuleCreateItem, 0, 3)} - for _, item := range checked.Items { - request.Items = append(request.Items, dto.FirewallRuleCreateItem{ - Rule: item.RequestedRule, CheckFlag: item.CheckFlag, Action: filter.CheckActionCreate, - SourceKind: constant.FirewallRuleSourceUser, - }) - } - - created, err := service.Create(ctx, request) - if err != nil || created.Succeeded != 1 || created.Failed != 1 || created.Skipped != 1 { - t.Fatalf("batch duplicate create: result=%#v err=%v", created, err) - } - if len(created.Errors) != 2 || created.Errors[0].Index != 1 || created.Errors[0].Status != "failed" || - created.Errors[0].Error == "" || created.Errors[0].Rule.DestinationPort != "8080" || - created.Errors[1].Index != 2 || created.Errors[1].Status != "skipped" || created.Errors[1].Error != "" || - created.Errors[1].Rule.DestinationPort != "8081" { - t.Fatalf("batch failure details missing: %#v", created.Errors) - } - stored, listErr := ruleRepo.List(ctx) - if listErr != nil || len(stored) != 1 || len(adapter.snapshot.Rules) != 1 || adapter.applyCount != 1 { - t.Fatalf( - "duplicate managed rule was persisted: stored=%#v snapshot=%#v applies=%d err=%v", - stored, adapter.snapshot, adapter.applyCount, listErr, - ) - } -} - -func TestFirewallRuleServiceIptablesBatchDoesNotPersistRuntimeMetadata(t *testing.T) { - rules := []filter.FirewallRule{executorTestRule("8080"), executorTestRule("8081")} - adapter := newFakeFilterAdapter(t, rules[0].Scope, nil) - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - service := &FirewallService{ - rules: &failingFirewallRuleRepo{ - IFirewallRuleRepo: ruleRepo, - updateErr: errors.New("metadata write failed"), - }, - adapters: firewallRuleRuntimeRegistry{filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil)}, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - ctx := context.Background() - checked, err := service.Check(ctx, "", dto.FirewallRuleCheck{Items: firewallRuleCheckItems(rules)}) - if err != nil { - t.Fatalf("batch check: %v", err) - } - request := dto.FirewallRuleCreate{Items: make([]dto.FirewallRuleCreateItem, 0, len(rules))} - for _, item := range checked.Items { - request.Items = append(request.Items, dto.FirewallRuleCreateItem{ - Rule: item.RequestedRule, CheckFlag: item.CheckFlag, Action: filter.CheckActionCreate, - SourceKind: constant.FirewallRuleSourceUser, - }) - } - created, err := service.Create(ctx, request) - if err != nil || created.Succeeded != 2 || created.Failed != 0 || created.Skipped != 0 { - t.Fatalf("batch result: %#v err=%v", created, err) - } - stored, listErr := ruleRepo.List(ctx) - if listErr != nil || len(stored) != 2 || len(adapter.snapshot.Rules) != 2 || adapter.applyCount != 1 || adapter.rollbackCount != 0 { - t.Fatalf( - "batch persisted runtime metadata: stored=%#v snapshot=%#v applies=%d rollbacks=%d err=%v", - stored, adapter.snapshot, adapter.applyCount, adapter.rollbackCount, listErr, - ) - } -} - -func TestFirewallRuleServiceCreateRejectsChangedFirewallState(t *testing.T) { - rule := executorTestRule("8080") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - service, ruleRepo := newTestFirewallExecutor(t, adapter) - service.selectedProvider = func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil } - ctx := context.Background() - - check, err := service.checkRule(ctx, "", dto.FirewallRuleCheckItem{Rule: rule}) - if err != nil { - t.Fatalf("check rule: %v", err) - } - other := executorTestRule("9090") - adapter.snapshot, err = filter.NewSnapshot(rule.Scope, []filter.ObservedRule{executorObservedRule(other, "", 1)}) - if err != nil { - t.Fatalf("change firewall snapshot: %v", err) - } - err = service.createFirewallRuleItem(ctx, dto.FirewallRuleCreateItem{ - Rule: check.RequestedRule, CheckFlag: check.CheckFlag, Action: filter.CheckActionCreate, - }) - if !errors.Is(err, filter.ErrRuleCheckRequired) { - t.Fatalf("expected changed snapshot to require another check, got %v", err) - } - rules, _ := ruleRepo.List(ctx) - if adapter.applyCount != 0 || len(rules) != 0 { - t.Fatalf("stale checked rule was created: applyCount=%d rules=%#v", adapter.applyCount, rules) - } -} - -func TestFirewallRuleServiceCreateRejectsRuleChangedAfterCheck(t *testing.T) { - rule := executorTestRule("8080") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - service, ruleRepo := newTestFirewallExecutor(t, adapter) - service.selectedProvider = func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil } - ctx := context.Background() - - check, err := service.checkRule(ctx, "", dto.FirewallRuleCheckItem{Rule: rule}) - if err != nil { - t.Fatalf("check rule: %v", err) - } - changed := check.RequestedRule - changed.DestinationPort = "8081" - err = service.createFirewallRuleItem(ctx, dto.FirewallRuleCreateItem{ - Rule: changed, CheckFlag: check.CheckFlag, Action: filter.CheckActionCreate, - }) - if !errors.Is(err, filter.ErrRuleCheckRequired) { - t.Fatalf("expected changed rule to require another check, got %v", err) - } - rules, _ := ruleRepo.List(ctx) - if adapter.applyCount != 0 || len(rules) != 0 { - t.Fatalf("changed unchecked rule was created: applyCount=%d rules=%#v", adapter.applyCount, rules) - } -} - -func TestFirewallRuleServiceCreateRejectsChangedManagedState(t *testing.T) { - rule := executorTestRule("8080") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - service, ruleRepo := newTestFirewallExecutor(t, adapter) - service.selectedProvider = func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil } - ctx := context.Background() - - check, err := service.checkRule(ctx, "", dto.FirewallRuleCheckItem{Rule: rule}) - if err != nil { - t.Fatalf("check rule: %v", err) - } - other := executorTestRule("9090") - record, err := firewallRuleModelForCreate(other, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}, constant.FirewallRuleOriginCreated) - if err != nil { - t.Fatalf("build managed record: %v", err) - } - if err := ruleRepo.Create(ctx, &record); err != nil { - t.Fatalf("change managed state: %v", err) - } - err = service.createFirewallRuleItem(ctx, dto.FirewallRuleCreateItem{ - Rule: check.RequestedRule, CheckFlag: check.CheckFlag, Action: filter.CheckActionCreate, - }) - if !errors.Is(err, filter.ErrRuleCheckRequired) { - t.Fatalf("expected changed managed state to require another check, got %v", err) - } - if adapter.applyCount != 0 { - t.Fatalf("rule was applied with a stale managed-state check: %d", adapter.applyCount) - } -} - -func TestFirewallRuleServiceRejectsUnavailableProductionAdapter(t *testing.T) { - db := newFirewallRuleTestDB(t) - service := &FirewallService{ - rules: repo.NewFirewallRuleRepo(db), adapters: firewallRuleRuntimeRegistry{}, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - _, err := service.checkRule(context.Background(), "", dto.FirewallRuleCheckItem{Rule: executorTestRule("80")}) - if !errors.Is(err, filter.ErrAdapterUnavailable) { - t.Fatalf("expected unavailable adapter error, got %v", err) - } -} - -func TestPrepareFirewallRuleUsesProviderRepresentation(t *testing.T) { - rule, err := filter.NormalizeRule(filter.FirewallRule{ - Scope: filter.Scope{Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, Zone: "public", Direction: filter.DirectionInput}, - Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, - }) - if err != nil { - t.Fatalf("normalize request: %v", err) - } - prepared, err := newFirewallRuleRuntime(filterfirewalld.NewAdapterWithReader(nil), nil).Prepare(rule) - if err != nil { - t.Fatalf("prepare firewalld request: %v", err) - } - if prepared.NativeKind != filter.NativeKindZonePort { - t.Fatalf("service planned against the wrong native identity: %#v", prepared) - } -} - -func TestNormalizeFirewallRuleScopeDefaultsIptablesChains(t *testing.T) { - input := filter.Scope{Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, Table: "filter", Direction: filter.DirectionInput}.Normalize() - if input.Chain != filter.IptablesInputChain { - t.Fatalf("unexpected input chain: %#v", input) - } -} - -func TestNormalizeSystemPortsKeepsFamilyAndRange(t *testing.T) { - ports, err := normalizeSystemPorts([]dto.FirewallSystemPort{ - {Family: "ipv4", Protocol: "tcp", Port: "80"}, - {Family: "ipv6", Protocol: "udp", Port: "8000:8100"}, - }) - if err != nil { - t.Fatalf("normalize system ports: %v", err) - } - if _, ok := ports["ipv4/tcp/80"]; !ok { - t.Fatalf("missing IPv4 port: %#v", ports) - } - if port, ok := ports["ipv6/udp/8000-8100"]; !ok || port.Family != "ipv6" || port.Port != "8000-8100" { - t.Fatalf("missing normalized IPv6 range: %#v", ports) - } -} - -func TestSystemPortRuleUsesRequestedFamily(t *testing.T) { - port := dto.FirewallSystemPort{Family: "ipv6", Protocol: "tcp", Port: "443"} - for _, provider := range []filter.Provider{ - filter.ProviderIptables, filter.ProviderNftables, filter.ProviderFirewalld, filter.ProviderUFW, - } { - rule := systemPortRule(provider, port) - if rule.Scope.Family != filter.FamilyIPv6 { - t.Fatalf("provider %s used family %s", provider, rule.Scope.Family) - } - } -} - -func TestSyncSystemPortsCreatesTracksAndDeletesAcceptedPort(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, Table: "filter", - Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - adapter := newFakeFilterAdapter(t, scope, nil) - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - engine := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - port := dto.FirewallSystemPort{Port: "8443", Protocol: "TCP"} - - if err := engine.SyncSystemPorts(context.Background(), nil, []dto.FirewallSystemPort{port}); err != nil { - t.Fatalf("create accepted port: %v", err) - } - stored, err := ruleRepo.List(context.Background(), repo.WithFirewallRuleSource( - constant.FirewallRuleSourceSecurity, constant.FirewallSystemAcceptedPortSourcePrefix+"tcp/8443", - )) - if err != nil || len(stored) != 1 { - t.Fatalf("accepted port ownership was not persisted: rules=%#v err=%v", stored, err) - } - if len(adapter.snapshot.Rules) != 1 { - t.Fatalf("accepted port was not applied: %#v", adapter.snapshot.Rules) - } - - if err := engine.SyncSystemPorts(context.Background(), []dto.FirewallSystemPort{port}, nil); err != nil { - t.Fatalf("delete accepted port: %v", err) - } - if len(adapter.snapshot.Rules) != 0 { - t.Fatalf("accepted port remained after deletion: %#v", adapter.snapshot.Rules) - } - present, err := ruleRepo.List(context.Background(), - repo.WithFirewallRuleSource(constant.FirewallRuleSourceSecurity, constant.FirewallSystemAcceptedPortSourcePrefix+"tcp/8443"), - ) - if err != nil || len(present) != 0 { - t.Fatalf("accepted port ownership remained present: rules=%#v err=%v", present, err) - } -} - -func TestSyncSystemPortsBatchesNativeAcceptedPortsByScope(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, Table: "filter", - Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - adapter := newFakeFilterAdapter(t, scope, nil) - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - engine := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - ports := []dto.FirewallSystemPort{ - {Family: "ipv4", Port: "8080", Protocol: "tcp"}, - {Family: "ipv4", Port: "8443", Protocol: "tcp"}, - {Family: "ipv4", Port: "5353", Protocol: "udp"}, - } - - if err := engine.SyncSystemPorts(context.Background(), nil, ports); err != nil { - t.Fatalf("batch create accepted ports: %v", err) - } - if adapter.applyCount != 1 || len(adapter.snapshot.Rules) != len(ports) { - t.Fatalf("accepted ports were not created in one native apply: applies=%d rules=%d", adapter.applyCount, len(adapter.snapshot.Rules)) - } - if err := engine.SyncSystemPorts(context.Background(), ports, nil); err != nil { - t.Fatalf("batch delete accepted ports: %v", err) - } - if adapter.applyCount != 2 || len(adapter.snapshot.Rules) != 0 { - t.Fatalf("accepted ports were not deleted in one native apply: applies=%d rules=%d", adapter.applyCount, len(adapter.snapshot.Rules)) - } -} - -func TestSyncSystemPortsCreatesAcceptedPortDespitePartialOppositeOverlap(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, - Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput, - } - existing := filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderFirewalld, Family: filter.FamilyIPv4, - Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput, - }, - NativeKind: filter.NativeKindRichRule, Protocol: "all", - SourceAddress: "1.1.1.1", Action: filter.ActionDrop, - } - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{executorObservedRule(existing, "", 1)}) - engine := &FirewallService{ - rules: repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)), - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderFirewalld: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderFirewalld, nil }, - } - - err := engine.SyncSystemPorts(context.Background(), nil, []dto.FirewallSystemPort{{Port: "443", Protocol: "tcp"}}) - if err != nil { - t.Fatalf("create partially overlapping accepted port: %v", err) - } - if adapter.applyCount != 1 || len(adapter.snapshot.Rules) != 2 { - t.Fatalf("accepted port was not created: applyCount=%d rules=%#v", adapter.applyCount, adapter.snapshot.Rules) - } -} - -func TestSyncSystemPortsRejectsFullyCoveredOppositeRule(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, - Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput, - } - existing := filter.FirewallRule{ - Scope: scope, NativeKind: filter.NativeKindRichRule, Protocol: "tcp", - DestinationPort: "443", Action: filter.ActionDrop, - } - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{executorObservedRule(existing, "", 1)}) - engine := &FirewallService{ - rules: repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)), - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderFirewalld: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderFirewalld, nil }, - } - - err := engine.SyncSystemPorts(context.Background(), nil, []dto.FirewallSystemPort{{Port: "443", Protocol: "tcp"}}) - if err == nil || !strings.Contains(err.Error(), "overlapping_rule_with_different_action") { - t.Fatalf("fully covered opposite rule returned %v", err) - } - if adapter.applyCount != 0 { - t.Fatalf("fully conflicting accepted port was applied %d times", adapter.applyCount) - } -} - -func TestSyncSystemPortsDoesNotTakeOverExistingManagedRule(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, Table: "filter", - Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - port := dto.FirewallSystemPort{Port: "443", Protocol: "tcp"} - rule := systemPortRule(filter.ProviderIptables, port) - observed := executorObservedRule(rule, "1panel-rule:user-rule", 1) - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{observed}) - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - userRecord, err := firewallRuleModelForCreate(rule, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}, constant.FirewallRuleOriginCreated) - if err != nil { - t.Fatal(err) - } - userRecord.UUID = "user-rule" - if err := ruleRepo.Create(context.Background(), &userRecord); err != nil { - t.Fatal(err) - } - engine := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - - if err := engine.SyncSystemPorts(context.Background(), nil, []dto.FirewallSystemPort{port}); err != nil { - t.Fatalf("reuse existing managed rule: %v", err) - } - systemOwned, err := ruleRepo.List(context.Background(), repo.WithFirewallRuleSource( - constant.FirewallRuleSourceSecurity, constant.FirewallSystemAcceptedPortSourcePrefix+"tcp/443", - )) - if err != nil || len(systemOwned) != 0 || adapter.applyCount != 0 { - t.Fatalf("existing user rule was taken over: rules=%#v applyCount=%d err=%v", systemOwned, adapter.applyCount, err) - } -} - -func TestSyncSystemPortsAdoptsAndDeletesLegacyAcceptedPort(t *testing.T) { - scope := filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, Table: "filter", - Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - port := dto.FirewallSystemPort{Port: "8080", Protocol: "tcp"} - external := executorObservedRule(systemPortRule(filter.ProviderIptables, port), "", 1) - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{external}) - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - engine := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderIptables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - - if err := engine.SyncSystemPorts(context.Background(), []dto.FirewallSystemPort{port}, nil); err != nil { - t.Fatalf("remove legacy accepted port: %v", err) - } - if len(adapter.snapshot.Rules) != 0 || adapter.applyCount != 2 { - t.Fatalf("legacy rule did not use adopt/delete workflow: snapshot=%#v applyCount=%d", adapter.snapshot, adapter.applyCount) - } -} -func TestFirewallExecutorCreatesAndVerifiesRule(t *testing.T) { - rule := executorTestRule("8080") - position := int64(1) - rule.OrderIndex = &position - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - request := dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - } - - if err := createExecutorRule(executor, adapter, request); err != nil { - t.Fatalf("commit create: %v", err) - } - if adapter.applyCount != 1 { - t.Fatalf("unexpected apply count: %d", adapter.applyCount) - } - rules, _ := ruleRepo.List(context.Background()) - if len(rules) != 1 || rules[0].DestinationPort != "8080" { - t.Fatalf("rule was not verified and bound: %#v", rules) - } -} - -func TestFirewallExecutorAdoptsWithoutAddingEquivalentRule(t *testing.T) { - rule := executorTestAddressRule("172.16.10.111") - external := executorObservedRule(rule, "", 1) - adapter := newFakeFilterAdapter(t, rule.Scope, []filter.ObservedRule{external}) - instanceKey, _ := filter.InstanceKey(external) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - AdoptInstanceKey: instanceKey, Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }) - if err != nil { - t.Fatalf("commit adoption: %v", err) - } - if len(adapter.snapshot.Rules) != 1 || adapter.snapshot.Rules[0].Marker == "" { - t.Fatalf("adoption changed rule count or missed marker: snapshot=%#v", adapter.snapshot) - } - rules, _ := ruleRepo.List(context.Background()) - if len(rules) != 1 || rules[0].Origin != constant.FirewallRuleOriginAdopted { - t.Fatalf("adopted ownership was not persisted: %#v", rules) - } -} - -func TestFirewallExecutorCleansUpVerificationFailure(t *testing.T) { - rule := executorTestRule("9090") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - adapter.verifyMatched = false - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - - err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }) - if !errors.Is(err, filter.ErrVerificationFailed) { - t.Fatalf("expected verification failure, got err=%v", err) - } - rules, _ := ruleRepo.List(context.Background()) - if len(rules) != 0 { - t.Fatalf("failed rule metadata was not cleaned up: %#v", rules) - } - if len(adapter.snapshot.Rules) != 0 || adapter.rollbackCount != 1 { - t.Fatalf("failed runtime rule was not rolled back: snapshot=%#v rollbacks=%d", adapter.snapshot, adapter.rollbackCount) - } -} - -func TestFirewallExecutorCreateDoesNotPersistRuntimeMetadata(t *testing.T) { - rule := executorTestRule("9091") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - executor.rules = &failingFirewallRuleRepo{ - IFirewallRuleRepo: ruleRepo, - updateErr: errors.New("commit failed"), - } - - err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }) - if err != nil { - t.Fatalf("runtime metadata update affected create: %v", err) - } - stored, _ := ruleRepo.List(context.Background()) - if len(stored) != 1 || len(adapter.snapshot.Rules) != 1 || adapter.rollbackCount != 0 { - t.Fatalf("create persisted runtime metadata: stored=%#v snapshot=%#v rollbacks=%d", stored, adapter.snapshot, adapter.rollbackCount) - } -} - -func TestFirewallExecutorRollsBackDeleteWhenPersistenceDeleteFails(t *testing.T) { - rule := executorTestRule("9092") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - if err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }); err != nil { - t.Fatalf("create managed rule: %v", err) - } - stored, _ := ruleRepo.List(context.Background()) - executor.rules = &failingFirewallRuleRepo{ - IFirewallRuleRepo: ruleRepo, - deleteErr: errors.New("delete commit failed"), - } - - err := executor.deleteRule(context.Background(), stored[0].UUID) - if err == nil || !strings.Contains(err.Error(), "delete commit failed") { - t.Fatalf("expected persistence delete failure, got %v", err) - } - remaining, _ := ruleRepo.List(context.Background()) - if len(remaining) != 1 || len(adapter.snapshot.Rules) != 1 || adapter.rollbackCount != 1 { - t.Fatalf("failed delete was not restored: stored=%#v snapshot=%#v rollbacks=%d", remaining, adapter.snapshot, adapter.rollbackCount) - } -} - -func TestFirewallExecutorRollsBackUpdateWhenPersistenceUpdateFails(t *testing.T) { - rule := executorTestRule("9093") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - if err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }); err != nil { - t.Fatalf("create managed rule: %v", err) - } - stored, _ := ruleRepo.List(context.Background()) - executor.rules = &failingFirewallRuleRepo{ - IFirewallRuleRepo: ruleRepo, - updateErr: errors.New("update commit failed"), - } - updated := rule - updated.DestinationPort = "9443" - - err := executor.updateRule(context.Background(), "", stored[0].UUID, updated) - if err == nil || !strings.Contains(err.Error(), "update commit failed") { - t.Fatalf("expected persistence update failure, got %v", err) - } - remaining, _ := ruleRepo.List(context.Background()) - if len(remaining) != 1 || remaining[0].DestinationPort != "9093" || len(adapter.snapshot.Rules) != 1 || - adapter.snapshot.Rules[0].Rule.DestinationPort != "9093" || adapter.rollbackCount != 1 { - t.Fatalf("failed update was not restored: stored=%#v snapshot=%#v rollbacks=%d", remaining, adapter.snapshot, adapter.rollbackCount) - } -} - -func TestFirewallExecutorDeletesOnlyVerifiedManagedRule(t *testing.T) { - rule := executorTestRule("8443") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }) - if err != nil { - t.Fatalf("create managed rule: %v", err) - } - stored, err := ruleRepo.List(context.Background()) - if err != nil || len(stored) != 1 { - t.Fatalf("load managed rule: rules=%#v err=%v", stored, err) - } - if err := executor.deleteRule(context.Background(), stored[0].UUID); err != nil { - t.Fatalf("delete managed rule: %v", err) - } - if len(adapter.snapshot.Rules) != 0 || adapter.applyCount != 2 { - t.Fatalf("unexpected delete result: snapshot=%#v applies=%d", adapter.snapshot, adapter.applyCount) - } - remaining, err := ruleRepo.List(context.Background()) - if err != nil || len(remaining) != 0 { - t.Fatalf("deleted rule was not archived: rules=%#v err=%v", remaining, err) - } -} - -func TestFirewallExecutorRefusesDeleteAfterManagedRuleDrifts(t *testing.T) { - rule := executorTestRule("9443") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }) - if err != nil { - t.Fatalf("create managed rule: %v", err) - } - stored, _ := ruleRepo.List(context.Background()) - adapter.snapshot.Rules[0].Rule.DestinationPort = "9444" - adapter.snapshot, _ = filter.NewSnapshot(adapter.snapshot.Scope, adapter.snapshot.Rules) - err = executor.deleteRule(context.Background(), stored[0].UUID) - if !errors.Is(err, filter.ErrRuleStale) { - t.Fatalf("expected drifted delete rejection, got %v", err) - } - if adapter.applyCount != 1 { - t.Fatalf("drifted rule was deleted: applyCount=%d", adapter.applyCount) - } -} - -func TestFirewallExecutorUpdatesManagedRule(t *testing.T) { - rule := executorTestRule("8080") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }) - if err != nil { - t.Fatalf("create managed rule: %v", err) - } - stored, _ := ruleRepo.List(context.Background()) - updated := rule - updated.DestinationPort = "8443" - updated.Description = "updated" - if err := executor.updateRule(context.Background(), "", stored[0].UUID, updated); err != nil { - t.Fatalf("update managed rule: %v", err) - } - if adapter.applyCount != 2 { - t.Fatalf("unexpected update apply count: %d", adapter.applyCount) - } - after, _ := ruleRepo.GetByUUID(context.Background(), stored[0].UUID) - desired, err := desiredFirewallRuleFromModel(after) - if err != nil { - t.Fatalf("decode updated rule: %v", err) - } - if desired.Rule.DestinationPort != "8443" || desired.Rule.Description != "updated" { - t.Fatalf("updated semantics were not persisted: %#v", after) - } -} - -func TestFirewallExecutorUpdatesUFWDescriptionWithoutNativeMutation(t *testing.T) { - rule := executorTestRule("8080") - rule.Scope = filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - rule.NativeKind = filter.NativeKindUFWRule - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - if err := createExecutorRule(executor, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }); err != nil { - t.Fatalf("create managed UFW rule: %v", err) - } - stored, _ := ruleRepo.List(context.Background()) - updated := rule - updated.Description = "updated description" - if err := executor.updateRule(context.Background(), "", stored[0].UUID, updated); err != nil { - t.Fatalf("update managed UFW description: %v", err) - } - if adapter.applyCount != 1 { - t.Fatalf("description-only update changed UFW: applyCount=%d", adapter.applyCount) - } - after, err := ruleRepo.GetByUUID(context.Background(), stored[0].UUID) - if err != nil { - t.Fatalf("load updated UFW rule: %v", err) - } - if after.Description != "updated description" || after.DestinationPort != "8080" { - t.Fatalf("unexpected persisted UFW rule: %#v", after) - } -} - -func TestIsUFWMetadataOnlyUpdate(t *testing.T) { - position := 3 - order := int64(position) - before := executorTestRule("8080") - before.Scope = filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, - Chain: filter.UFWInputChain, Direction: filter.DirectionInput, - } - before.NativeKind = filter.NativeKindUFWRule - before.OrderIndex = &order - after := before - after.Description = "new description" - - metadataOnly, err := isUFWMetadataOnlyUpdate(before, after, filter.Locator{Position: &position}) - if err != nil || !metadataOnly { - t.Fatalf("description-only UFW update = %v, err=%v", metadataOnly, err) - } - - tests := []struct { - name string - mutate func(*filter.FirewallRule, *filter.Locator) - wantErr bool - }{ - {name: "rule changed", mutate: func(rule *filter.FirewallRule, _ *filter.Locator) { rule.DestinationPort = "8081" }}, - {name: "order changed", mutate: func(rule *filter.FirewallRule, _ *filter.Locator) { changed := int64(4); rule.OrderIndex = &changed }}, - {name: "missing locator", mutate: func(_ *filter.FirewallRule, locator *filter.Locator) { locator.Position = nil }}, - {name: "invalid rule", mutate: func(rule *filter.FirewallRule, _ *filter.Locator) { rule.DestinationPort = "invalid" }, wantErr: true}, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - candidate := after - locator := filter.Locator{Position: &position} - test.mutate(&candidate, &locator) - got, err := isUFWMetadataOnlyUpdate(before, candidate, locator) - if test.wantErr { - if err == nil { - t.Fatal("invalid rule did not return an error") - } - return - } - if err != nil || got { - t.Fatalf("metadata-only update = %v, err=%v", got, err) - } - }) - } -} - -func TestFirewallRuntimeRejectsUnavailableManagedScopeForMutation(t *testing.T) { - rule := executorTestRule("8080") - for _, notice := range []filter.ScopeNoticeCode{ - filter.ScopeNoticeManagedScopeInactive, - filter.ScopeNoticeManagedScopeMissing, - } { - t.Run(string(notice), func(t *testing.T) { - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - adapter.snapshot.Notices = []filter.ScopeNotice{{Code: notice}} - runtime := newFirewallRuleRuntime(adapter, nil) - if _, err := runtime.ObserveMutation(context.Background(), rule.Scope); !errors.Is(err, filter.ErrProviderUnavailable) { - t.Fatalf("mutation observe error = %v", err) - } - if _, err := runtime.Observe(context.Background(), rule.Scope); err != nil { - t.Fatalf("read-only observe rejected unavailable scope: %v", err) - } - }) - } -} - -func TestFilterChainOperationValidationAllowsOnlyUnifiedOperations(t *testing.T) { - validate := validator.New() - for _, operation := range []string{"init-base", "bind-base", "unbind-base"} { - request := dto.FilterChainOperation{Name: constant.FirewallBasicChain, Operate: operation} - if err := validate.Struct(request); err != nil { - t.Fatalf("operation %q rejected by API contract: %v", operation, err) - } - } - for _, operation := range []string{"", "init-ipv6-base", "init-forward", "repair-anything"} { - request := dto.FilterChainOperation{Name: constant.FirewallBasicChain, Operate: operation} - if err := validate.Struct(request); err == nil { - t.Fatalf("operation %q accepted by API contract", operation) - } - } -} - -func TestLoadSystemFirewallFamilyInfoExcludesServiceBackends(t *testing.T) { - for _, provider := range []string{constant.FirewallProviderFirewalld, constant.FirewallProviderUFW, "unsupported"} { - status := loadSystemFirewallFamilyInfo(provider, constant.FirewallFamilyIPv6) - if status.Available || status.Initialized || status.Bound { - t.Fatalf("%s IPv6 status = %#v, want unavailable", provider, status) - } - } -} - -func TestSupportsManagedFilterChains(t *testing.T) { - for _, provider := range []string{constant.FirewallProviderIptables, constant.FirewallProviderNftables} { - if !supportsManagedFilterChains(provider) { - t.Fatalf("%s should support managed filter chains", provider) - } - } - for _, provider := range []string{constant.FirewallProviderFirewalld, constant.FirewallProviderUFW} { - if supportsManagedFilterChains(provider) { - t.Fatalf("%s should not support managed filter chains", provider) - } - } -} - -func TestFirewallRuleServiceChecksManagedUpdateWithoutApplying(t *testing.T) { - rule := executorTestRule("8080") - adapter := newFakeFilterAdapter(t, rule.Scope, nil) - service, ruleRepo := newTestFirewallExecutor(t, adapter) - if err := createExecutorRule(service, adapter, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceUser, - }); err != nil { - t.Fatalf("create managed rule: %v", err) - } - stored, err := ruleRepo.List(context.Background()) - if err != nil || len(stored) != 1 { - t.Fatalf("load managed rule: rules=%#v err=%v", stored, err) - } - updated := rule - updated.DestinationPort = "8443" - result, err := service.checkRule(context.Background(), "", dto.FirewallRuleCheckItem{ - UUID: stored[0].UUID, - Rule: updated, - }) - if err != nil { - t.Fatalf("check managed update: %v", err) - } - if result.Decision != filter.CheckDecisionReady || result.Reason != "update_ready" || result.RequestedRule.DestinationPort != "8443" { - t.Fatalf("unexpected managed update check: %#v", result) - } - if adapter.applyCount != 1 { - t.Fatalf("update check changed the firewall: applyCount=%d", adapter.applyCount) - } -} - -func TestFirewallExecutorUpdatesManagedRuleAndPositionTogether(t *testing.T) { - first := executorTestRule("8080") - first.UUID = "first" - second := executorTestRule("8081") - second.UUID = "second" - adapter := newFakeFilterAdapter(t, first.Scope, []filter.ObservedRule{ - executorObservedRule(first, "1panel-rule:first", 1), - executorObservedRule(second, "1panel-rule:second", 2), - }) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - for _, rule := range []filter.FirewallRule{first, second} { - record, err := firewallRuleModelForCreate(rule, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}, constant.FirewallRuleOriginCreated) - if err != nil { - t.Fatalf("build managed rule: %v", err) - } - record.UUID = rule.UUID - if err := ruleRepo.Create(context.Background(), &record); err != nil { - t.Fatalf("create managed record: %v", err) - } - } - target := int64(2) - updated := first - updated.DestinationPort = "8443" - updated.OrderIndex = &target - if err := executor.updateRule(context.Background(), "", first.UUID, updated); err != nil { - t.Fatalf("positioned update: %v", err) - } - if adapter.snapshot.Rules[1].Marker != "1panel-rule:first" || adapter.snapshot.Rules[1].Rule.DestinationPort != "8443" { - t.Fatalf("rule content and position were not updated together: %#v", adapter.snapshot.Rules) - } - stored, err := ruleRepo.GetByUUID(context.Background(), first.UUID) - if err != nil || stored.DestinationPort != "8443" { - t.Fatalf("request-only position was persisted: stored=%#v err=%v", stored, err) - } -} - -func TestFirewallExecutorReordersManagedRule(t *testing.T) { - scope := executorTestRule("8080").Scope - first := executorTestRule("8080") - first.UUID = "first" - second := executorTestRule("8081") - second.UUID = "second" - observed := []filter.ObservedRule{ - executorObservedRule(first, "1panel-rule:first", 1), - executorObservedRule(second, "1panel-rule:second", 2), - } - adapter := newFakeFilterAdapter(t, scope, observed) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - for _, rule := range []filter.FirewallRule{first, second} { - record, err := firewallRuleModelForCreate(rule, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}, constant.FirewallRuleOriginCreated) - if err != nil { - t.Fatalf("build managed rule: %v", err) - } - record.UUID = rule.UUID - if err := ruleRepo.Create(context.Background(), &record); err != nil { - t.Fatalf("create managed record: %v", err) - } - } - target := int64(2) - err := executor.reorderRule(context.Background(), "", "first", &target, nil) - if err != nil || adapter.applyCount != 1 || adapter.snapshot.Rules[1].Marker != "1panel-rule:first" { - t.Fatalf("managed reorder failed: applyCount=%d snapshot=%#v err=%v", adapter.applyCount, adapter.snapshot.Rules, err) - } - firstRecord, err := ruleRepo.GetByUUID(context.Background(), "first") - if err != nil || firstRecord.Sequence == nil || *firstRecord.Sequence != 2*model.FirewallRuleSequenceStep { - t.Fatalf("reordered rule sequence was not persisted: record=%#v err=%v", firstRecord, err) - } - secondRecord, err := ruleRepo.GetByUUID(context.Background(), "second") - if err != nil || secondRecord.Sequence == nil || *secondRecord.Sequence != model.FirewallRuleSequenceStep { - t.Fatalf("neighbor sequence was not rebalanced: record=%#v err=%v", secondRecord, err) - } -} - -func TestFirewallExecutorReorderUsesSparseSequenceMidpoint(t *testing.T) { - first := executorTestRule("8080") - first.UUID = "first" - second := executorTestRule("8081") - second.UUID = "second" - third := executorTestRule("8082") - third.UUID = "third" - adapter := newFakeFilterAdapter(t, first.Scope, []filter.ObservedRule{ - executorObservedRule(first, "1panel-rule:first", 1), - executorObservedRule(second, "1panel-rule:second", 2), - executorObservedRule(third, "1panel-rule:third", 3), - }) - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - for index, rule := range []filter.FirewallRule{first, second, third} { - record, err := firewallRuleModelForCreate(rule, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}, constant.FirewallRuleOriginCreated) - if err != nil { - t.Fatalf("build managed rule: %v", err) - } - record.UUID = rule.UUID - sequence := int64(index+1) * model.FirewallRuleSequenceStep - record.Sequence = &sequence - if err := ruleRepo.Create(context.Background(), &record); err != nil { - t.Fatalf("create managed record: %v", err) - } - } - - target := int64(2) - if err := executor.reorderRule(context.Background(), "", third.UUID, &target, nil); err != nil { - t.Fatalf("reorder managed rule: %v", err) - } - want := map[string]int64{ - first.UUID: model.FirewallRuleSequenceStep, - third.UUID: model.FirewallRuleSequenceStep + model.FirewallRuleSequenceStep/2, - second.UUID: 2 * model.FirewallRuleSequenceStep, - } - for uuid, wantSequence := range want { - record, err := ruleRepo.GetByUUID(context.Background(), uuid) - if err != nil || record.Sequence == nil || *record.Sequence != wantSequence { - t.Fatalf("sequence for %s = %#v, want %d (err=%v)", uuid, record.Sequence, wantSequence, err) - } - } -} - -func TestFirewallExecutorUsesExplicitPositionDuringUpdate(t *testing.T) { - first := executorTestRule("8080") - first.UUID = "first" - second := executorTestRule("8081") - second.UUID = "second" - adapter := newFakeFilterAdapter(t, first.Scope, []filter.ObservedRule{ - executorObservedRule(first, "1panel-rule:first", 1), - executorObservedRule(second, "1panel-rule:second", 2), - }) - adapter.capabilities = filter.Capabilities{ - Scopes: filter.MVPScopePatterns(), Marker: true, ExplicitPosition: true, - } - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - for _, rule := range []filter.FirewallRule{first, second} { - record, err := firewallRuleModelForCreate(rule, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}, constant.FirewallRuleOriginCreated) - if err != nil { - t.Fatalf("build managed rule: %v", err) - } - record.UUID = rule.UUID - if err := ruleRepo.Create(context.Background(), &record); err != nil { - t.Fatalf("create managed record: %v", err) - } - } - target := int64(2) - updated := first - updated.DestinationPort = "8443" - updated.OrderIndex = &target - if err := executor.updateRule(context.Background(), "", first.UUID, updated); err != nil { - t.Fatalf("explicit-position update: %v", err) - } - if adapter.snapshot.Rules[1].Marker != "1panel-rule:first" || adapter.snapshot.Rules[1].Rule.DestinationPort != "8443" { - t.Fatalf("rule was not updated at requested position: snapshot=%#v", adapter.snapshot) - } -} - -func TestFirewallExecutorRejectsPriorityReorderWithoutExplicitPriority(t *testing.T) { - rule := filter.FirewallRule{ - UUID: "zone-port", Scope: filter.Scope{Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, Zone: "public", Direction: filter.DirectionInput}, - NativeKind: filter.NativeKindZonePort, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, - } - adapter := newFakeFilterAdapter(t, rule.Scope, []filter.ObservedRule{executorObservedRule(rule, "1panel-rule:zone-port", 1)}) - adapter.capabilities = filter.Capabilities{Scopes: filter.MVPScopePatterns(), Marker: true, ExplicitPriority: true} - executor, ruleRepo := newTestFirewallExecutor(t, adapter) - record, err := firewallRuleModelForCreate(rule, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}, constant.FirewallRuleOriginCreated) - if err != nil { - t.Fatalf("build zone port record: %v", err) - } - record.UUID = rule.UUID - if err := ruleRepo.Create(context.Background(), &record); err != nil { - t.Fatalf("create zone port record: %v", err) - } - priority := -10 - err = executor.reorderRule(context.Background(), "", rule.UUID, nil, &priority) - if !errors.Is(err, filter.ErrUnsupportedScope) || adapter.applyCount != 0 { - t.Fatalf("native zone port was converted by reorder: applyCount=%d err=%v", adapter.applyCount, err) - } -} - -func TestMergeFirewallInventoryMapsPersistenceOwnershipAndUsage(t *testing.T) { - domainRule := filter.FirewallRule{ - Scope: filter.Scope{Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: filter.DirectionInput}, - NativeKind: filter.NativeKindRule, - Protocol: "tcp", - DestinationPort: "22", - ConnectionStates: []string{"established", "new"}, - Action: filter.ActionAccept, - Description: "ssh", - } - position := 1 - marker := "1panel-rule:managed" - stored, err := model.FirewallRuleFromDomain(domainRule) - if err != nil { - t.Fatalf("build stored rule: %v", err) - } - stored.UUID = "managed" - stored.Origin = constant.FirewallRuleOriginCreated - stored.Owner = constant.FirewallRuleSourceUser - observed := filter.ObservedRule{ - Rule: domainRule, - Locator: filter.Locator{Provider: filter.ProviderIptables, ScopeKey: domainRule.Scope.Key(), Position: &position}, - Marker: marker, ParseStatus: filter.ParseStatusSupported, - } - usage := map[string]filter.RuntimeUsage{filter.RuntimeUsageKey(domainRule): {UsedBy: []string{"sshd"}}} - - items, err := mergeFirewallInventory([]filter.ObservedRule{observed}, []model.FirewallRule{stored}, nil, usage) - if err != nil { - t.Fatalf("merge inventory: %v", err) - } - if len(items) != 1 || items[0].State != filter.InventoryStateManaged || items[0].Desired == nil || items[0].Usage == nil || !items[0].Usage.Used { - t.Fatalf("unexpected merged inventory: %#v", items) - } - if items[0].Desired.Rule.ConnectionStates[0] != "established" { - t.Fatalf("persistence metadata was not normalized: %#v", items[0].Desired.Rule) - } -} - -func TestProtectFirewallSnapshotMarksConfiguredAndRequiredPorts(t *testing.T) { - scope := filter.Scope{Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC_BEFORE", Direction: filter.DirectionInput} - positionOne, positionTwo := 1, 2 - rules := []filter.ObservedRule{ - { - Rule: filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "22", Action: filter.ActionAccept}, - Locator: filter.Locator{Provider: filter.ProviderIptables, ScopeKey: scope.Key(), Position: &positionOne}, ParseStatus: filter.ParseStatusSupported, - }, - { - Rule: filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}, - Locator: filter.Locator{Provider: filter.ProviderIptables, ScopeKey: scope.Key(), Position: &positionTwo}, ParseStatus: filter.ParseStatusSupported, - }, - } - snapshot, err := filter.NewSnapshot(scope, rules) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - snapshot.Notices = []filter.ScopeNotice{{Code: filter.ScopeNoticeManagedScopeInactive}} - - protected, err := filter.ProtectSnapshot(snapshot, []firewall.PortWhitelist{{Port: "22", Protocol: "tcp"}}) - if err != nil { - t.Fatalf("protect snapshot: %v", err) - } - if !protected.Rules[0].Protected || protected.Rules[1].Protected { - t.Fatalf("unexpected protected ports: %#v", protected.Rules) - } - if protected.Revision != snapshot.Revision { - t.Fatal("runtime safety classification changed the provider-state revision") - } - if len(protected.Notices) != 1 || protected.Notices[0].Code != filter.ScopeNoticeManagedScopeInactive { - t.Fatalf("snapshot notices were lost: %#v", protected.Notices) - } -} - -func newTestFirewallExecutor(t *testing.T, adapter *fakeFilterAdapter) (*FirewallService, *repo.FirewallRuleRepo) { - t.Helper() - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - return &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - adapter.Provider(): newFirewallRuleRuntime(adapter, nil), - }, - }, ruleRepo -} - -func createExecutorRule(executor *FirewallService, adapter *fakeFilterAdapter, request dto.FirewallRuleCreateItem) error { - authorization := firewallRuleCreateAuthorization{Operation: filter.ChangeCreate} - if request.AdoptInstanceKey != "" { - candidate, err := filter.FindCandidate(adapter.snapshot.Rules, request.AdoptInstanceKey) - if err != nil { - return err - } - locator := candidate.Locator - authorization = firewallRuleCreateAuthorization{Operation: filter.ChangeAdopt, Locator: &locator} - } - return executor.createRule( - context.Background(), newFirewallRuleRuntime(adapter, nil), adapter.snapshot, request, authorization, - ) -} - -type fakeFilterAdapter struct { - snapshot filter.Snapshot - multiSnapshots []filter.Snapshot - rollbackSnapshot filter.Snapshot - lastChange filter.DesiredChange - applyCount int - observeCount int - observeScopesCount int - rollbackCount int - verifyMatched bool - capabilities filter.Capabilities - nativeDetail string - nativeDetailName string - nativeDetailPermanent bool -} - -type failingFirewallRuleRepo struct { - repo.IFirewallRuleRepo - updateErr error - deleteErr error -} - -func (r *failingFirewallRuleRepo) UpdateWithRevision(ctx context.Context, ruleUUID string, revision uint, updates map[string]interface{}) error { - if r.updateErr != nil { - return r.updateErr - } - return r.IFirewallRuleRepo.UpdateWithRevision(ctx, ruleUUID, revision, updates) -} - -func (r *failingFirewallRuleRepo) DeleteWithRevision(ctx context.Context, ruleUUID string, revision uint) error { - if r.deleteErr != nil { - return r.deleteErr - } - return r.IFirewallRuleRepo.DeleteWithRevision(ctx, ruleUUID, revision) -} - -func newFakeFilterAdapter(t *testing.T, scope filter.Scope, rules []filter.ObservedRule) *fakeFilterAdapter { - t.Helper() - snapshot, err := filter.NewSnapshot(scope, rules) - if err != nil { - t.Fatalf("new fake snapshot: %v", err) - } - return &fakeFilterAdapter{snapshot: snapshot, verifyMatched: true} -} - -func (f *fakeFilterAdapter) Provider() filter.Provider { return f.snapshot.Scope.Provider } -func (f *fakeFilterAdapter) Capabilities(context.Context) (filter.Capabilities, error) { - if f.capabilities.Scopes != nil { - return f.capabilities, nil - } - return filter.Capabilities{Scopes: filter.MVPScopePatterns(), Marker: true, OwnedChains: true}, nil -} -func (f *fakeFilterAdapter) Observe(context.Context, filter.Scope) (filter.Snapshot, error) { - f.observeCount++ - return f.snapshot, nil -} -func (f *fakeFilterAdapter) ObserveScopes(context.Context, []filter.Scope) ([]filter.Snapshot, error) { - f.observeScopesCount++ - return append([]filter.Snapshot(nil), f.multiSnapshots...), nil -} -func (f *fakeFilterAdapter) Compile(snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, error) { - f.rollbackSnapshot = snapshot - f.rollbackSnapshot.Rules = append([]filter.ObservedRule(nil), snapshot.Rules...) - if len(changes) > 1 { - plan := filter.BackendPlan{ - Provider: f.Provider(), Scope: snapshot.Scope, SnapshotRevision: snapshot.Revision, - Rules: make([]filter.NativeRulePlan, 0, len(changes)), - } - position := len(snapshot.Rules) - for _, change := range changes { - if change.Operation != filter.ChangeCreate && change.Operation != filter.ChangeDelete { - return filter.BackendPlan{}, filter.ErrInvalidRule - } - rule := change.After - if change.Operation == filter.ChangeDelete { - rule = change.Before - } - if rule == nil { - return filter.BackendPlan{}, filter.ErrInvalidRule - } - position++ - if change.Locator != nil && change.Locator.Position != nil { - position = *change.Locator.Position - } - marker := "1panel-rule:" + rule.UUID - expected := executorObservedRule(*rule, marker, position) - plan.Rules = append(plan.Rules, filter.NativeRulePlan{ - RuleUUID: rule.UUID, Operation: change.Operation, Expected: expected, - }) - } - f.lastChange = changes[len(changes)-1] - return plan, nil - } - change := changes[0] - f.lastChange = change - rule := change.After - if change.Operation == filter.ChangeDelete { - rule = change.Before - } - marker := "1panel-rule:" + rule.UUID - position := len(snapshot.Rules) + 1 - if change.Locator != nil && change.Locator.Position != nil { - position = *change.Locator.Position - } - if rule.OrderIndex != nil && (change.Operation == filter.ChangeCreate || change.Operation == filter.ChangeReorder || change.Operation == filter.ChangeUpdate) { - position = int(*rule.OrderIndex) - } - expected := executorObservedRule(*rule, marker, position) - return filter.BackendPlan{ - Provider: f.Provider(), Scope: snapshot.Scope, SnapshotRevision: snapshot.Revision, - Rules: []filter.NativeRulePlan{{RuleUUID: rule.UUID, Operation: change.Operation, Expected: expected}}, - }, nil -} -func (f *fakeFilterAdapter) Apply(_ context.Context, plan filter.BackendPlan) (filter.ApplyResult, error) { - f.applyCount++ - if len(plan.Rules) > 1 { - applied := make([]filter.ObservedRule, 0, len(plan.Rules)) - rules := append([]filter.ObservedRule(nil), f.snapshot.Rules...) - if plan.Rules[0].Operation == filter.ChangeDelete { - deleted := make(map[string]struct{}, len(plan.Rules)) - for _, rulePlan := range plan.Rules { - deleted[rulePlan.Expected.Marker] = struct{}{} - } - remaining := make([]filter.ObservedRule, 0, len(rules)-len(deleted)) - for _, observed := range rules { - if _, exists := deleted[observed.Marker]; !exists { - remaining = append(remaining, observed) - } - } - f.snapshot, _ = filter.NewSnapshot(f.snapshot.Scope, remaining) - return filter.ApplyResult{}, nil - } - for _, rulePlan := range plan.Rules { - if rulePlan.Operation != filter.ChangeCreate { - return filter.ApplyResult{}, filter.ErrInvalidRule - } - rules = append(rules, rulePlan.Expected) - applied = append(applied, rulePlan.Expected) - } - f.snapshot, _ = filter.NewSnapshot(f.snapshot.Scope, rules) - return filter.ApplyResult{Applied: applied}, nil - } - expected := plan.Rules[0].Expected - if plan.Rules[0].Operation == filter.ChangeReorder || plan.Rules[0].Operation == filter.ChangeUpdate { - current := -1 - for index, observed := range f.snapshot.Rules { - if observed.Marker == expected.Marker { - current = index - break - } - } - if current < 0 || expected.Locator.Position == nil { - return filter.ApplyResult{}, filter.ErrRuleStale - } - rules := append([]filter.ObservedRule(nil), f.snapshot.Rules...) - moving := rules[current] - rules = append(rules[:current], rules[current+1:]...) - target := *expected.Locator.Position - 1 - if target > len(rules) { - target = len(rules) - } - rules = append(rules, filter.ObservedRule{}) - copy(rules[target+1:], rules[target:]) - moving.Rule = expected.Rule - rules[target] = moving - for index := range rules { - position := index + 1 - rules[index].Locator.Position = &position - } - f.snapshot, _ = filter.NewSnapshot(f.snapshot.Scope, rules) - return filter.ApplyResult{Applied: []filter.ObservedRule{rules[target]}}, nil - } - if plan.Rules[0].Operation == filter.ChangeDelete { - remaining := make([]filter.ObservedRule, 0, len(f.snapshot.Rules)) - for _, observed := range f.snapshot.Rules { - if observed.Marker != expected.Marker { - remaining = append(remaining, observed) - } - } - f.snapshot, _ = filter.NewSnapshot(f.snapshot.Scope, remaining) - return filter.ApplyResult{}, nil - } - replaced := false - for index := range f.snapshot.Rules { - if f.snapshot.Rules[index].Locator.Position != nil && expected.Locator.Position != nil && *f.snapshot.Rules[index].Locator.Position == *expected.Locator.Position { - f.snapshot.Rules[index] = expected - replaced = true - break - } - } - if !replaced { - f.snapshot.Rules = append(f.snapshot.Rules, expected) - } - f.snapshot, _ = filter.NewSnapshot(f.snapshot.Scope, f.snapshot.Rules) - return filter.ApplyResult{Applied: []filter.ObservedRule{expected}}, nil -} -func (f *fakeFilterAdapter) Verify(context.Context, filter.BackendPlan) (filter.VerifyResult, error) { - return filter.VerifyResult{Snapshot: f.snapshot, Matched: f.verifyMatched}, nil -} -func (f *fakeFilterAdapter) Rollback(context.Context, filter.BackendPlan) error { - f.rollbackCount++ - f.snapshot = f.rollbackSnapshot - return nil -} - -func (f *fakeFilterAdapter) NativeDetail(_ context.Context, name string, permanent bool) (string, error) { - f.nativeDetailName = name - f.nativeDetailPermanent = permanent - return f.nativeDetail, nil -} - -func firewallRuleCheckItems(rules []filter.FirewallRule) []dto.FirewallRuleCheckItem { - items := make([]dto.FirewallRuleCheckItem, 0, len(rules)) - for _, rule := range rules { - items = append(items, dto.FirewallRuleCheckItem{Rule: rule}) - } - return items -} - -func executorTestRule(port string) filter.FirewallRule { - return filter.FirewallRule{ - Scope: filter.Scope{Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: filter.DirectionInput}, - NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: port, Action: filter.ActionAccept, - } -} - -func executorTestAddressRule(address string) filter.FirewallRule { - rule := executorTestRule("") - rule.Protocol = "all" - rule.SourceAddress = address - rule.Action = filter.ActionDrop - return rule -} - -func executorObservedRule(rule filter.FirewallRule, marker string, position int) filter.ObservedRule { - return filter.ObservedRule{ - Rule: rule, Marker: marker, ParseStatus: filter.ParseStatusSupported, - Locator: filter.Locator{Provider: rule.Scope.Provider, ScopeKey: rule.Scope.Key(), Position: &position}, - } -} - -func newFirewallRuleTestDB(t *testing.T) *gorm.DB { - t.Helper() - dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString()) - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) - if err != nil { - t.Fatalf("open sqlite: %v", err) - } - sqlDB, err := db.DB() - if err != nil { - t.Fatalf("load sql db: %v", err) - } - sqlDB.SetMaxOpenConns(1) - t.Cleanup(func() { _ = sqlDB.Close() }) - if err := db.AutoMigrate(&model.FirewallRule{}); err != nil { - t.Fatalf("migrate firewall rule: %v", err) - } - return db -} diff --git a/agent/app/service/firewall_setting.go b/agent/app/service/firewall_setting.go index f18774275f56..608d9923c101 100644 --- a/agent/app/service/firewall_setting.go +++ b/agent/app/service/firewall_setting.go @@ -2,6 +2,7 @@ package service import ( "context" + "errors" "fmt" "strings" @@ -24,7 +25,9 @@ type IFirewallSettingService interface { type FirewallSettingService struct{} -func NewIFirewallSettingService() IFirewallSettingService { return &FirewallSettingService{} } +func NewIFirewallSettingService() IFirewallSettingService { + return &FirewallSettingService{} +} func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings, error) { result := dto.FirewallSettings{PingStatus: ping.LoadStatus()} @@ -132,10 +135,7 @@ func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings Name: name, Installed: installed[name], Supported: dockerInstalled, Active: dockerInstalled && installed[name] && result.Docker.Selected == name, } - var guard dockerGuardRuntime = docker_guard.NewManager() - if name == constant.FirewallProviderNftables { - guard = docker_guard.NewNftablesManager() - } + guard := docker_guard.NewRuntime(name) ipv4, ipv6 := guard.Status(docker_guard.FamilyIPv4), guard.Status(docker_guard.FamilyIPv6) option.Initialized = ipv4.Initialized || ipv6.Initialized option.Bound = ipv4.Bound || ipv6.Bound @@ -196,10 +196,7 @@ func (s *FirewallSettingService) Operate(ctx context.Context, request dto.Firewa } func (s *FirewallSettingService) operateDocker(ctx context.Context, request dto.FirewallBackendOperation) error { - var guard dockerGuardRuntime = docker_guard.NewManager() - if request.Backend == constant.FirewallProviderNftables { - guard = docker_guard.NewNftablesManager() - } + guard := docker_guard.NewRuntime(request.Backend) if request.Operation == "cleanup" { if err := guard.Cleanup(); err != nil { return err @@ -224,7 +221,7 @@ func (s *FirewallSettingService) operateDocker(ctx context.Context, request dto. return err } if request.Operation == "initialize" { - if err := NewIDockerPortGuardService().Operate(ctx, dto.DockerPortGuardOperation{Operation: "initialize"}); err != nil { + if err := newDockerPortGuardService().Operate(ctx, dto.DockerPortGuardOperation{Operation: "initialize"}); err != nil { _ = settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, previous) return err } @@ -233,10 +230,7 @@ func (s *FirewallSettingService) operateDocker(ctx context.Context, request dto. } func dockerGuardBackendInitialized(backend string) (bool, error) { - var guard dockerGuardRuntime = docker_guard.NewManager() - if backend == constant.FirewallProviderNftables { - guard = docker_guard.NewNftablesManager() - } + guard := docker_guard.NewRuntime(backend) for _, family := range []string{docker_guard.FamilyIPv4, docker_guard.FamilyIPv6} { initialized, err := guard.Initialized(family) if err != nil { @@ -372,7 +366,7 @@ func (s *FirewallSettingService) operateForwarding(request dto.FirewallBackendOp return err } if request.Operation == "initialize" { - return NewIForwardingService().Enable() + return newForwardingService().Enable() } recordForwardingSyncError(nil) return nil @@ -381,7 +375,7 @@ func (s *FirewallSettingService) operateForwarding(request dto.FirewallBackendOp func forwardingBackendInitialized(backend string) (bool, error) { manager, err := newForwardingManagerFor(backend) if err != nil { - if strings.Contains(err.Error(), "is not installed") { + if errors.Is(err, lifecycle.ErrNotInstalled) { return false, nil } return false, err diff --git a/agent/app/service/firewall_setting_test.go b/agent/app/service/firewall_setting_test.go deleted file mode 100644 index 079f96f0be84..000000000000 --- a/agent/app/service/firewall_setting_test.go +++ /dev/null @@ -1,174 +0,0 @@ -package service - -import ( - "context" - "fmt" - "os" - "path/filepath" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/global" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" - "github.com/glebarez/sqlite" - "github.com/google/uuid" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -func TestFirewallSettingServiceKeepsSelectionsWhenBackendsAreMissing(t *testing.T) { - previousDB := global.DB - dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString()) - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) - if err != nil { - t.Fatalf("open settings database: %v", err) - } - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatalf("migrate settings database: %v", err) - } - global.DB = db - t.Cleanup(func() { global.DB = previousDB }) - - binDir := t.TempDir() - if err := os.WriteFile(filepath.Join(binDir, "docker"), []byte("#!/bin/sh\nexit 0\n"), 0o755); err != nil { - t.Fatalf("create fake Docker executable: %v", err) - } - t.Setenv("PATH", binDir) - for key, backend := range map[string]string{ - constant.FirewallSystemBackendKey: constant.FirewallProviderNftables, - constant.FirewallForwardingBackendKey: constant.FirewallProviderNftables, - constant.FirewallDockerBackendKey: constant.FirewallProviderNftables, - } { - if err := settingRepo.UpdateOrCreate(key, backend); err != nil { - t.Fatalf("save selected backend %s: %v", key, err) - } - } - - settings, err := (&FirewallSettingService{}).Load(context.Background()) - if err != nil { - t.Fatalf("load firewall settings: %v", err) - } - for subsystem, selected := range map[string]string{ - "system": settings.System.Selected, - "forwarding": settings.Forwarding.Selected, - "docker": settings.Docker.Selected, - } { - if selected != constant.FirewallProviderNftables { - t.Fatalf("%s selected backend = %q, want persisted nftables", subsystem, selected) - } - } - for _, option := range settings.Docker.Options { - if option.Installed || option.Active || !option.Supported { - t.Fatalf("unexpected Docker backend availability without firewall commands: %#v", option) - } - } -} - -func TestFirewallSettingServiceSeparatesDockerAndFirewallAvailability(t *testing.T) { - previousDB := global.DB - dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString()) - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) - if err != nil { - t.Fatalf("open settings database: %v", err) - } - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatalf("migrate settings database: %v", err) - } - global.DB = db - t.Cleanup(func() { global.DB = previousDB }) - - binDir := t.TempDir() - for _, name := range []string{"iptables", "iptables-restore"} { - if err := os.WriteFile(filepath.Join(binDir, name), []byte("#!/bin/sh\nexit 0\n"), 0o755); err != nil { - t.Fatalf("create fake %s executable: %v", name, err) - } - } - t.Setenv("PATH", binDir) - - settings, err := (&FirewallSettingService{}).Load(context.Background()) - if err != nil { - t.Fatalf("load firewall settings: %v", err) - } - for _, option := range settings.Docker.Options { - if option.Name != constant.FirewallProviderIptables { - continue - } - if !option.Installed || option.Supported || option.Active { - t.Fatalf("Docker and firewall availability were not separated: %#v", option) - } - return - } - t.Fatal("iptables Docker backend option was not returned") -} - -func TestFirewallSettingServiceUsesExplicitEmptyDockerAndIptablesForwardingDefaults(t *testing.T) { - previousDB := global.DB - dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString()) - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) - if err != nil { - t.Fatalf("open settings database: %v", err) - } - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatalf("migrate settings database: %v", err) - } - global.DB = db - t.Cleanup(func() { global.DB = previousDB }) - t.Setenv("PATH", t.TempDir()) - - settings, err := (&FirewallSettingService{}).Load(context.Background()) - if err != nil { - t.Fatalf("load firewall settings: %v", err) - } - if settings.Docker.Selected != "" { - t.Fatalf("Docker selected backend = %q, want empty", settings.Docker.Selected) - } - if settings.Forwarding.Selected != constant.FirewallProviderIptables { - t.Fatalf("forwarding selected backend = %q, want iptables", settings.Forwarding.Selected) - } -} - -func TestFirewallSettingServiceSelectsGuardBackendWithoutDockerRestart(t *testing.T) { - previousDB := global.DB - dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString()) - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) - if err != nil { - t.Fatalf("open settings database: %v", err) - } - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatalf("migrate settings database: %v", err) - } - global.DB = db - t.Cleanup(func() { global.DB = previousDB }) - - err = (&FirewallSettingService{}).Operate(context.Background(), dto.FirewallBackendOperation{ - Subsystem: "docker", - Backend: "nftables", - Operation: "select", - }) - if err != nil { - t.Fatalf("select Docker guard backend: %v", err) - } - if got := selectedDockerFirewallBackend("iptables"); got != "nftables" { - t.Fatalf("selected Docker guard backend = %q, want nftables", got) - } - if _, ok := (&DockerPortGuardService{}).guardRuntime(selectedDockerFirewallBackend("iptables")).(*docker_guard.NftablesManager); !ok { - t.Fatal("Docker port guard did not use the selected nftables runtime") - } -} - -func TestFirewallSettingServiceRejectsServiceBackendInitialization(t *testing.T) { - for _, backend := range []string{"firewalld", "ufw"} { - for _, operation := range []string{"initialize", "cleanup"} { - err := (&FirewallSettingService{}).Operate(context.Background(), dto.FirewallBackendOperation{ - Subsystem: "system", - Backend: backend, - Operation: operation, - }) - if err == nil { - t.Fatalf("%s %s unexpectedly succeeded", backend, operation) - } - } - } -} diff --git a/agent/app/service/firewall_sync.go b/agent/app/service/firewall_sync.go index cabddceae802..d8bac03bd0d1 100644 --- a/agent/app/service/firewall_sync.go +++ b/agent/app/service/firewall_sync.go @@ -4,7 +4,6 @@ import ( "context" "errors" "fmt" - "slices" "sort" "strings" "sync" @@ -20,6 +19,7 @@ import ( "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" + firewallsync "github.com/1Panel-dev/1Panel/agent/utils/firewall/sync" "gorm.io/gorm" ) @@ -29,9 +29,9 @@ var ( ) const ( - firewallRuleSyncReady = "ready" - firewallRuleSyncExisting = "existing" - firewallRuleSyncBlocked = "blocked" + firewallRuleSyncReady = firewallsync.StatusReady + firewallRuleSyncExisting = firewallsync.StatusExisting + firewallRuleSyncBlocked = firewallsync.StatusBlocked ) func firewallSyncSubsystem(value string) string { @@ -42,13 +42,13 @@ func firewallSyncSubsystem(value string) string { return value } -type firewallRuleSyncOutcome string +type firewallRuleSyncOutcome = firewallsync.Outcome const ( - firewallRuleSyncApplied firewallRuleSyncOutcome = "applied" - firewallRuleSyncSkipped firewallRuleSyncOutcome = "skipped" - firewallRuleSyncRemoved firewallRuleSyncOutcome = "removed" - firewallRuleSyncFailed firewallRuleSyncOutcome = "failed" + firewallRuleSyncApplied = firewallsync.OutcomeApplied + firewallRuleSyncSkipped = firewallsync.OutcomeSkipped + firewallRuleSyncRemoved = firewallsync.OutcomeRemoved + firewallRuleSyncFailed = firewallsync.OutcomeFailed ) type firewallRuleSyncEntry struct { @@ -146,7 +146,7 @@ func (s *FirewallService) loadFirewallRuleSyncPlan( expectedMarkers[scopeKey+"\x00"+entry.desired.Marker] = struct{}{} } - scopes := firewallRuleSyncScopes(target) + scopes := filter.ManagedInputScopes(target) if hasCompileErrors { scopes = scopesWithFirewallSyncCandidates(scopes, entriesByScope) } @@ -203,14 +203,18 @@ func (s *FirewallService) loadFirewallRuleSyncPlan( source: model.FirewallRule{UUID: strings.TrimPrefix(observed.Marker, "1panel-rule:")}, rule: observed.Rule, remove: ©, } - status, reason := "remove", "managed rule exists only in target backend" + status := firewallsync.StatusRemove + reasonCode := firewallsync.ReasonManagedOnlyInTarget + reason := firewallsync.ReasonMessage(reasonCode) if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque { - status, reason = firewallRuleSyncBlocked, "managed runtime rule cannot be safely removed" + status = firewallRuleSyncBlocked + reasonCode = firewallsync.ReasonUnsafeRemoval + reason = firewallsync.ReasonMessage(reasonCode) entry.err = filter.ErrProtectedRule } entry.item = dto.FirewallRuleSyncItem{ SourceUUID: entry.source.UUID, Rule: &entry.rule, - Status: status, Reason: reason, + Status: status, ReasonCode: reasonCode, Reason: reason, } entries = append(entries, entry) scopeEntries = append(scopeEntries, entry) @@ -232,76 +236,20 @@ func (s *FirewallService) loadFirewallRuleSyncPlan( }, nil } -func firewallProviderHasOrderedManagedRules(provider filter.Provider) bool { - return provider == filter.ProviderIptables || provider == filter.ProviderNftables || provider == filter.ProviderUFW -} - func planFirewallManagedOrder( snapshot filter.Snapshot, entries []*firewallRuleSyncEntry, ) { - if !firewallProviderHasOrderedManagedRules(snapshot.Scope.Provider) { - return - } byMarker := make(map[string]*firewallRuleSyncEntry, len(entries)) + desiredMarkers := make([]string, 0, len(entries)) for _, entry := range entries { if entry.remove == nil && entry.err == nil && entry.desired.Marker != "" { byMarker[entry.desired.Marker] = entry + desiredMarkers = append(desiredMarkers, entry.desired.Marker) } } - if len(byMarker) < 2 { - return - } - - actual := make([]string, 0, len(byMarker)) - segments := make(map[string]int, len(byMarker)) - segment := 0 - for _, observed := range snapshot.Rules { - _, expected := byMarker[observed.Marker] - if expected { - actual = append(actual, observed.Marker) - if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque { - segment++ - segments[observed.Marker] = segment - segment++ - } else { - segments[observed.Marker] = segment - } - continue - } - if strings.HasPrefix(observed.Marker, "1panel-rule:") && - !observed.Protected && observed.ParseStatus != filter.ParseStatusOpaque { - continue - } - segment++ - } - - desired := make([]string, 0, len(actual)) - for _, entry := range entries { - marker := entry.desired.Marker - if _, exists := segments[marker]; exists { - desired = append(desired, marker) - } - } - if slices.Equal(actual, desired) { - return - } - drifted := make(map[string]struct{}, len(desired)) - for index := range desired { - if actual[index] != desired[index] { - drifted[actual[index]] = struct{}{} - drifted[desired[index]] = struct{}{} - } - } - feasible, previousSegment := true, -1 - for _, marker := range desired { - if segments[marker] < previousSegment { - feasible = false - break - } - previousSegment = segments[marker] - } - for _, marker := range desired { + drifted, feasible := firewallsync.ManagedOrderDrift(snapshot, desiredMarkers) + for _, marker := range desiredMarkers { if _, exists := drifted[marker]; !exists { continue } @@ -529,7 +477,7 @@ func (s *FirewallService) loadStoredFirewallRuleSyncCandidates( if err != nil { return "", nil, false, err } - sortFirewallPolicies(stored, selected) + model.SortFirewallRules(stored, selected) entries := make([]*firewallRuleSyncEntry, 0, len(stored)) hasCompileErrors := false for _, record := range stored { @@ -558,137 +506,12 @@ func scopesWithFirewallSyncCandidates(scopes []filter.Scope, candidates map[stri return result } -func firewallRuleSyncScopes(provider filter.Provider) []filter.Scope { - base := filter.Scope{Provider: provider, Direction: filter.DirectionInput} - switch provider { - case filter.ProviderIptables, filter.ProviderNftables: - result := make([]filter.Scope, 0, 6) - for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} { - for _, chain := range []string{filter.BasicBeforeChain, filter.IptablesInputChain, filter.BasicAfterChain} { - scope := base - scope.Family, scope.Table, scope.Chain = family, "filter", chain - result = append(result, scope) - } - } - return result - case filter.ProviderFirewalld: - base.Family, base.Zone = filter.FamilyInet, filter.FirewalldInputZone - return []filter.Scope{base} - case filter.ProviderUFW: - result := make([]filter.Scope, 0, 2) - for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} { - scope := base - scope.Family, scope.Chain = family, filter.UFWInputChain - result = append(result, scope) - } - return result - default: - return nil - } -} - -func firewallPolicyRulesForProvider(stored model.FirewallRule, provider filter.Provider) ([]filter.FirewallRule, error) { - if stored.CompatibilityError != "" { - return nil, fmt.Errorf("%w: %s", filter.ErrUnsupportedScope, stored.CompatibilityError) - } - connectionStates := make([]string, 0) - if stored.ConnectionStates != "" { - connectionStates = strings.Split(stored.ConnectionStates, ",") - } - base := filter.FirewallRule{ - Protocol: stored.Protocol, SourceAddress: stored.SourceAddress, SourcePort: stored.SourcePort, - DestinationAddress: stored.DestinationAddress, DestinationPort: stored.DestinationPort, - Interface: stored.Interface, ConnectionStates: connectionStates, - Action: filter.Action(stored.Action), Description: stored.Description, - } - if provider == filter.ProviderFirewalld { - base.Priority = stored.Priority - } - families := []filter.Family{filter.Family(stored.Family)} - if provider != filter.ProviderFirewalld && len(families) == 1 && families[0] == filter.FamilyInet { - hasIPv4, hasIPv6 := firewallRuleAddressFamilies(base) - switch { - case hasIPv4 && hasIPv6: - return nil, fmt.Errorf("%w: inet policy contains both IPv4 and IPv6 addresses", filter.ErrUnsupportedScope) - case hasIPv6 || strings.EqualFold(base.Protocol, "icmpv6"): - families = []filter.Family{filter.FamilyIPv6} - case hasIPv4: - families = []filter.Family{filter.FamilyIPv4} - default: - families = []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} - } - } - result := make([]filter.FirewallRule, 0, len(families)) - for _, family := range families { - rule := base - rule.Scope = filter.Scope{Provider: provider, Family: family, Direction: filter.DirectionInput} - switch provider { - case filter.ProviderIptables, filter.ProviderNftables: - rule.Scope.Table, rule.Scope.Chain = "filter", filter.IptablesInputChain - case filter.ProviderFirewalld: - rule.Scope.Zone = filter.FirewalldInputZone - case filter.ProviderUFW: - rule.Scope.Chain = filter.UFWInputChain - default: - return nil, fmt.Errorf("%w: unsupported firewall provider %q", filter.ErrProviderUnavailable, provider) - } - expanded, err := filter.ExpandAtomicRules(rule) - if err != nil { - return nil, err - } - result = append(result, expanded...) - } - return result, nil -} - -func sortFirewallPolicies(stored []model.FirewallRule, provider filter.Provider) { - sort.SliceStable(stored, func(i, j int) bool { - left, right := stored[i], stored[j] - if provider == filter.ProviderFirewalld { - switch { - case left.Priority == nil && right.Priority != nil: - return false - case left.Priority != nil && right.Priority == nil: - return true - case left.Priority != nil && right.Priority != nil && *left.Priority != *right.Priority: - return *left.Priority < *right.Priority - } - } else { - switch { - case left.Sequence == nil && right.Sequence != nil: - return false - case left.Sequence != nil && right.Sequence == nil: - return true - case left.Sequence != nil && right.Sequence != nil && *left.Sequence != *right.Sequence: - return *left.Sequence < *right.Sequence - } - } - return left.UUID < right.UUID - }) -} - -func firewallRuleAddressFamilies(rule filter.FirewallRule) (bool, bool) { - hasIPv4, hasIPv6 := false, false - for _, address := range []string{rule.SourceAddress, rule.DestinationAddress} { - address = strings.TrimSpace(address) - if address == "" { - continue - } - if strings.Contains(address, ":") { - hasIPv6 = true - } else { - hasIPv4 = true - } - } - return hasIPv4, hasIPv6 -} - func (s *FirewallService) classifyFirewallRuleSyncCandidate( clientIP string, snapshot filter.Snapshot, entry *firewallRuleSyncEntry, item filter.InventoryItem, -) (string, string) { +) (firewallsync.Status, string) { switch item.Match { case filter.InventoryMatchExact: return firewallRuleSyncExisting, "rule already matches database policy" @@ -818,64 +641,24 @@ func buildDatabaseSyncPlan[T any]( key func(T) string, actualItem func(T) dto.FirewallRuleSyncItem, ) databaseSyncPlan { - items := make([]dto.FirewallRuleSyncItem, 0, len(desired)+len(actual)) - actualByKey := make(map[string][]int, len(actual)) - for index, value := range actual { - actualByKey[key(value)] = append(actualByKey[key(value)], index) - } - matched := make([]bool, len(actual)) + candidates := make([]firewallsync.Desired[T, dto.FirewallRuleSyncItem], 0, len(desired)) for _, candidate := range desired { - item := candidate.item - switch { - case candidate.err != nil: - item.Status, item.Reason = firewallRuleSyncBlocked, candidate.err.Error() - default: - match := unmatchedDatabaseSyncIndex(actualByKey[key(candidate.value)], matched) - if match >= 0 { - matched[match] = true - item.Status, item.Reason = firewallRuleSyncExisting, "rule already exists in target backend" - } else { - item.Status = firewallRuleSyncReady - } - } - items = append(items, item) + candidates = append(candidates, firewallsync.Desired[T, dto.FirewallRuleSyncItem]{ + Value: candidate.value, Payload: candidate.item, Err: candidate.err, + }) } - for index, value := range actual { - if matched[index] { - continue - } - item := actualItem(value) - item.Status, item.Reason = "remove", "rule exists only in target backend" + diff := firewallsync.Diff(candidates, actual, key, actualItem) + items := make([]dto.FirewallRuleSyncItem, 0, len(diff)) + for _, diffItem := range diff { + item := diffItem.Payload + item.Status, item.ReasonCode, item.Reason = diffItem.Status, diffItem.ReasonCode, diffItem.Reason items = append(items, item) } return databaseSyncPlan{subsystem: subsystem, target: target, items: items} } -func unmatchedDatabaseSyncIndex(indices []int, matched []bool) int { - for _, index := range indices { - if !matched[index] { - return index - } - } - return -1 -} - func databaseSyncStatesEqual[T any](left, right []T, key func(T) string) bool { - if len(left) != len(right) { - return false - } - counts := make(map[string]int, len(left)) - for _, value := range left { - counts[key(value)]++ - } - for _, value := range right { - valueKey := key(value) - if counts[valueKey] == 0 { - return false - } - counts[valueKey]-- - } - return true + return firewallsync.StatesEqual(left, right, key) } func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview { @@ -894,7 +677,7 @@ func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview { case firewallRuleSyncBlocked: result.Blocked++ result.Total++ - case "remove": + case firewallsync.StatusRemove: result.Removed++ } } @@ -904,7 +687,7 @@ func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview { func (p databaseSyncPlan) baseResult() dto.FirewallRuleSyncResult { result := dto.FirewallRuleSyncResult{Subsystem: p.subsystem, TargetProvider: p.target} for _, item := range p.items { - if item.Status != "remove" { + if item.Status != firewallsync.StatusRemove { result.Total++ } } @@ -921,7 +704,7 @@ func (p databaseSyncPlan) completedResult() dto.FirewallRuleSyncResult { result.Skipped++ case firewallRuleSyncBlocked: appendDatabaseSyncFailure(&result, item, errors.New(item.Reason)) - case "remove": + case firewallsync.StatusRemove: result.Removed++ } } @@ -1116,10 +899,10 @@ func (s *DockerPortGuardService) syncRules( } plan := buildDockerDatabaseSyncPlan(filter.Provider(target), policies, targetPolicies) result, reconcileErr := plan.reconcile(func() error { - if err := reconcileDockerGuardSyncTarget(target, runtimePolicies, targetRuntime); err != nil { + if err := docker_guard.ReconcileTarget(target, runtimePolicies, targetRuntime); err != nil { return err } - if err := verifyDockerGuardRuleSync(targetRuntime, runtimePolicies); err != nil { + if err := docker_guard.Verify(targetRuntime, runtimePolicies); err != nil { return err } if len(policies) == 0 { @@ -1153,7 +936,7 @@ func buildDockerDatabaseSyncPlan( }) } return buildDatabaseSyncPlan( - "docker", target, desired, actual, dockerGuardPolicySyncKey, + "docker", target, desired, actual, docker_guard.PolicySyncKey, func(policy docker_guard.Policy) dto.FirewallRuleSyncItem { return dto.FirewallRuleSyncItem{SourceUUID: policy.UUID, DockerRule: dockerGuardRuntimeRuleSyncDTO(policy)} }, @@ -1178,7 +961,7 @@ func (r *firewallScopeReconciler) reconcile() { continue } switch entry.item.Status { - case "remove": + case firewallsync.StatusRemove: removes = append(removes, entry) case firewallRuleSyncReady: switch entry.match { @@ -1217,7 +1000,7 @@ func (r *firewallScopeReconciler) applyGroups(entries []*firewallRuleSyncEntry) if len(entries) == 0 { return } - batch := r.runtime.adapter.Provider() == filter.ProviderIptables || r.runtime.adapter.Provider() == filter.ProviderNftables + batch := r.runtime.Provider() == filter.ProviderIptables || r.runtime.Provider() == filter.ProviderNftables if batch { r.apply(entries) return @@ -1237,7 +1020,7 @@ func (r *firewallScopeReconciler) apply(entries []*firewallRuleSyncEntry) { continue } if !changed { - if entry.item.Status == "remove" { + if entry.item.Status == firewallsync.StatusRemove { entry.outcome = firewallRuleSyncRemoved } else { entry.outcome = firewallRuleSyncSkipped @@ -1261,7 +1044,7 @@ func (r *firewallScopeReconciler) apply(entries []*firewallRuleSyncEntry) { return } for _, entry := range active { - if entry.item.Status == "remove" { + if entry.item.Status == firewallsync.StatusRemove { entry.outcome = firewallRuleSyncRemoved } else { entry.outcome = firewallRuleSyncApplied @@ -1294,7 +1077,7 @@ func (r *firewallScopeReconciler) restoreOrder() { changed := false for step := 0; step < len(desiredMarkers); step++ { - marker, position, converged, err := nextFirewallManagedOrderChange(r.snapshot, desiredMarkers) + marker, position, converged, err := firewallsync.NextManagedOrderChange(r.snapshot, desiredMarkers) if err != nil { r.failOrder(reorderEntries, err) return @@ -1309,18 +1092,18 @@ func (r *firewallScopeReconciler) restoreOrder() { } return } - observed, _, exists := firewallRuleSyncObservedByMarker(r.snapshot, marker) + observed, _, exists := firewallsync.ObservedByMarker(r.snapshot, marker) if !exists { r.failOrder(reorderEntries, filter.ErrRuleStale) return } - after := firewallRuleSyncObservedRule(observed) + after := firewallsync.ObservedRule(observed) target := int64(position) after.OrderIndex = &target - before := firewallRuleSyncObservedRule(observed) + before := firewallsync.ObservedRule(observed) locator := observed.Locator operation := filter.ChangeReorder - if r.runtime.adapter.Provider() == filter.ProviderUFW { + if r.runtime.Provider() == filter.ProviderUFW { operation = filter.ChangeUpdate } _, verification, executeErr := r.runtime.Execute(r.ctx, r.snapshot, []filter.DesiredChange{{ @@ -1339,59 +1122,6 @@ func (r *firewallScopeReconciler) restoreOrder() { r.failOrder(reorderEntries, fmt.Errorf("%w: managed rule order did not converge", filter.ErrVerificationFailed)) } -func nextFirewallManagedOrderChange( - snapshot filter.Snapshot, - desiredMarkers []string, -) (string, int, bool, error) { - expected := make(map[string]struct{}, len(desiredMarkers)) - for _, marker := range desiredMarkers { - expected[marker] = struct{}{} - } - actual := make([]string, 0, len(desiredMarkers)) - positions := make([]int, 0, len(desiredMarkers)) - for index, observed := range snapshot.Rules { - if _, exists := expected[observed.Marker]; !exists { - continue - } - actual = append(actual, observed.Marker) - position := index + 1 - if observed.Locator.Position != nil { - position = *observed.Locator.Position - } - positions = append(positions, position) - } - if len(actual) != len(desiredMarkers) { - return "", 0, false, filter.ErrRuleStale - } - for index := range desiredMarkers { - if actual[index] != desiredMarkers[index] { - return desiredMarkers[index], positions[index], false, nil - } - } - return "", 0, true, nil -} - -func firewallRuleSyncObservedByMarker(snapshot filter.Snapshot, marker string) (filter.ObservedRule, int, bool) { - for index, observed := range snapshot.Rules { - if observed.Marker == marker { - position := index + 1 - if observed.Locator.Position != nil { - position = *observed.Locator.Position - } - return observed, position, true - } - } - return filter.ObservedRule{}, 0, false -} - -func firewallRuleSyncObservedRule(observed filter.ObservedRule) filter.FirewallRule { - rule := observed.Rule - if rule.UUID == "" && strings.HasPrefix(observed.Marker, "1panel-rule:") { - rule.UUID = strings.TrimSpace(strings.TrimPrefix(observed.Marker, "1panel-rule:")) - } - return rule -} - func (r *firewallScopeReconciler) failOrder(entries []*firewallRuleSyncEntry, err error) { for _, entry := range entries { if entry.outcome != firewallRuleSyncFailed { @@ -1416,7 +1146,7 @@ func firewallRuleSyncChange( if observed.Protected { return filter.DesiredChange{}, false, filter.ErrProtectedRule } - before := firewallRuleSyncObservedRule(observed) + before := firewallsync.ObservedRule(observed) locator := observed.Locator return filter.DesiredChange{ Operation: filter.ChangeDelete, Before: &before, Locator: &locator, @@ -1452,7 +1182,7 @@ func firewallRuleSyncChange( if err := filter.GuardMutation(snapshot, *item.Observed, entry.rule, clientIP, protectedPorts...); err != nil { return filter.DesiredChange{}, false, err } - before := firewallRuleSyncObservedRule(*item.Observed) + before := firewallsync.ObservedRule(*item.Observed) locator := item.Observed.Locator change.Operation = filter.ChangeUpdate change.Before = &before @@ -1468,29 +1198,16 @@ func firewallRuleSyncInsertionPosition( entries []*firewallRuleSyncEntry, target *firewallRuleSyncEntry, ) *int64 { - if !firewallProviderHasOrderedManagedRules(snapshot.Scope.Provider) { - return nil - } - targetIndex := slices.Index(entries, target) - for index := targetIndex - 1; index >= 0; index-- { - entry := entries[index] + desiredMarkers := make([]string, 0, len(entries)) + for _, entry := range entries { if entry.remove != nil || entry.err != nil { continue } - if _, position, exists := firewallRuleSyncObservedByMarker(snapshot, entry.desired.Marker); exists { - value := int64(position + 1) - return &value - } + desiredMarkers = append(desiredMarkers, entry.desired.Marker) } - for index := targetIndex + 1; index < len(entries); index++ { - entry := entries[index] - if entry.remove != nil || entry.err != nil { - continue - } - if _, position, exists := firewallRuleSyncObservedByMarker(snapshot, entry.desired.Marker); exists { - value := int64(position) - return &value - } + position, exists := firewallsync.InsertionPosition(snapshot, desiredMarkers, target.desired.Marker) + if !exists { + return nil } - return nil + return &position } diff --git a/agent/app/service/firewall_sync_test.go b/agent/app/service/firewall_sync_test.go deleted file mode 100644 index 373f86da99b2..000000000000 --- a/agent/app/service/firewall_sync_test.go +++ /dev/null @@ -1,683 +0,0 @@ -package service - -import ( - "context" - "errors" - "path/filepath" - "testing" - "time" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/app/repo" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/global" - agenti18n "github.com/1Panel-dev/1Panel/agent/i18n" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - "github.com/glebarez/sqlite" - "github.com/go-playground/validator/v10" - "gorm.io/gorm" -) - -type fakeFirewallDatabaseSyncAdapter struct { - previewResult dto.FirewallRuleSyncPreview - syncResult dto.FirewallRuleSyncResult - previewErr error - syncErr error - previewCalls int - syncCalls int -} - -func (f *fakeFirewallDatabaseSyncAdapter) previewRuleSync( - context.Context, - dto.FirewallRuleSyncRequest, -) (dto.FirewallRuleSyncPreview, error) { - f.previewCalls++ - return f.previewResult, f.previewErr -} - -func (f *fakeFirewallDatabaseSyncAdapter) syncRules( - context.Context, - dto.FirewallRuleSyncRequest, -) (dto.FirewallRuleSyncResult, error) { - f.syncCalls++ - return f.syncResult, f.syncErr -} - -func TestFirewallRuleSyncCoordinatorDispatchesSubsystemAdapters(t *testing.T) { - forwardingErr := errors.New("forwarding preview") - dockerErr := errors.New("docker sync") - forwarding := &fakeFirewallDatabaseSyncAdapter{previewErr: forwardingErr} - docker := &fakeFirewallDatabaseSyncAdapter{syncErr: dockerErr} - service := &FirewallService{forwardingSync: forwarding, dockerSync: docker} - - _, err := service.PreviewRuleSync(context.Background(), "client-ip", dto.FirewallRuleSyncRequest{ - Subsystem: "forwarding", TargetProvider: filter.ProviderNftables, - }) - if !errors.Is(err, forwardingErr) { - t.Fatalf("forwarding preview was not dispatched to its adapter: %v", err) - } - _, err = service.SyncRules(context.Background(), "client-ip", dto.FirewallRuleSyncRequest{ - Subsystem: "docker", TargetProvider: filter.ProviderNftables, - }) - if !errors.Is(err, dockerErr) { - t.Fatalf("Docker synchronization was not dispatched to its adapter: %v", err) - } - if forwarding.previewCalls != 1 || forwarding.syncCalls != 0 || docker.previewCalls != 0 || docker.syncCalls != 1 { - t.Fatalf("unexpected adapter calls: forwarding=%#v docker=%#v", forwarding, docker) - } -} - -func TestFirewallRuleSyncRequestValidationAllowsDatabaseSource(t *testing.T) { - validate := validator.New() - for _, subsystem := range []string{"forwarding", "docker"} { - request := dto.FirewallRuleSyncRequest{Subsystem: subsystem, TargetProvider: filter.ProviderNftables} - if err := validate.Struct(request); err != nil { - t.Fatalf("%s database synchronization rejected missing source provider: %v", subsystem, err) - } - } - if err := validate.Struct(dto.FirewallRuleSyncRequest{Subsystem: "system", SourceProvider: filter.ProviderIptables}); err == nil { - t.Fatal("synchronization request without target provider was accepted") - } -} - -func TestFirewallRuleSyncFailureMessagesIncludeDetailsAndGroupDuplicates(t *testing.T) { - messages := firewallRuleSyncFailureMessages([]dto.FirewallRuleSyncFailure{ - {SourceUUID: "rule-1", Error: "iptables-restore failed: invalid port"}, - {SourceUUID: "rule-2", Error: "iptables-restore failed: invalid port"}, - {SourceUUID: "rule-3", Error: "permission denied"}, - }) - want := []string{ - "UUID [rule-1, rule-2]: iptables-restore failed: invalid port", - "UUID [rule-3]: permission denied", - } - if len(messages) != len(want) { - t.Fatalf("failure messages = %#v, want %#v", messages, want) - } - for index := range want { - if messages[index] != want[index] { - t.Fatalf("failure message %d = %q, want %q", index, messages[index], want[index]) - } - } -} - -func TestFirewallRuleSyncPreviewApplyAndRetry(t *testing.T) { - ctx := context.Background() - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - sourceRule := filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - }, - Protocol: "tcp", DestinationPort: "8443", Action: filter.ActionAccept, Description: "sync-test", - } - sourceRecord, err := model.FirewallRuleFromDomain(sourceRule) - if err != nil { - t.Fatalf("encode source rule: %v", err) - } - sourceRecord.Origin = constant.FirewallRuleOriginCreated - sourceRecord.Owner = model.FirewallRuleOwner(constant.FirewallRuleSourceApp, "test-app") - if err := ruleRepo.Create(ctx, &sourceRecord); err != nil { - t.Fatalf("persist source rule: %v", err) - } - - targetScope := filter.Scope{ - Provider: filter.ProviderNftables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - adapter := newFakeFilterAdapter(t, targetScope, nil) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil }, - } - request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables} - - preview, err := service.PreviewRuleSync(ctx, "", request) - if err != nil { - t.Fatalf("preview rule sync: %v", err) - } - if preview.Total != 1 || preview.Ready != 1 || preview.Existing != 0 || preview.Blocked != 0 { - t.Fatalf("unexpected preview: %#v", preview) - } - - result, err := service.syncRules(ctx, "", request) - if err != nil { - t.Fatalf("apply rule sync: %v", err) - } - if result.Total != 1 || result.Succeeded != 1 || result.Skipped != 0 || result.Failed != 0 { - t.Fatalf("unexpected sync result: %#v", result) - } - targetRecords, err := ruleRepo.List(ctx) - if err != nil { - t.Fatalf("list target records: %v", err) - } - if len(targetRecords) != 1 || targetRecords[0].Owner != sourceRecord.Owner { - t.Fatalf("target ownership was not preserved: %#v", targetRecords) - } - - retry, err := service.syncRules(ctx, "", request) - if err != nil { - t.Fatalf("retry rule sync: %v", err) - } - if retry.Succeeded != 0 || retry.Skipped != 1 || retry.Failed != 0 { - t.Fatalf("rule sync is not idempotent: %#v", retry) - } - - setupFirewallTaskTestDB(t) - taskResult, err := service.SyncRules(ctx, "", dto.FirewallRuleSyncRequest{ - Subsystem: "system", TargetProvider: filter.ProviderNftables, - TaskID: "firewall-sync-task", - }) - if err != nil { - t.Fatal(err) - } - if !taskResult.Queued || taskResult.TaskID != "firewall-sync-task" { - t.Fatalf("plain synchronization was not queued as a task: %#v", taskResult) - } - waitFirewallSyncTask(t, taskResult.TaskID) - sourceRecords, err := ruleRepo.List(ctx) - if err != nil || len(sourceRecords) != 1 { - t.Fatalf("plain synchronization reset the source backend: %#v err=%v", sourceRecords, err) - } -} - -func TestFirewallRuleSyncKeepsExpandedRulesInTheSameScope(t *testing.T) { - ctx := context.Background() - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - sourceRule := filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - }, - Protocol: "tcp", DestinationPort: "80,443", Action: filter.ActionAccept, - } - sourceRecord, err := model.FirewallRuleFromDomain(sourceRule) - if err != nil { - t.Fatal(err) - } - sourceRecord.Origin = constant.FirewallRuleOriginCreated - sourceRecord.Owner = constant.FirewallRuleSourceUser - if err := ruleRepo.Create(ctx, &sourceRecord); err != nil { - t.Fatal(err) - } - - targetScope := filter.Scope{ - Provider: filter.ProviderNftables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - adapter := newFakeFilterAdapter(t, targetScope, nil) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil }, - } - request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables} - - result, err := service.syncRules(ctx, "", request) - if err != nil || result.Succeeded != 2 || result.Failed != 0 { - t.Fatalf("expanded rule synchronization failed: result=%#v err=%v", result, err) - } - if len(adapter.snapshot.Rules) != 2 { - t.Fatalf("expanded rules overwrote each other: %#v", adapter.snapshot.Rules) - } - if adapter.applyCount != 1 { - t.Fatalf("same-scope rules were applied in %d calls, want one batch", adapter.applyCount) - } - if adapter.observeCount != len(firewallRuleSyncScopes(filter.ProviderNftables)) { - t.Fatalf("synchronization observed scopes %d times, want one read per scope", adapter.observeCount) - } - ports := make(map[string]struct{}, len(adapter.snapshot.Rules)) - markers := make(map[string]struct{}, len(adapter.snapshot.Rules)) - for _, observed := range adapter.snapshot.Rules { - ports[observed.Rule.DestinationPort] = struct{}{} - markers[observed.Marker] = struct{}{} - } - if _, exists := ports["80"]; !exists { - t.Fatalf("expanded port 80 is missing: %#v", adapter.snapshot.Rules) - } - if _, exists := ports["443"]; !exists { - t.Fatalf("expanded port 443 is missing: %#v", adapter.snapshot.Rules) - } - if len(markers) != 2 { - t.Fatalf("expanded rules reused one runtime marker: %#v", adapter.snapshot.Rules) - } - - retry, err := service.syncRules(ctx, "", request) - if err != nil || retry.Succeeded != 0 || retry.Skipped != 2 || retry.Failed != 0 { - t.Fatalf("expanded rule synchronization is not idempotent: result=%#v err=%v", retry, err) - } -} - -func TestFirewallRuleSyncRestoresManagedOrderFromDatabase(t *testing.T) { - ctx := context.Background() - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - targetScope := filter.Scope{ - Provider: filter.ProviderNftables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - adapter := newFakeFilterAdapter(t, targetScope, nil) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil }, - } - - records := make([]model.FirewallRule, 0, 2) - for index, port := range []string{"8080", "8081"} { - rule := filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - }, - Protocol: "tcp", DestinationPort: port, Action: filter.ActionAccept, - } - record, err := model.FirewallRuleFromDomain(rule) - if err != nil { - t.Fatal(err) - } - record.UUID = []string{"first", "second"}[index] - record.Origin = constant.FirewallRuleOriginCreated - record.Owner = constant.FirewallRuleSourceUser - sequence := int64(index+1) * model.FirewallRuleSequenceStep - record.Sequence = &sequence - if err := ruleRepo.Create(ctx, &record); err != nil { - t.Fatal(err) - } - records = append(records, record) - } - compiledFirst, err := service.compileStoredFirewallRules(ctx, records[0], filter.ProviderNftables) - if err != nil { - t.Fatal(err) - } - compiledSecond, err := service.compileStoredFirewallRules(ctx, records[1], filter.ProviderNftables) - if err != nil { - t.Fatal(err) - } - observedSecond := executorObservedRule(compiledSecond[0].Rule, compiledSecond[0].Marker, 1) - observedFirst := executorObservedRule(compiledFirst[0].Rule, compiledFirst[0].Marker, 2) - // Native inventory identifies managed rules through their marker and does - // not populate the domain rule UUID. - observedSecond.Rule.UUID = "" - observedFirst.Rule.UUID = "" - adapter.snapshot, err = filter.NewSnapshot(targetScope, []filter.ObservedRule{observedSecond, observedFirst}) - if err != nil { - t.Fatal(err) - } - - request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables} - preview, err := service.PreviewRuleSync(ctx, "", request) - if err != nil { - t.Fatal(err) - } - if preview.Ready != 2 || preview.Existing != 0 || preview.Blocked != 0 { - t.Fatalf("order drift was not included in preview: %#v", preview) - } - result, err := service.syncRules(ctx, "", request) - if err != nil { - t.Fatal(err) - } - if result.Succeeded != 2 || result.Failed != 0 || adapter.applyCount != 1 { - t.Fatalf("managed order was not synchronized: result=%#v applies=%d", result, adapter.applyCount) - } - if adapter.snapshot.Rules[0].Marker != compiledFirst[0].Marker || adapter.snapshot.Rules[1].Marker != compiledSecond[0].Marker { - t.Fatalf("runtime order does not match database order: %#v", adapter.snapshot.Rules) - } - for index, uuid := range []string{"first", "second"} { - stored, loadErr := ruleRepo.GetByUUID(ctx, uuid) - want := int64(index+1) * model.FirewallRuleSequenceStep - if loadErr != nil || stored.Sequence == nil || *stored.Sequence != want { - t.Fatalf("synchronization rewrote database sequence for %s: record=%#v err=%v", uuid, stored, loadErr) - } - } -} - -func TestFirewallRuleSyncBlocksManagedOrderAcrossExternalRule(t *testing.T) { - ctx := context.Background() - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - targetScope := filter.Scope{ - Provider: filter.ProviderNftables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - adapter := newFakeFilterAdapter(t, targetScope, nil) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil }, - } - - compiled := make([]filter.DesiredRule, 0, 2) - for index, port := range []string{"8080", "8081"} { - rule := filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - }, - Protocol: "tcp", DestinationPort: port, Action: filter.ActionAccept, - } - record, err := model.FirewallRuleFromDomain(rule) - if err != nil { - t.Fatal(err) - } - record.UUID = []string{"first", "second"}[index] - record.Origin = constant.FirewallRuleOriginCreated - record.Owner = constant.FirewallRuleSourceUser - sequence := int64(index+1) * model.FirewallRuleSequenceStep - record.Sequence = &sequence - if err := ruleRepo.Create(ctx, &record); err != nil { - t.Fatal(err) - } - rules, compileErr := service.compileStoredFirewallRules(ctx, record, filter.ProviderNftables) - if compileErr != nil { - t.Fatal(compileErr) - } - compiled = append(compiled, rules[0]) - } - external := compiled[0].Rule - external.UUID = "" - external.DestinationPort = "9090" - adapter.snapshot, _ = filter.NewSnapshot(targetScope, []filter.ObservedRule{ - executorObservedRule(compiled[1].Rule, compiled[1].Marker, 1), - executorObservedRule(external, "", 2), - executorObservedRule(compiled[0].Rule, compiled[0].Marker, 3), - }) - - request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables} - preview, err := service.PreviewRuleSync(ctx, "", request) - if err != nil { - t.Fatal(err) - } - if preview.Blocked != 2 || preview.Ready != 0 { - t.Fatalf("unsafe order change was not blocked: %#v", preview) - } - result, err := service.syncRules(ctx, "", request) - if err != nil { - t.Fatal(err) - } - if result.Failed != 2 || adapter.applyCount != 0 || adapter.snapshot.Rules[1].Marker != "" { - t.Fatalf("blocked order change mutated runtime rules: result=%#v applies=%d rules=%#v", result, adapter.applyCount, adapter.snapshot.Rules) - } -} - -func TestFirewallRuleSyncRejectsProviderSourceAndKeepsDatabasePolicy(t *testing.T) { - ctx := context.Background() - db := newFirewallRuleTestDB(t) - ruleRepo := repo.NewFirewallRuleRepo(db) - sourceRule := filter.FirewallRule{ - Scope: filter.Scope{ - Provider: filter.ProviderIptables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - }, - Protocol: "tcp", DestinationPort: "8443", Action: filter.ActionAccept, - } - sourceRecord, err := model.FirewallRuleFromDomain(sourceRule) - if err != nil { - t.Fatal(err) - } - sourceRecord.Origin = constant.FirewallRuleOriginCreated - sourceRecord.Owner = constant.FirewallRuleSourceUser - if err := ruleRepo.Create(ctx, &sourceRecord); err != nil { - t.Fatal(err) - } - targetScope := filter.Scope{ - Provider: filter.ProviderNftables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - target := newFakeFilterAdapter(t, targetScope, nil) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderNftables: newFirewallRuleRuntime(target, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil }, - } - _, err = service.PreviewRuleSync(ctx, "", dto.FirewallRuleSyncRequest{ - Subsystem: "system", SourceProvider: filter.ProviderIptables, TargetProvider: filter.ProviderNftables, - }) - if err == nil { - t.Fatal("system database synchronization accepted a source provider") - } - result, err := service.syncRules(ctx, "", dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables}) - if err != nil || result.Succeeded != 1 { - t.Fatalf("database synchronization failed: result=%#v err=%v", result, err) - } - records, err := ruleRepo.List(ctx) - if err != nil || len(records) != 1 || records[0].UUID != sourceRecord.UUID { - t.Fatalf("database policy was duplicated or removed: %#v err=%v", records, err) - } -} - -func TestFirewallRuleSyncRemovesManagedRuntimeRuleMissingFromDatabase(t *testing.T) { - ctx := context.Background() - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - scope := filter.Scope{ - Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, - Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput, - } - rule := filter.FirewallRule{ - UUID: "orphan", Scope: scope, NativeKind: filter.NativeKindRichRule, - Protocol: "tcp", DestinationPort: "9443", Action: filter.ActionAccept, - } - observed := executorObservedRule(rule, "1panel-rule:orphan", 1) - observed.Rule.UUID = "" - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{observed}) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderFirewalld: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderFirewalld, nil }, - } - request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderFirewalld} - - preview, err := service.PreviewRuleSync(ctx, "", request) - if err != nil { - t.Fatal(err) - } - if preview.Total != 0 || preview.Removed != 1 || len(preview.Items) != 1 || preview.Items[0].Status != "remove" { - t.Fatalf("unexpected orphan preview: %#v", preview) - } - result, err := service.syncRules(ctx, "", request) - if err != nil { - t.Fatal(err) - } - if result.Total != 0 || result.Removed != 1 || result.Failed != 0 || len(adapter.snapshot.Rules) != 0 { - t.Fatalf("orphan runtime rule was not removed: result=%#v rules=%#v", result, adapter.snapshot.Rules) - } - stored, err := ruleRepo.List(ctx) - if err != nil || len(stored) != 0 { - t.Fatalf("runtime cleanup changed database policies: %#v err=%v", stored, err) - } -} - -func TestFirewallRuleSyncCompileFailureDoesNotCreateOrphanRemoval(t *testing.T) { - ctx := context.Background() - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - scope := filter.Scope{ - Provider: filter.ProviderFirewalld, Family: filter.FamilyInet, - Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput, - } - rule := filter.FirewallRule{ - UUID: "incompatible", Scope: scope, NativeKind: filter.NativeKindRichRule, - Protocol: "tcp", DestinationPort: "9443", Action: filter.ActionAccept, - } - record, err := model.FirewallRuleFromDomain(rule) - if err != nil { - t.Fatal(err) - } - record.UUID = rule.UUID - record.Origin = constant.FirewallRuleOriginCreated - record.Owner = constant.FirewallRuleSourceUser - record.CompatibilityError = "manual recreation required" - if err := ruleRepo.Create(ctx, &record); err != nil { - t.Fatal(err) - } - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{ - executorObservedRule(rule, "1panel-rule:"+rule.UUID, 1), - }) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderFirewalld: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderFirewalld, nil }, - } - - preview, err := service.PreviewRuleSync(ctx, "", dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderFirewalld}) - if err != nil { - t.Fatal(err) - } - if preview.Blocked != 1 || preview.Removed != 0 || len(preview.Items) != 1 { - t.Fatalf("compile failure classified its runtime rule as an orphan: %#v", preview) - } -} - -func TestFirewallRuleSyncBlockedPlanDoesNotRemoveOrphans(t *testing.T) { - ctx := context.Background() - ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)) - scope := filter.Scope{ - Provider: filter.ProviderNftables, Family: filter.FamilyIPv4, - Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput, - } - blockedRule := filter.FirewallRule{ - Scope: scope, Protocol: "all", Action: filter.ActionDrop, - } - record, err := model.FirewallRuleFromDomain(blockedRule) - if err != nil { - t.Fatal(err) - } - record.Origin = constant.FirewallRuleOriginCreated - record.Owner = constant.FirewallRuleSourceUser - if err := ruleRepo.Create(ctx, &record); err != nil { - t.Fatal(err) - } - orphanRule := filter.FirewallRule{ - UUID: "orphan", Scope: scope, Protocol: "tcp", DestinationPort: "9443", Action: filter.ActionAccept, - } - adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{ - executorObservedRule(orphanRule, "1panel-rule:orphan", 1), - }) - service := &FirewallService{ - rules: ruleRepo, - adapters: firewallRuleRuntimeRegistry{ - filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil), - }, - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil }, - } - - result, err := service.syncRules(ctx, "203.0.113.10", dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables}) - if err != nil { - t.Fatal(err) - } - if result.Failed != 1 || result.Removed != 0 || adapter.applyCount != 0 || len(adapter.snapshot.Rules) != 1 { - t.Fatalf("blocked synchronization mutated target rules: result=%#v rules=%#v applies=%d", result, adapter.snapshot.Rules, adapter.applyCount) - } -} - -func waitFirewallSyncTask(t *testing.T, taskID string) { - t.Helper() - taskRepo := repo.NewITaskRepo() - deadline := time.Now().Add(5 * time.Second) - for { - completed := false - record, loadErr := taskRepo.GetFirst(taskRepo.WithByID(taskID)) - if loadErr == nil && record.Status != constant.StatusExecuting { - if record.Status != constant.StatusSuccess { - t.Fatalf("firewall synchronization task failed: %#v", record) - } - completed = true - } - firewallRuleSyncTaskMu.Lock() - activeTaskID := firewallRuleSyncTaskID - firewallRuleSyncTaskMu.Unlock() - if completed && activeTaskID == "" { - return - } - if time.Now().After(deadline) { - t.Fatal("timed out waiting for firewall synchronization task") - } - time.Sleep(20 * time.Millisecond) - } -} - -func setupFirewallTaskTestDB(t *testing.T) { - t.Helper() - agenti18n.Init() - previousDB := global.TaskDB - previousDir := global.Dir.TaskDir - taskDir := t.TempDir() - db, err := gorm.Open(sqlite.Open(filepath.Join(taskDir, "task.db")), &gorm.Config{}) - if err != nil { - t.Fatal(err) - } - if err := db.AutoMigrate(&model.Task{}); err != nil { - t.Fatal(err) - } - global.TaskDB = db - global.Dir.TaskDir = taskDir - t.Cleanup(func() { - global.TaskDB = previousDB - global.Dir.TaskDir = previousDir - }) -} - -func TestSortFirewallPoliciesUsesProviderPlacement(t *testing.T) { - sequenceOne, sequenceTwo := model.FirewallRuleSequenceStep, 2*model.FirewallRuleSequenceStep - priorityLow, priorityHigh := -100, 100 - policies := []model.FirewallRule{ - {UUID: "high", Priority: &priorityHigh, Sequence: &sequenceOne}, - {UUID: "none"}, - {UUID: "low", Priority: &priorityLow, Sequence: &sequenceTwo}, - } - - positional := append([]model.FirewallRule(nil), policies...) - sortFirewallPolicies(positional, filter.ProviderUFW) - if positional[0].UUID != "high" || positional[1].UUID != "low" || positional[2].UUID != "none" { - t.Fatalf("positional policies were not sorted by sequence: %#v", positional) - } - - weighted := append([]model.FirewallRule(nil), policies...) - sortFirewallPolicies(weighted, filter.ProviderFirewalld) - if weighted[0].UUID != "low" || weighted[1].UUID != "high" || weighted[2].UUID != "none" { - t.Fatalf("firewalld policies were not sorted by priority: %#v", weighted) - } -} - -func TestFirewallRuleSyncRejectsCurrentProviderAsSource(t *testing.T) { - service := &FirewallService{ - rules: repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)), - selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil }, - } - _, err := service.PreviewRuleSync(context.Background(), "", dto.FirewallRuleSyncRequest{ - SourceProvider: filter.ProviderIptables, - TargetProvider: filter.ProviderIptables, - }) - if err == nil { - t.Fatal("expected identical source and target providers to be rejected") - } -} - -func TestDatabaseRuleSyncTargetUsesExplicitTargetOnly(t *testing.T) { - target, err := databaseRuleSyncTarget(dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables}, "Docker") - if err != nil || target != filter.ProviderNftables { - t.Fatalf("target = %q, err = %v", target, err) - } - if _, err := databaseRuleSyncTarget(dto.FirewallRuleSyncRequest{ - SourceProvider: filter.ProviderIptables, - TargetProvider: filter.ProviderNftables, - }, "Docker"); err == nil { - t.Fatal("expected database-backed synchronization to reject a source provider") - } - if _, err := databaseRuleSyncTarget(dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderUFW}, "Docker"); err == nil { - t.Fatal("expected database-backed synchronization to reject a non-netfilter target") - } -} diff --git a/agent/app/service/forward_service.go b/agent/app/service/forward.go similarity index 97% rename from agent/app/service/forward_service.go rename to agent/app/service/forward.go index 48e44b7222ee..75e8edf15e40 100644 --- a/agent/app/service/forward_service.go +++ b/agent/app/service/forward.go @@ -319,7 +319,7 @@ func (s *ForwardingService) loadRuleSyncCandidates( Family: record.Family, Protocol: record.Protocol, Port: record.Port, TargetIP: record.TargetIP, TargetPort: record.TargetPort, Interface: record.Interface, } - normalized, normalizeErr := forwardingproviders.NormalizeRule(rule) + normalized, normalizeErr := forwarding.NormalizeRule(rule) candidates = append(candidates, forwardingRuleSyncCandidate{rule: normalized, err: normalizeErr}) } targetStatus, err := target.Status() @@ -358,7 +358,7 @@ func verifyForwardingRuleSync(target *forwarding.Manager, desired []forwarding.R func normalizeForwardingRuntimeRules(rules []forwarding.Rule) ([]forwarding.Rule, error) { normalized := make([]forwarding.Rule, 0, len(rules)) for _, rule := range rules { - item, err := forwardingproviders.NormalizeRule(rule) + item, err := forwarding.NormalizeRule(rule) if err != nil { return nil, fmt.Errorf("normalize target forwarding rule %s: %w", rule.Identity(), err) } @@ -473,7 +473,7 @@ func mergeForwardingInventory( items := make([]forwardingInventoryItem, 0, len(stored)+len(runtime)) byIdentity := make(map[string]int, len(stored)+len(runtime)) for _, record := range stored { - rule, err := forwardingproviders.NormalizeRule(forwarding.Rule{ + rule, err := forwarding.NormalizeRule(forwarding.Rule{ Family: record.Family, Protocol: record.Protocol, Port: record.Port, TargetIP: record.TargetIP, TargetPort: record.TargetPort, Interface: record.Interface, }) @@ -485,7 +485,7 @@ func mergeForwardingInventory( items = append(items, forwardingInventoryItem{ID: record.ID, Rule: rule, IsDesired: true}) } for _, observed := range runtime { - rule, err := forwardingproviders.NormalizeRule(observed) + rule, err := forwarding.NormalizeRule(observed) if err != nil { return nil, fmt.Errorf("normalize runtime forwarding rule: %w", err) } @@ -518,7 +518,7 @@ func lastForwardingSyncError() string { func applyForwardingOperations(current []forwarding.Rule, requested []dto.ForwardRuleOperation) ([]forwarding.Rule, error) { desired := make([]forwarding.Rule, 0, len(current)+len(requested)) for _, rule := range current { - normalized, err := forwardingproviders.NormalizeRule(rule) + normalized, err := forwarding.NormalizeRule(rule) if err != nil { return nil, fmt.Errorf("normalize persisted forwarding rule: %w", err) } @@ -526,7 +526,7 @@ func applyForwardingOperations(current []forwarding.Rule, requested []dto.Forwar } for _, operation := range requested { for _, protocol := range strings.Split(operation.Protocol, "/") { - rule, err := forwardingproviders.NormalizeRule(forwarding.Rule{ + rule, err := forwarding.NormalizeRule(forwarding.Rule{ Family: operation.Family, Protocol: protocol, Port: operation.Port, TargetIP: operation.TargetIP, TargetPort: operation.TargetPort, Interface: operation.Interface, }) @@ -605,7 +605,10 @@ func newForwardingManagerFor(backend string) (*forwarding.Manager, error) { return forwarding.NewManager(candidate.adapter, candidate.runtime), nil } } - return nil, fmt.Errorf("%w: selected forwarding backend %s is not installed", errForwardingBackendUnavailable, backend) + return nil, fmt.Errorf( + "%w: selected forwarding backend %s %w", + errForwardingBackendUnavailable, backend, lifecycle.ErrNotInstalled, + ) } return selectForwardingManager(candidates) } diff --git a/agent/app/service/forwarding_contract_test.go b/agent/app/service/forwarding_contract_test.go deleted file mode 100644 index 5e397b499e74..000000000000 --- a/agent/app/service/forwarding_contract_test.go +++ /dev/null @@ -1,686 +0,0 @@ -package service - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "reflect" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - forwardClient "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" - "github.com/go-playground/validator/v10" -) - -type fakeForwardingAdapter struct { - name string - rules []forwardClient.Rule - listErr error - operateErr error - enableErr error - init bool - initErr error - familyInit map[string]bool - reconciled []forwardClient.Rule - reconciles int -} - -func (f *fakeForwardingAdapter) Name() string { return f.name } - -func (f *fakeForwardingAdapter) List() ([]forwardClient.Rule, error) { - return append([]forwardClient.Rule(nil), f.rules...), f.listErr -} - -func (f *fakeForwardingAdapter) Reconcile(rules []forwardClient.Rule) error { - f.reconciles++ - f.reconciled = append([]forwardClient.Rule(nil), rules...) - if f.operateErr == nil { - f.rules = append([]forwardClient.Rule(nil), rules...) - } - return f.operateErr -} - -func (f *fakeForwardingAdapter) Enable() error { return f.enableErr } -func (f *fakeForwardingAdapter) Cleanup() error { return nil } -func (f *fakeForwardingAdapter) FamilyStatus(family string) (bool, bool, error) { - if f.familyInit != nil { - initialized := f.familyInit[family] - return initialized, initialized, f.initErr - } - return f.init, f.init, f.initErr -} -func (f *fakeForwardingAdapter) InitStatus() (bool, bool, error) { - return f.init, f.init, f.initErr -} -func (f *fakeForwardingAdapter) Replay() error { return nil } - -type fakeForwardingRuleRepo struct { - rules []model.ForwardingRule - listErr error -} - -func (r *fakeForwardingRuleRepo) List(context.Context) ([]model.ForwardingRule, error) { - return append([]model.ForwardingRule(nil), r.rules...), r.listErr -} - -func (r *fakeForwardingRuleRepo) ReplaceAll(_ context.Context, rules []model.ForwardingRule) error { - r.rules = append([]model.ForwardingRule(nil), rules...) - return nil -} - -func forwardingServiceWithAdapter(adapter forwardClient.Adapter) *ForwardingService { - rules, listErr := adapter.List() - return &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return forwardClient.NewManager(adapter, nil), nil - }, - rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels(rules), listErr: listErr}, - enabled: func() (bool, error) { return true, nil }, - persistBackend: func(string) error { return nil }, - } -} - -func TestForwardingAndFilterInterfacesAreSeparated(t *testing.T) { - filterType := reflect.TypeOf((*lifecycle.Client)(nil)).Elem() - for _, method := range []string{"ListForward", "PortForward", "EnableForward"} { - if _, ok := filterType.MethodByName(method); ok { - t.Fatalf("filter interface still exposes %s", method) - } - } - firewallServiceType := reflect.TypeOf((*IFirewallService)(nil)).Elem() - if _, ok := firewallServiceType.MethodByName("OperateForwardRule"); ok { - t.Fatal("firewall service still owns forwarding writes") - } - for _, method := range []string{"PreviewRuleSync", "SyncRules"} { - if _, ok := firewallServiceType.MethodByName(method); !ok { - t.Fatalf("firewall service missing centralized %s", method) - } - } - forwardingServiceType := reflect.TypeOf((*IForwardingService)(nil)).Elem() - for _, method := range []string{"LoadBaseInfo", "SearchRules", "OperateRules", "Enable", "Restore"} { - if _, ok := forwardingServiceType.MethodByName(method); !ok { - t.Fatalf("forwarding service missing %s", method) - } - } - for _, serviceType := range []reflect.Type{ - forwardingServiceType, - reflect.TypeOf((*IDockerPortGuardService)(nil)).Elem(), - } { - for _, method := range []string{"PreviewRuleSync", "SyncRules"} { - if _, ok := serviceType.MethodByName(method); ok { - t.Fatalf("subsystem service still exposes centralized %s", method) - } - } - } -} - -func TestForwardingInitNoLongerUsesFilterRequest(t *testing.T) { - request := dto.FilterChainOperation{Name: "1PANEL_FORWARD", Operate: "init-forward"} - if err := validator.New().Struct(request); err == nil { - t.Fatal("forwarding initialization must use the dedicated forwarding endpoint") - } -} - -func TestForwardingRestoreHonorsPersistedState(t *testing.T) { - adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80", - }}} - service := forwardingServiceWithAdapter(adapter) - service.enabled = func() (bool, error) { return false, nil } - if err := service.Restore(context.Background()); err != nil { - t.Fatal(err) - } - if adapter.reconciled != nil { - t.Fatalf("disabled forwarding was restored: %#v", adapter.reconciled) - } - - service.enabled = func() (bool, error) { return true, nil } - if err := service.Restore(context.Background()); err != nil { - t.Fatal(err) - } - if len(adapter.reconciled) != 1 || adapter.reconciled[0].Port != "8080" { - t.Fatalf("unexpected restored rules: %#v", adapter.reconciled) - } -} - -func TestForwardingBackendSelection(t *testing.T) { - nft := &fakeForwardingAdapter{name: "nftables"} - iptables := &fakeForwardingAdapter{name: "iptables"} - - manager, err := selectForwardingManager([]forwardingCandidate{{adapter: nft}, {adapter: iptables}}) - if err != nil { - t.Fatal(err) - } - if manager.Name() != "nftables" { - t.Fatalf("uninitialized backends selected %q, want nftables", manager.Name()) - } - - iptables.init = true - manager, err = selectForwardingManager([]forwardingCandidate{{adapter: nft}, {adapter: iptables}}) - if err != nil { - t.Fatal(err) - } - if manager.Name() != "iptables" { - t.Fatalf("initialized backend selected %q, want iptables", manager.Name()) - } - - nft.init = true - if _, err := selectForwardingManager([]forwardingCandidate{{adapter: nft}, {adapter: iptables}}); !errors.Is(err, errForwardingBackendConflict) { - t.Fatalf("both initialized backends returned %v, want conflict", err) - } -} - -func TestForwardingBackendSelectionRejectsSplitFamilies(t *testing.T) { - nft := &fakeForwardingAdapter{name: "nftables", familyInit: map[string]bool{forwardClient.FamilyIPv4: true}} - iptables := &fakeForwardingAdapter{name: "iptables", familyInit: map[string]bool{forwardClient.FamilyIPv6: true}} - if _, err := selectForwardingManager([]forwardingCandidate{{adapter: nft}, {adapter: iptables}}); !errors.Is(err, errForwardingBackendConflict) { - t.Fatalf("split-family backends returned %v, want conflict", err) - } -} - -func TestForwardingBackendSelectionReturnsStatusError(t *testing.T) { - wantErr := errors.New("status failed") - _, err := selectForwardingManager([]forwardingCandidate{{adapter: &fakeForwardingAdapter{name: "nftables", initErr: wantErr}}}) - if !errors.Is(err, wantErr) { - t.Fatalf("got %v want %v", err, wantErr) - } -} - -func TestForwardingRuleSyncReplaysPersistedRulesIntoTarget(t *testing.T) { - sourceRule := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - target := &fakeForwardingAdapter{name: "nftables"} - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return forwardClient.NewManager(target, nil), nil - }, - rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{sourceRule})}, - enabled: func() (bool, error) { return true, nil }, - persistBackend: func(string) error { return nil }, - markEnabled: func() error { return nil }, - } - request := dto.FirewallRuleSyncRequest{Subsystem: "forwarding", TargetProvider: "nftables"} - preview, err := service.previewRuleSync(context.Background(), request) - if err != nil { - t.Fatal(err) - } - if preview.Total != 1 || preview.Ready != 1 || preview.TargetProvider != "nftables" || preview.Items[0].ForwardRule == nil { - t.Fatalf("unexpected preview: %#v", preview) - } - result, err := service.syncRules(context.Background(), request) - if err != nil { - t.Fatal(err) - } - if result.Succeeded != 1 || result.Failed != 0 || len(target.reconciled) != 1 || target.reconciled[0].Identity() != sourceRule.Identity() { - t.Fatalf("unexpected sync result=%#v target=%#v", result, target.reconciled) - } - target.init = true - retry, err := service.previewRuleSync(context.Background(), request) - if err != nil { - t.Fatal(err) - } - if retry.Ready != 0 || retry.Existing != 1 { - t.Fatalf("synchronized forwarding rule was not recognized: %#v", retry) - } -} - -func TestForwardingRuleSyncReportsActivationFailure(t *testing.T) { - wantErr := errors.New("enable forwarding failed") - sourceRule := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - target := &fakeForwardingAdapter{name: "nftables", enableErr: wantErr} - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return forwardClient.NewManager(target, nil), nil - }, - rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{sourceRule})}, - persistBackend: func(string) error { return nil }, - markEnabled: func() error { return nil }, - } - result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "forwarding", TargetProvider: "nftables", - }) - if err != nil { - t.Fatal(err) - } - if result.Succeeded != 0 || result.Failed != 1 || len(result.Errors) != 1 || - result.Errors[0].Error != wantErr.Error() { - t.Fatalf("activation failure was not reported: %#v", result) - } - if target.reconciles != 0 { - t.Fatalf("failed activation reconciled target %d times", target.reconciles) - } -} - -func TestForwardingRuleSyncTreatsWildcardInterfaceAsDatabaseDefault(t *testing.T) { - databaseRule := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - runtimeRule := databaseRule - runtimeRule.Interface = "*" - target := &fakeForwardingAdapter{name: "iptables", init: true, rules: []forwardClient.Rule{runtimeRule}} - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return forwardClient.NewManager(target, nil), nil - }, - rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{databaseRule})}, - enabled: func() (bool, error) { return true, nil }, - persistBackend: func(string) error { return nil }, - markEnabled: func() error { return nil }, - } - preview, err := service.previewRuleSync(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "forwarding", TargetProvider: "iptables", - }) - if err != nil { - t.Fatal(err) - } - if preview.Total != 1 || preview.Existing != 1 || preview.Ready != 0 || preview.Removed != 0 || len(preview.Items) != 1 { - t.Fatalf("wildcard interface produced duplicate sync actions: %#v", preview) - } -} - -func TestForwardingRuleSyncReconcilesTargetToDatabaseState(t *testing.T) { - existing := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - missing := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8443", TargetIP: "10.0.0.3", TargetPort: "443", - } - extra := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "udp", Port: "5353", TargetIP: "10.0.0.4", TargetPort: "53", - } - target := &fakeForwardingAdapter{name: "nftables", init: true, rules: []forwardClient.Rule{existing, extra}} - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return forwardClient.NewManager(target, nil), nil - }, - rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{existing, missing})}, - enabled: func() (bool, error) { return true, nil }, - persistBackend: func(string) error { return nil }, - markEnabled: func() error { return nil }, - } - preview, err := service.previewRuleSync(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "forwarding", TargetProvider: "nftables", - }) - if err != nil { - t.Fatal(err) - } - if preview.Total != 2 || preview.Removed != 1 || len(preview.Items) != 3 || preview.Items[2].Status != "remove" { - t.Fatalf("extra target rule was not included in preview: %#v", preview) - } - result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "forwarding", TargetProvider: "nftables", - }) - if err != nil { - t.Fatal(err) - } - if result.Succeeded != 1 || result.Skipped != 1 || result.Removed != 1 || result.Failed != 0 || target.reconciles != 1 { - t.Fatalf("unexpected sync result=%#v reconciles=%d", result, target.reconciles) - } - if len(target.rules) != 2 || forwardingRuleIndex(target.rules, existing) < 0 || - forwardingRuleIndex(target.rules, missing) < 0 || forwardingRuleIndex(target.rules, extra) >= 0 { - t.Fatalf("target did not converge to database state: %#v", target.rules) - } -} - -func TestForwardingRuleSyncRemovesExtraRulesWhenDatabaseRulesAlreadyExist(t *testing.T) { - existing := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - extra := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "udp", Port: "5353", TargetIP: "10.0.0.4", TargetPort: "53", - } - target := &fakeForwardingAdapter{name: "nftables", init: true, rules: []forwardClient.Rule{existing, extra}} - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return forwardClient.NewManager(target, nil), nil - }, - rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{existing})}, - enabled: func() (bool, error) { return true, nil }, - persistBackend: func(string) error { return nil }, - markEnabled: func() error { return nil }, - } - result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "forwarding", TargetProvider: "nftables", - }) - if err != nil { - t.Fatal(err) - } - if result.Succeeded != 0 || result.Skipped != 1 || result.Failed != 0 || target.reconciles != 1 || - len(target.rules) != 1 || forwardingRuleIndex(target.rules, existing) < 0 { - t.Fatalf("extra target rule was not removed: result=%#v target=%#v", result, target) - } -} - -func TestForwardingRuleSyncClearsInitializedTargetWhenDatabaseIsEmpty(t *testing.T) { - extra := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - target := &fakeForwardingAdapter{name: "nftables", init: true, rules: []forwardClient.Rule{extra}} - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return forwardClient.NewManager(target, nil), nil - }, - rules: &fakeForwardingRuleRepo{}, - } - result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{ - Subsystem: "forwarding", TargetProvider: "nftables", - }) - if err != nil { - t.Fatal(err) - } - if result.Total != 0 || result.Removed != 1 || target.reconciles != 1 || len(target.rules) != 0 { - t.Fatalf("empty database did not clear initialized target: result=%#v target=%#v", result, target) - } -} - -func TestForwardingRuleSyncRejectsUnselectedTarget(t *testing.T) { - current := &fakeForwardingAdapter{name: "iptables"} - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return forwardClient.NewManager(current, nil), nil - }, - rules: &fakeForwardingRuleRepo{}, - } - request := dto.FirewallRuleSyncRequest{Subsystem: "forwarding", TargetProvider: filter.ProviderNftables} - if _, err := service.previewRuleSync(context.Background(), request); !errors.Is(err, filter.ErrProviderUnavailable) { - t.Fatalf("preview error = %v, want provider unavailable", err) - } - if _, err := service.syncRules(context.Background(), request); !errors.Is(err, filter.ErrProviderUnavailable) { - t.Fatalf("sync error = %v, want provider unavailable", err) - } - if current.reconciles != 0 { - t.Fatalf("unselected target request modified the current backend: %#v", current) - } -} - -func TestForwardingDisplayName(t *testing.T) { - for backend, want := range map[string]string{ - "iptables": "iptables-forward", - "nftables": "nftables-forward", - "unknown": "unknown", - } { - if got := forwardingDisplayName(backend); got != want { - t.Fatalf("forwardingDisplayName(%q) = %q, want %q", backend, got, want) - } - } -} - -func TestForwardingBaseInfoIncludesFamilyStatus(t *testing.T) { - adapter := &fakeForwardingAdapter{ - name: "nftables", - familyInit: map[string]bool{forwardClient.FamilyIPv4: true}, - } - base, err := forwardingServiceWithAdapter(adapter).LoadBaseInfo() - if err != nil { - t.Fatal(err) - } - if !base.IPv4.Available || !base.IPv4.Initialized || !base.IPv4.Bound { - t.Fatalf("unexpected IPv4 status: %#v", base.IPv4) - } - if !base.IPv6.Available || base.IPv6.Initialized || base.IPv6.Bound { - t.Fatalf("unexpected IPv6 status: %#v", base.IPv6) - } -} - -func TestForwardingBaseInfoReportsMissingBackend(t *testing.T) { - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return nil, errForwardingBackendUnavailable - }, - installedProviders: func() []string { return nil }, - } - base, err := service.LoadBaseInfo() - if err != nil { - t.Fatal(err) - } - if base.IsExist || base.Name != "-" || base.Backend != "-" { - t.Fatalf("unexpected missing forwarding backend status: %#v", base) - } -} - -func TestForwardingBaseInfoReportsUnavailableSelection(t *testing.T) { - wantErr := fmt.Errorf("%w: selected forwarding backend nftables is not installed", errForwardingBackendUnavailable) - service := &ForwardingService{ - managerFactory: func() (*forwardClient.Manager, error) { - return nil, wantErr - }, - installedProviders: func() []string { return []string{constant.FirewallProviderIptables} }, - } - base, err := service.LoadBaseInfo() - if err != nil { - t.Fatal(err) - } - if !base.IsExist || base.Message != wantErr.Error() || base.Name != "-" || base.Backend != "-" { - t.Fatalf("unexpected unavailable forwarding selection status: %#v", base) - } -} - -func TestForwardingSearchPreservesAPIShapeAndPagination(t *testing.T) { - adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{ - {Num: "1", Family: forwardClient.FamilyIPv6, Protocol: "tcp", Port: "8080", TargetIP: "2001:db8::2", TargetPort: "80", Interface: "eth0"}, - {Num: "2", Protocol: "udp", Port: "5353", TargetIP: "127.0.0.1", TargetPort: "53"}, - }} - service := forwardingServiceWithAdapter(adapter) - service.rules.(*fakeForwardingRuleRepo).rules[0].ID = 42 - total, value, err := service.SearchRules(dto.ForwardRuleSearch{PageInfo: dto.PageInfo{Page: 1, PageSize: 10}, Info: "2001:db8"}) - if err != nil { - t.Fatal(err) - } - if total != 1 { - t.Fatalf("got total %d want 1", total) - } - items, ok := value.([]dto.ForwardRule) - if !ok || len(items) != 1 || items[0].ID != 42 || items[0].Port != "8080" || - items[0].Family != forwardClient.FamilyIPv6 || !items[0].IsDesired || !items[0].IsRuntime || - items[0].SyncStatus != forwardingSyncConverged { - t.Fatalf("unexpected items: %#v", value) - } - data, err := json.Marshal(items[0]) - if err != nil { - t.Fatal(err) - } - var fields map[string]interface{} - if err := json.Unmarshal(data, &fields); err != nil { - t.Fatal(err) - } - wantFields := []string{"id", "chain", "family", "address", "port", "protocol", "strategy", "num", "targetIP", "targetPort", "interface", "usedStatus", "description", "isDesired", "isRuntime", "syncStatus"} - for _, field := range wantFields { - if _, ok := fields[field]; !ok { - t.Fatalf("forward response dropped compatibility field %q: %s", field, data) - } - } -} - -func TestForwardingSearchTrimsAndMatchesAllDisplayedFields(t *testing.T) { - adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{ - {Family: forwardClient.FamilyIPv4, Protocol: "udp", Port: "5353", TargetIP: "127.0.0.1", TargetPort: "53", Interface: "eth0"}, - }} - service := forwardingServiceWithAdapter(adapter) - for _, keyword := range []string{" UDP ", "ipv4", "ETH0", "converged"} { - total, value, err := service.SearchRules(dto.ForwardRuleSearch{ - PageInfo: dto.PageInfo{Page: 1, PageSize: 10}, Info: keyword, - }) - if err != nil { - t.Fatalf("search %q: %v", keyword, err) - } - if items := value.([]dto.ForwardRule); total != 1 || len(items) != 1 { - t.Fatalf("search %q returned total=%d items=%#v", keyword, total, items) - } - } -} - -func TestForwardingSearchMergesDesiredAndRuntimeState(t *testing.T) { - desired := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - runtimeOnly := forwardClient.Rule{ - Family: forwardClient.FamilyIPv6, Protocol: "udp", Port: "5353", TargetIP: "2001:db8::2", TargetPort: "53", - } - adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{runtimeOnly}} - service := forwardingServiceWithAdapter(adapter) - service.rules = &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{desired})} - - total, value, err := service.SearchRules(dto.ForwardRuleSearch{PageInfo: dto.PageInfo{Page: 1, PageSize: 10}}) - if err != nil { - t.Fatal(err) - } - items, ok := value.([]dto.ForwardRule) - if !ok || total != 2 || len(items) != 2 { - t.Fatalf("unexpected forwarding inventory: %#v", value) - } - if !items[0].IsDesired || items[0].IsRuntime || items[0].SyncStatus != forwardingSyncMissing { - t.Fatalf("unexpected missing desired rule: %#v", items[0]) - } - if items[1].IsDesired || !items[1].IsRuntime || items[1].SyncStatus != forwardingSyncRuntimeOnly { - t.Fatalf("unexpected runtime-only rule: %#v", items[1]) - } -} - -func TestForwardingDuplicateIdentityIncludesAddressFamily(t *testing.T) { - adapter := &fakeForwardingAdapter{name: "nftables", rules: []forwardClient.Rule{{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - }}} - service := forwardingServiceWithAdapter(adapter) - err := service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{{ - Operation: "add", Family: forwardClient.FamilyIPv6, Protocol: "tcp", Port: "8080", TargetIP: "2001:db8::2", TargetPort: "80", - }}}) - if err != nil { - t.Fatal(err) - } - if len(adapter.reconciled) != 2 || adapter.reconciled[1].Family != forwardClient.FamilyIPv6 { - t.Fatalf("unexpected reconciled rules: %#v", adapter.reconciled) - } -} - -func TestForwardingOperatePreservesDuplicateAndOrderingContracts(t *testing.T) { - existing := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{ - {Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80"}, - }} - service := forwardingServiceWithAdapter(existing) - err := service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{{ - Operation: "add", Protocol: "tcp", Port: "8080", TargetPort: "80", - }}}) - if err == nil { - t.Fatal("duplicate forwarding rule must be rejected") - } - adapter := &fakeForwardingAdapter{name: "iptables"} - service = forwardingServiceWithAdapter(adapter) - err = service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{ - {Operation: "add", Protocol: "tcp/udp", Port: "9000", TargetIP: "10.0.0.2", TargetPort: "90"}, - {Operation: "remove", Num: "1", Protocol: "tcp", Port: "8001", TargetIP: "10.0.0.2", TargetPort: "81"}, - {Operation: "remove", Num: "3", Protocol: "tcp", Port: "8003", TargetIP: "10.0.0.2", TargetPort: "83"}, - }}) - if err != nil { - t.Fatal(err) - } - want := []forwardClient.Rule{ - {Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "9000", TargetIP: "10.0.0.2", TargetPort: "90"}, - {Family: forwardClient.FamilyIPv4, Protocol: "udp", Port: "9000", TargetIP: "10.0.0.2", TargetPort: "90"}, - } - if !reflect.DeepEqual(adapter.reconciled, want) { - t.Fatalf("desired forwarding rules changed\ngot %#v\nwant %#v", adapter.reconciled, want) - } -} - -func TestForwardingSearchReturnsAdapterError(t *testing.T) { - wantErr := errors.New("list failed") - service := forwardingServiceWithAdapter(&fakeForwardingAdapter{name: "nftables", listErr: wantErr}) - _, _, err := service.SearchRules(dto.ForwardRuleSearch{PageInfo: dto.PageInfo{Page: 1, PageSize: 20}}) - if !errors.Is(err, wantErr) { - t.Fatalf("got %v want %v", err, wantErr) - } -} - -func TestForwardingOperateReturnsAdapterListError(t *testing.T) { - wantErr := errors.New("list failed") - service := forwardingServiceWithAdapter(&fakeForwardingAdapter{name: "iptables", listErr: wantErr}) - err := service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{{ - Operation: "remove", Protocol: "tcp", Port: "8080", TargetPort: "80", - }}}) - if !errors.Is(err, wantErr) { - t.Fatalf("got %v want %v", err, wantErr) - } -} - -func TestForwardingForceDeleteOnlySuppressesRemoveErrors(t *testing.T) { - wantErr := errors.New("operate failed") - removeAdapter := &fakeForwardingAdapter{name: "iptables", operateErr: wantErr} - removeService := forwardingServiceWithAdapter(removeAdapter) - err := removeService.OperateRules(dto.ForwardRuleOperate{ForceDelete: true, Rules: []dto.ForwardRuleOperation{{ - Operation: "remove", Protocol: "tcp", Port: "8080", TargetPort: "80", - }}}) - if err != nil { - t.Fatalf("forced remove returned %v", err) - } - - addAdapter := &fakeForwardingAdapter{name: "iptables", operateErr: wantErr} - addService := forwardingServiceWithAdapter(addAdapter) - err = addService.OperateRules(dto.ForwardRuleOperate{ForceDelete: true, Rules: []dto.ForwardRuleOperation{{ - Operation: "add", Protocol: "tcp", Port: "8080", TargetPort: "80", - }}}) - if !errors.Is(err, wantErr) { - t.Fatalf("got %v want %v", err, wantErr) - } -} - -func TestForwardingForceDeleteExposesRuntimeOnlyRuleAndSyncError(t *testing.T) { - recordForwardingSyncError(nil) - t.Cleanup(func() { recordForwardingSyncError(nil) }) - wantErr := errors.New("runtime reconcile failed") - rule := forwardClient.Rule{ - Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{rule}} - service := forwardingServiceWithAdapter(adapter) - adapter.operateErr = wantErr - - err := service.OperateRules(dto.ForwardRuleOperate{ForceDelete: true, Rules: []dto.ForwardRuleOperation{{ - Operation: "remove", Family: rule.Family, Protocol: rule.Protocol, Port: rule.Port, - TargetIP: rule.TargetIP, TargetPort: rule.TargetPort, - }}}) - if err != nil { - t.Fatalf("forced remove returned %v", err) - } - base, err := service.LoadBaseInfo() - if err != nil { - t.Fatal(err) - } - if base.SyncError != wantErr.Error() { - t.Fatalf("sync error = %q, want %q", base.SyncError, wantErr) - } - _, value, err := service.SearchRules(dto.ForwardRuleSearch{PageInfo: dto.PageInfo{Page: 1, PageSize: 10}}) - if err != nil { - t.Fatal(err) - } - items := value.([]dto.ForwardRule) - if len(items) != 1 || items[0].IsDesired || !items[0].IsRuntime || items[0].SyncStatus != forwardingSyncRuntimeOnly { - t.Fatalf("unexpected runtime-only inventory after forced delete: %#v", items) - } -} - -func TestForwardingOperateKeepsDesiredStateWhenRuntimeReconcileFails(t *testing.T) { - wantErr := errors.New("runtime reconcile failed") - adapter := &fakeForwardingAdapter{name: "iptables", operateErr: wantErr} - rules := &fakeForwardingRuleRepo{} - service := forwardingServiceWithAdapter(adapter) - service.rules = rules - - err := service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{{ - Operation: "add", Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - }}}) - if !errors.Is(err, wantErr) { - t.Fatalf("got %v want %v", err, wantErr) - } - if len(rules.rules) != 1 || rules.rules[0].Port != "8080" { - t.Fatalf("database desired state was lost after runtime failure: %#v", rules.rules) - } - -} diff --git a/agent/init/migration/migrate_test.go b/agent/init/migration/migrate_test.go deleted file mode 100644 index 4fa1dff370a4..000000000000 --- a/agent/init/migration/migrate_test.go +++ /dev/null @@ -1,37 +0,0 @@ -package migration - -import "testing" - -func TestAgentDBMigrationsRegisterFirewallUpgradeSteps(t *testing.T) { - migrations := agentDBMigrations() - indexes := make(map[string]int, len(migrations)) - for index, migration := range migrations { - if migration == nil || migration.ID == "" { - t.Fatalf("invalid migration at index %d: %#v", index, migration) - } - if previous, exists := indexes[migration.ID]; exists { - t.Fatalf("duplicate migration ID %q at indexes %d and %d", migration.ID, previous, index) - } - indexes[migration.ID] = index - } - - tables, hasTables := indexes["20260819-add-firewall-v2-tables"] - status, hasStatus := indexes["20260818-init-docker-port-guard-status"] - selections, hasSelections := indexes["20260826-normalize-firewall-backend-selections"] - policy, hasPolicy := indexes["20260826-simplify-firewall-rule-policy"] - if !hasTables || !hasStatus || !hasSelections || !hasPolicy { - t.Fatalf("firewall upgrade migrations are not registered: %#v", indexes) - } - if _, exists := indexes["20260826-remove-firewall-rule-provider"]; exists { - t.Fatal("firewall provider removal is still registered as a separate migration") - } - if tables > status { - t.Fatalf("firewall table migration index %d runs after status migration index %d", tables, status) - } - if status > selections { - t.Fatalf("firewall status migration index %d runs after selection migration index %d", status, selections) - } - if selections > policy { - t.Fatalf("firewall selection migration index %d runs after policy migration index %d", selections, policy) - } -} diff --git a/agent/init/migration/migrations/firewall_test.go b/agent/init/migration/migrations/firewall_test.go deleted file mode 100644 index 0ee14dad7a0a..000000000000 --- a/agent/init/migration/migrations/firewall_test.go +++ /dev/null @@ -1,327 +0,0 @@ -package migrations - -import ( - "path/filepath" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/glebarez/sqlite" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -func TestAddFirewallRuleTableCreatesUpgradeSchema(t *testing.T) { - db := newFirewallMigrationTestDB(t) - for i := 0; i < 2; i++ { - if err := AddFirewallRuleTable.Migrate(db); err != nil { - t.Fatalf("migrate firewall v2 tables on pass %d: %v", i+1, err) - } - } - - for _, table := range []interface{}{&model.FirewallRule{}, &model.DockerPortGuardPolicy{}, &model.ForwardingRule{}} { - if !db.Migrator().HasTable(table) { - t.Fatalf("migration did not create table for %T", table) - } - } - for _, column := range []string{ - "provider", "scope_key", "location", "native_kind", "order_index", "order_bucket", "rule_key", "match_key", - } { - if db.Migrator().HasColumn("firewall_rules", column) { - t.Fatalf("new firewall desired-rule schema persisted derived column %s", column) - } - } - for _, column := range []string{"priority", "sequence"} { - if !db.Migrator().HasColumn("firewall_rules", column) { - t.Fatalf("new firewall desired-rule schema omitted placement column %s", column) - } - } - - policy := model.DockerPortGuardPolicy{ - UUID: "policy-1", Family: "ipv4", HostIP: "0.0.0.0", HostPort: 8080, - Protocol: "tcp", Mode: "allow_all", - } - if err := db.Create(&policy).Error; err != nil { - t.Fatalf("insert first Docker guard policy: %v", err) - } - duplicatePolicy := policy - duplicatePolicy.BaseModel = model.BaseModel{} - duplicatePolicy.UUID = "policy-2" - if err := db.Create(&duplicatePolicy).Error; err == nil { - t.Fatal("Docker guard endpoint uniqueness was not created") - } - - forward := model.ForwardingRule{ - Family: "ipv4", Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", - } - if err := db.Create(&forward).Error; err != nil { - t.Fatalf("insert first forwarding rule: %v", err) - } - duplicateForward := forward - duplicateForward.BaseModel = model.BaseModel{} - if err := db.Create(&duplicateForward).Error; err == nil { - t.Fatal("forwarding identity uniqueness was not created") - } -} - -func TestInitDockerPortGuardStatusCreatesDefaultOnce(t *testing.T) { - db := newFirewallMigrationTestDB(t) - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatal(err) - } - - for i := 0; i < 2; i++ { - if err := InitDockerPortGuardStatus.Migrate(db); err != nil { - t.Fatalf("initialize Docker port guard status on pass %d: %v", i+1, err) - } - } - - var settings []model.Setting - if err := db.Where("key = ?", constant.FirewallDockerPortGuardStatusKey).Find(&settings).Error; err != nil { - t.Fatal(err) - } - if len(settings) != 1 || settings[0].Value != constant.StatusDisable { - t.Fatalf("unexpected default Docker port guard settings: %#v", settings) - } -} - -func TestInitDockerPortGuardStatusPreservesUpgradeValue(t *testing.T) { - db := newFirewallMigrationTestDB(t) - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatal(err) - } - existing := model.Setting{Key: constant.FirewallDockerPortGuardStatusKey, Value: constant.StatusEnable} - if err := db.Create(&existing).Error; err != nil { - t.Fatal(err) - } - - if err := InitDockerPortGuardStatus.Migrate(db); err != nil { - t.Fatal(err) - } - var after model.Setting - if err := db.Where("key = ?", constant.FirewallDockerPortGuardStatusKey).First(&after).Error; err != nil { - t.Fatal(err) - } - if after.ID != existing.ID || after.Value != constant.StatusEnable { - t.Fatalf("migration replaced persisted Docker guard status: before=%#v after=%#v", existing, after) - } -} - -func TestNormalizeFirewallBackendSelections(t *testing.T) { - db := newFirewallMigrationTestDB(t) - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatal(err) - } - settings := []model.Setting{ - {Key: constant.FirewallDockerBackendKey, Value: constant.FirewallProviderNftables}, - {Key: constant.FirewallForwardingBackendKey, Value: constant.FirewallProviderFirewalld}, - } - if err := db.Create(&settings).Error; err != nil { - t.Fatal(err) - } - - for pass := 1; pass <= 2; pass++ { - if err := NormalizeFirewallBackendSelections.Migrate(db); err != nil { - t.Fatalf("normalize firewall backend selections on pass %d: %v", pass, err) - } - } - for key, want := range map[string]string{ - constant.FirewallDockerBackendKey: constant.FirewallProviderNftables, - constant.FirewallForwardingBackendKey: constant.FirewallProviderIptables, - } { - var matches []model.Setting - if err := db.Where("key = ?", key).Find(&matches).Error; err != nil { - t.Fatal(err) - } - if len(matches) != 1 || matches[0].Value != want { - t.Fatalf("setting %s after normalization = %#v, want %q", key, matches, want) - } - } -} - -func TestNormalizeFirewallBackendSelectionsPreservesValidForwardingBackend(t *testing.T) { - db := newFirewallMigrationTestDB(t) - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatal(err) - } - existing := model.Setting{ - Key: constant.FirewallForwardingBackendKey, Value: constant.FirewallProviderNftables, - } - if err := db.Create(&existing).Error; err != nil { - t.Fatal(err) - } - - for pass := 1; pass <= 2; pass++ { - if err := NormalizeFirewallBackendSelections.Migrate(db); err != nil { - t.Fatalf("normalize firewall backend selections on pass %d: %v", pass, err) - } - } - var after model.Setting - if err := db.Where("key = ?", constant.FirewallForwardingBackendKey).First(&after).Error; err != nil { - t.Fatal(err) - } - if after.ID != existing.ID || after.Value != constant.FirewallProviderNftables { - t.Fatalf("migration replaced valid forwarding backend: before=%#v after=%#v", existing, after) - } -} - -func TestNormalizeFirewallBackendSelectionsCreatesMissingDefaults(t *testing.T) { - db := newFirewallMigrationTestDB(t) - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatal(err) - } - if err := NormalizeFirewallBackendSelections.Migrate(db); err != nil { - t.Fatal(err) - } - for key, want := range map[string]string{ - constant.FirewallDockerBackendKey: "", - constant.FirewallForwardingBackendKey: constant.FirewallProviderIptables, - } { - var matches []model.Setting - if err := db.Where("key = ?", key).Find(&matches).Error; err != nil { - t.Fatal(err) - } - if len(matches) != 1 || matches[0].Value != want { - t.Fatalf("missing setting %s after normalization = %#v, want %q", key, matches, want) - } - } -} - -func TestSimplifyFirewallRulePolicyKeepsDesiredRules(t *testing.T) { - db := newFirewallMigrationTestDB(t) - if err := db.Exec(`CREATE TABLE firewall_rules ( - uuid text PRIMARY KEY, - provider text NOT NULL, - family text NOT NULL, - protocol text NOT NULL, - scope_key text NOT NULL, - location text NOT NULL, - native_kind text NOT NULL, - priority integer, - order_index integer, - order_bucket text, - match_key text, - rule_key text NOT NULL - )`).Error; err != nil { - t.Fatal(err) - } - if err := db.Exec( - "INSERT INTO firewall_rules (uuid, provider, family, protocol, scope_key, location, native_kind, priority, order_index, order_bucket, match_key, rule_key) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", - "policy-1", constant.FirewallProviderIptables, constant.FirewallFamilyIPv4, "tcp", - "iptables:ipv4:filter:1PANEL_BASIC:input", "1PANEL_BASIC", "rule", 10, 3, "legacy", "instance:legacy", "legacy-key", - ).Error; err != nil { - t.Fatal(err) - } - if err := db.Exec( - "INSERT INTO firewall_rules (uuid, provider, family, protocol, scope_key, location, native_kind, priority, order_index, order_bucket, match_key, rule_key) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", - "native-1", constant.FirewallProviderFirewalld, constant.FirewallFamilyInet, "all", - "firewalld:inet:public:input", "public", "zone_service", nil, nil, "legacy", "", "native-key", - ).Error; err != nil { - t.Fatal(err) - } - if err := db.Exec( - "INSERT INTO firewall_rules (uuid, provider, family, protocol, scope_key, location, native_kind, priority, order_index, order_bucket, match_key, rule_key) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", - "rich-1", constant.FirewallProviderFirewalld, constant.FirewallFamilyInet, "tcp", - "firewalld:inet:public:input", "public", "rich_rule", -100, nil, "rich_pre", "", "rich-key", - ).Error; err != nil { - t.Fatal(err) - } - for _, statement := range []string{ - "CREATE UNIQUE INDEX uk_firewall_rules_scope_rule ON firewall_rules(scope_key, rule_key)", - "CREATE UNIQUE INDEX uk_firewall_rules_scope_match ON firewall_rules(scope_key, match_key) WHERE match_key <> ''", - } { - if err := db.Exec(statement).Error; err != nil { - t.Fatal(err) - } - } - for pass := 1; pass <= 2; pass++ { - if err := SimplifyFirewallRulePolicy.Migrate(db); err != nil { - t.Fatalf("simplify firewall rule policy on pass %d: %v", pass, err) - } - } - for _, column := range []string{ - "provider", "scope_key", "location", "native_kind", "order_index", "order_bucket", "rule_key", "match_key", - } { - if db.Migrator().HasColumn("firewall_rules", column) { - t.Fatalf("derived column %s was retained in firewall desired rules", column) - } - } - var positional struct { - Priority *int - Sequence *int64 - } - if err := db.Table("firewall_rules").Where("uuid = ?", "policy-1").First(&positional).Error; err != nil { - t.Fatal(err) - } - if positional.Priority != nil || positional.Sequence != nil { - t.Fatalf("positional policy migration inferred unsupported placement: %#v", positional) - } - var count int64 - if err := db.Table("firewall_rules").Count(&count).Error; err != nil { - t.Fatal(err) - } - if count != 3 { - t.Fatalf("policy migration removed desired rules, count=%d", count) - } - var native struct { - CompatibilityError string - } - if err := db.Table("firewall_rules").Where("uuid = ?", "native-1").First(&native).Error; err != nil { - t.Fatal(err) - } - if native.CompatibilityError == "" { - t.Fatal("legacy provider-native rule was not quarantined from synchronization") - } - var rich struct { - Priority *int - Sequence *int64 - } - if err := db.Table("firewall_rules").Where("uuid = ?", "rich-1").First(&rich).Error; err != nil { - t.Fatal(err) - } - if rich.Priority != nil || rich.Sequence != nil { - t.Fatalf("firewalld policy migration retained unsupported placement: %#v", rich) - } -} - -func TestSimplifyFirewallRulePolicyContinuesWhenObsoleteColumnCannotBeDropped(t *testing.T) { - db := newFirewallMigrationTestDB(t) - if err := db.Exec(`CREATE TABLE firewall_rules ( - uuid text PRIMARY KEY, - provider text NOT NULL, - scope_key text NOT NULL, - location text NOT NULL, - native_kind text NOT NULL, - priority integer - )`).Error; err != nil { - t.Fatal(err) - } - if err := db.Exec("CREATE INDEX unexpected_scope_index ON firewall_rules(scope_key)").Error; err != nil { - t.Fatal(err) - } - - if err := SimplifyFirewallRulePolicy.Migrate(db); err != nil { - t.Fatalf("obsolete column cleanup blocked the migration: %v", err) - } - for _, column := range []string{"provider", "location", "native_kind"} { - if db.Migrator().HasColumn("firewall_rules", column) { - t.Fatalf("obsolete column %s was not dropped after an unrelated drop failure", column) - } - } - for _, column := range []string{"compatibility_error", "sequence"} { - if !db.Migrator().HasColumn("firewall_rules", column) { - t.Fatalf("migration omitted required column %s", column) - } - } -} - -func newFirewallMigrationTestDB(t *testing.T) *gorm.DB { - t.Helper() - db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "migration.db")), &gorm.Config{ - Logger: logger.Default.LogMode(logger.Silent), - }) - if err != nil { - t.Fatal(err) - } - return db -} diff --git a/agent/init/migration/migrations/utils/firewall_transfer.go b/agent/init/migration/migrations/utils/firewall_transfer.go index e3a6ceb633d8..7981673bff31 100644 --- a/agent/init/migration/migrations/utils/firewall_transfer.go +++ b/agent/init/migration/migrations/utils/firewall_transfer.go @@ -11,7 +11,6 @@ import ( "github.com/1Panel-dev/1Panel/agent/global" "github.com/1Panel-dev/1Panel/agent/utils/cmd" "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" - forwardingproviders "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding/providers" "github.com/1Panel-dev/1Panel/agent/utils/firewall/iptables_helper" "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" "gorm.io/gorm" @@ -109,7 +108,7 @@ func importLegacyForwardingRules(ctx context.Context, db *gorm.DB, rules []forwa models := make([]model.ForwardingRule, 0, len(rules)) seen := make(map[string]struct{}, len(rules)) for _, rule := range rules { - normalized, err := forwardingproviders.NormalizeRule(rule) + normalized, err := forwarding.NormalizeRule(rule) if err != nil { return fmt.Errorf("normalize legacy forwarding rule: %w", err) } @@ -160,7 +159,7 @@ func loadLegacyFirewallForwarding() (firewallTransferSource, error) { return firewallTransferSource{}, err } for _, item := range firewalldRules { - if _, err := forwardingproviders.NormalizeRule(item.rule); err != nil { + if _, err := forwarding.NormalizeRule(item.rule); err != nil { if global.LOG != nil { global.LOG.Warnf("skip unsupported legacy firewalld forwarding rule %q: %v", item.spec, err) } diff --git a/agent/init/migration/migrations/utils/firewall_transfer_test.go b/agent/init/migration/migrations/utils/firewall_transfer_test.go deleted file mode 100644 index 561ef730d587..000000000000 --- a/agent/init/migration/migrations/utils/firewall_transfer_test.go +++ /dev/null @@ -1,316 +0,0 @@ -package utils - -import ( - "context" - "errors" - "path/filepath" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" - "github.com/glebarez/sqlite" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -func TestFirewallTransferImportsCleansAndMarksCompletion(t *testing.T) { - db := newFirewallTransferTestDB(t) - loadCalls, cleanupCalls := 0, 0 - transfer := &firewallTransfer{ - db: db, - load: func() (firewallTransferSource, error) { - loadCalls++ - legacy := legacyFirewalldForward{ - rule: forwarding.Rule{Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80"}, - spec: "port=8080:proto=tcp:toport=80:toaddr=10.0.0.2", - } - return firewallTransferSource{ - rules: []forwarding.Rule{legacy.rule, legacy.rule}, - firewalld: []legacyFirewalldForward{legacy}, - provider: "iptables", - cleanupOld: func(items []legacyFirewalldForward) error { - cleanupCalls++ - if len(items) != 1 { - return errors.New("unexpected cleanup inventory") - } - return nil - }, - }, nil - }, - } - - if err := transfer.run(context.Background()); err != nil { - t.Fatal(err) - } - if loadCalls != 1 || cleanupCalls != 1 { - t.Fatalf("unexpected calls: load=%d cleanup=%d", loadCalls, cleanupCalls) - } - completed, err := firewallTransferCompleted(db) - if err != nil || !completed { - t.Fatalf("firewall transfer completion = %v, err=%v", completed, err) - } - var setting model.Setting - if err := db.Where("key = ?", "IptablesForwardStatus").First(&setting).Error; err != nil { - t.Fatal(err) - } - if setting.Value != constant.StatusEnable { - t.Fatalf("forwarding status = %q", setting.Value) - } - setting = model.Setting{} - if err := db.Where("key = ?", "ForwardingBackend").First(&setting).Error; err != nil { - t.Fatal(err) - } - if setting.Value != "iptables" { - t.Fatalf("forwarding backend = %q", setting.Value) - } - - if err := transfer.run(context.Background()); err != nil { - t.Fatal(err) - } - if loadCalls != 1 || cleanupCalls != 1 { - t.Fatalf("completed transfer ran again: load=%d cleanup=%d", loadCalls, cleanupCalls) - } -} - -func TestFirewallTransferFailureRemainsRetryable(t *testing.T) { - db := newFirewallTransferTestDB(t) - wantErr := errors.New("cleanup failed") - transfer := &firewallTransfer{ - db: db, - load: func() (firewallTransferSource, error) { - legacy := legacyFirewalldForward{rule: forwarding.Rule{ - Protocol: "udp", Port: "5353", TargetIP: "127.0.0.1", TargetPort: "53", - }} - return firewallTransferSource{ - rules: []forwarding.Rule{legacy.rule}, - firewalld: []legacyFirewalldForward{legacy}, - provider: "iptables", - cleanupOld: func([]legacyFirewalldForward) error { - return wantErr - }, - }, nil - }, - } - if err := transfer.run(context.Background()); !errors.Is(err, wantErr) { - t.Fatalf("transfer error = %v, want %v", err, wantErr) - } - completed, err := firewallTransferCompleted(db) - if err != nil { - t.Fatal(err) - } - if completed { - t.Fatal("failed transfer was marked complete") - } - var count int64 - if err := db.Model(&model.ForwardingRule{}).Count(&count).Error; err != nil { - t.Fatal(err) - } - if count != 1 { - t.Fatalf("retry inventory count = %d, want 1", count) - } -} - -func TestFirewallTransferRetriesCleanupWithoutDuplicatingData(t *testing.T) { - db := newFirewallTransferTestDB(t) - cleanupCalls := 0 - transfer := &firewallTransfer{ - db: db, - load: func() (firewallTransferSource, error) { - legacy := legacyFirewalldForward{ - rule: forwarding.Rule{Protocol: "tcp", Port: "8443", TargetIP: "10.0.0.8", TargetPort: "443"}, - spec: "port=8443:proto=tcp:toport=443:toaddr=10.0.0.8", - } - return firewallTransferSource{ - rules: []forwarding.Rule{legacy.rule}, firewalld: []legacyFirewalldForward{legacy}, provider: "iptables", - cleanupOld: func([]legacyFirewalldForward) error { - cleanupCalls++ - if cleanupCalls == 1 { - return errors.New("temporary cleanup failure") - } - return nil - }, - }, nil - }, - } - - if err := transfer.run(context.Background()); err == nil { - t.Fatal("first cleanup failure was ignored") - } - if err := transfer.run(context.Background()); err != nil { - t.Fatalf("retry firewall transfer: %v", err) - } - var count int64 - if err := db.Model(&model.ForwardingRule{}).Count(&count).Error; err != nil { - t.Fatal(err) - } - if count != 1 || cleanupCalls != 2 { - t.Fatalf("retry result count=%d cleanupCalls=%d", count, cleanupCalls) - } - completed, err := firewallTransferCompleted(db) - if err != nil || !completed { - t.Fatalf("retry completion = %v, err=%v", completed, err) - } -} - -func TestFirewallTransferEmptyInventoryMarksCompletionWithoutEnabling(t *testing.T) { - db := newFirewallTransferTestDB(t) - transfer := &firewallTransfer{ - db: db, - load: func() (firewallTransferSource, error) { - return firewallTransferSource{}, nil - }, - } - if err := transfer.run(context.Background()); err != nil { - t.Fatal(err) - } - completed, err := firewallTransferCompleted(db) - if err != nil || !completed { - t.Fatalf("empty transfer completion = %v, err=%v", completed, err) - } - var count int64 - if err := db.Model(&model.Setting{}).Where("key IN ?", []string{"IptablesForwardStatus", "ForwardingBackend"}).Count(&count).Error; err != nil { - t.Fatal(err) - } - if count != 0 { - t.Fatalf("empty transfer created %d forwarding settings", count) - } -} - -func TestFirewallTransferInvalidInventoryIsAtomicAndRetryable(t *testing.T) { - db := newFirewallTransferTestDB(t) - transfer := &firewallTransfer{ - db: db, - load: func() (firewallTransferSource, error) { - return firewallTransferSource{rules: []forwarding.Rule{ - {Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80"}, - {Protocol: "tcp", Port: "invalid", TargetIP: "127.0.0.1", TargetPort: "80"}, - }}, nil - }, - } - if err := transfer.run(context.Background()); err == nil { - t.Fatal("invalid legacy forwarding inventory was accepted") - } - completed, err := firewallTransferCompleted(db) - if err != nil { - t.Fatal(err) - } - var count int64 - if err := db.Model(&model.ForwardingRule{}).Count(&count).Error; err != nil { - t.Fatal(err) - } - if completed || count != 0 { - t.Fatalf("invalid transfer completed=%v imported=%d", completed, count) - } -} - -func TestFirewallTransferUpdatesExistingSettings(t *testing.T) { - db := newFirewallTransferTestDB(t) - settings := []model.Setting{ - {Key: "IptablesForwardStatus", Value: constant.StatusDisable}, - {Key: "ForwardingBackend", Value: "nftables"}, - } - if err := db.Create(&settings).Error; err != nil { - t.Fatal(err) - } - transfer := &firewallTransfer{ - db: db, - load: func() (firewallTransferSource, error) { - return firewallTransferSource{ - rules: []forwarding.Rule{{Protocol: "udp", Port: "5353", TargetIP: "127.0.0.1", TargetPort: "53"}}, - provider: "iptables", - }, nil - }, - } - if err := transfer.run(context.Background()); err != nil { - t.Fatal(err) - } - for key, want := range map[string]string{ - "IptablesForwardStatus": constant.StatusEnable, - "ForwardingBackend": "iptables", - } { - var matches []model.Setting - if err := db.Where("key = ?", key).Find(&matches).Error; err != nil { - t.Fatal(err) - } - if len(matches) != 1 || matches[0].Value != want { - t.Fatalf("setting %s after transfer = %#v, want %q", key, matches, want) - } - } -} - -func TestFirewallTransferValidatesDependencies(t *testing.T) { - if err := (&firewallTransfer{}).run(context.Background()); err == nil { - t.Fatal("nil transfer database was accepted") - } - db := newFirewallTransferTestDB(t) - if err := (&firewallTransfer{db: db}).run(context.Background()); err == nil { - t.Fatal("nil legacy loader was accepted") - } - transfer := &firewallTransfer{ - db: db, - load: func() (firewallTransferSource, error) { - legacy := legacyFirewalldForward{rule: forwarding.Rule{ - Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80", - }} - return firewallTransferSource{rules: []forwarding.Rule{legacy.rule}, firewalld: []legacyFirewalldForward{legacy}}, nil - }, - } - if err := transfer.run(context.Background()); err == nil { - t.Fatal("missing firewalld cleanup was accepted") - } - completed, err := firewallTransferCompleted(db) - if err != nil || completed { - t.Fatalf("dependency failure completion=%v err=%v", completed, err) - } -} - -func TestParseLegacyFirewalldForwarding(t *testing.T) { - rules := parseLegacyFirewalldForwarding( - "port=8080:proto=tcp:toport=80:toaddr=10.0.0.2\n" + - "port=8443:proto=tcp:toport=443:toaddr=\ninvalid\n", - ) - if len(rules) != 2 { - t.Fatalf("parsed %d firewalld rules", len(rules)) - } - if rules[0].rule.Family != forwarding.FamilyIPv4 || rules[0].rule.TargetIP != "10.0.0.2" { - t.Fatalf("unexpected remote rule: %#v", rules[0]) - } - if rules[1].rule.TargetIP != "127.0.0.1" || rules[1].rule.TargetPort != "443" { - t.Fatalf("unexpected local rule: %#v", rules[1]) - } -} - -func TestParseLegacyIptablesForwarding(t *testing.T) { - stdout := "1 0 0 DNAT 6 -- eth0 * 0.0.0.0/0 0.0.0.0/0 tcp dpt:8080 to:10.0.0.2:80\n" + - "2 0 0 REDIRECT 17 -- * * 0.0.0.0/0 0.0.0.0/0 udp dpts:5353:5354 redir ports 53\n" - rules := parseLegacyIptablesForwarding(stdout) - if len(rules) != 2 { - t.Fatalf("parsed %d iptables rules", len(rules)) - } - if rules[0].Protocol != "tcp" || rules[0].Interface != "eth0" || rules[0].Port != "8080" || - rules[0].TargetIP != "10.0.0.2" || rules[0].TargetPort != "80" { - t.Fatalf("unexpected DNAT rule: %#v", rules[0]) - } - if rules[1].Protocol != "udp" || rules[1].Port != "5353-5354" || - rules[1].TargetIP != "127.0.0.1" || rules[1].TargetPort != "53" { - t.Fatalf("unexpected REDIRECT rule: %#v", rules[1]) - } -} - -func newFirewallTransferTestDB(t *testing.T) *gorm.DB { - t.Helper() - db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "migration.db")), &gorm.Config{ - Logger: logger.Default.LogMode(logger.Silent), - }) - if err != nil { - t.Fatal(err) - } - if err := db.AutoMigrate(&model.Setting{}, &model.ForwardingRule{}); err != nil { - t.Fatal(err) - } - if err := db.Exec("CREATE TABLE migrations (id VARCHAR(255) PRIMARY KEY)").Error; err != nil { - t.Fatal(err) - } - return db -} diff --git a/agent/init/migration/migrations/utils/host_firewall_transfer_test.go b/agent/init/migration/migrations/utils/host_firewall_transfer_test.go deleted file mode 100644 index 84673753ca9a..000000000000 --- a/agent/init/migration/migrations/utils/host_firewall_transfer_test.go +++ /dev/null @@ -1,365 +0,0 @@ -package utils - -import ( - "context" - "errors" - "path/filepath" - "sort" - "testing" - - "github.com/1Panel-dev/1Panel/agent/app/model" - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - "github.com/glebarez/sqlite" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -func TestTransferHostFirewallImportsIptablesRecords(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - records := []legacyHostFirewallRecord{ - {Type: "port", Protocol: "tcp/udp", SrcIP: "Anywhere", DstPort: "80", Strategy: "accept", Description: "web"}, - {Type: "address", SrcIP: "10.0.0.8", Strategy: "drop", Description: "blocked host"}, - {Chain: "1PANEL_BASIC_BEFORE", Protocol: "tcp", DstPort: "22", Strategy: "accept", Description: "ssh first"}, - {Chain: "1PANEL_INPUT", Protocol: "tcp", DstPort: "9000", Strategy: "accept", Description: "unsupported chain"}, - } - if err := db.Table("firewalls").Create(&records).Error; err != nil { - t.Fatal(err) - } - - if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil { - t.Fatal(err) - } - rules := loadTransferredHostFirewallRules(t, db) - if len(rules) != 4 { - t.Fatalf("transferred %d rules, want 4", len(rules)) - } - assertTransferredRuleDefaults(t, rules) - - protocols := make([]string, 0, 2) - foundAddress, foundAdvanced := false, false - for _, rule := range rules { - switch rule.Description { - case "web": - protocols = append(protocols, rule.Protocol) - if rule.Family != string(filter.FamilyIPv4) || rule.DestinationPort != "80" { - t.Fatalf("unexpected port rule: %#v", rule) - } - case "blocked host": - foundAddress = rule.SourceAddress == "10.0.0.8/32" && rule.Action == string(filter.ActionDrop) - case "ssh first": - foundAdvanced = rule.DestinationPort == "22" - case "unsupported chain": - t.Fatal("unsupported legacy chain was imported") - } - } - sort.Strings(protocols) - if len(protocols) != 2 || protocols[0] != "tcp" || protocols[1] != "udp" { - t.Fatalf("port protocols = %#v", protocols) - } - if !foundAddress || !foundAdvanced { - t.Fatalf("address=%v advanced=%v", foundAddress, foundAdvanced) - } - assertHostFirewallTransferCompleted(t, db) - - if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil { - t.Fatal(err) - } - if got := len(loadTransferredHostFirewallRules(t, db)); got != 4 { - t.Fatalf("retry transferred %d rules, want 4", got) - } -} - -func TestTransferHostFirewallMapsFirewalldRepresentations(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - records := []legacyHostFirewallRecord{ - {Type: "port", Protocol: "tcp/udp", DstPort: "80", Strategy: "accept", Description: "zone ports"}, - {Type: "port", Protocol: "tcp", DstPort: "81", Strategy: "drop", Description: "dual family deny"}, - {Type: "port", Protocol: "tcp", SrcIP: "2001:db8::8", DstPort: "443", Strategy: "accept", Description: "v6 source"}, - {Type: "address", SrcIP: "10.0.0.9", Strategy: "drop", Description: "v4 address"}, - } - if err := db.Table("firewalls").Create(&records).Error; err != nil { - t.Fatal(err) - } - - if err := transferHostFirewall(context.Background(), db, filter.ProviderFirewalld); err != nil { - t.Fatal(err) - } - rules := loadTransferredHostFirewallRules(t, db) - if len(rules) != 6 { - t.Fatalf("transferred %d rules, want 6", len(rules)) - } - zonePorts, denyFamilies := 0, make(map[string]bool) - for _, rule := range rules { - switch rule.Description { - case "zone ports": - zonePorts++ - if rule.Family != string(filter.FamilyInet) { - t.Fatalf("unexpected zone port: %#v", rule) - } - case "dual family deny": - denyFamilies[rule.Family] = true - case "v6 source": - if rule.Family != string(filter.FamilyIPv6) || rule.SourceAddress != "2001:db8::8/128" { - t.Fatalf("unexpected v6 rule: %#v", rule) - } - } - } - if zonePorts != 2 || !denyFamilies[string(filter.FamilyIPv4)] || !denyFamilies[string(filter.FamilyIPv6)] { - t.Fatalf("zonePorts=%d denyFamilies=%#v", zonePorts, denyFamilies) - } -} - -func TestTransferHostFirewallExpandsUFWFamilies(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - records := []legacyHostFirewallRecord{ - {Type: "port", Protocol: "tcp", DstPort: "8080", Strategy: "accept", Description: "dual family port"}, - {Type: "address", SrcIP: "10.0.0.1-10.0.0.2", Strategy: "drop", Description: "from to"}, - } - if err := db.Table("firewalls").Create(&records).Error; err != nil { - t.Fatal(err) - } - - if err := transferHostFirewall(context.Background(), db, filter.ProviderUFW); err != nil { - t.Fatal(err) - } - rules := loadTransferredHostFirewallRules(t, db) - if len(rules) != 3 { - t.Fatalf("transferred %d rules, want 3", len(rules)) - } - portFamilies := make(map[string]bool) - foundFromTo := false - for _, rule := range rules { - if rule.Description == "dual family port" { - portFamilies[rule.Family] = true - } - if rule.Description == "from to" { - foundFromTo = rule.SourceAddress == "10.0.0.1/32" && rule.DestinationAddress == "10.0.0.2/32" - } - } - if !portFamilies[string(filter.FamilyIPv4)] || !portFamilies[string(filter.FamilyIPv6)] || !foundFromTo { - t.Fatalf("portFamilies=%#v foundFromTo=%v", portFamilies, foundFromTo) - } -} - -func TestTransferHostFirewallWithoutLegacyTableOnlyMarksCompletion(t *testing.T) { - db := newHostFirewallTransferTestDB(t, false) - if err := transferHostFirewall(context.Background(), db, filter.ProviderNftables); err != nil { - t.Fatal(err) - } - if got := len(loadTransferredHostFirewallRules(t, db)); got != 0 { - t.Fatalf("transferred %d rules without a legacy table", got) - } - assertHostFirewallTransferCompleted(t, db) -} - -func TestTransferHostFirewallRestoresMissingDescription(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - record := legacyHostFirewallRecord{ - Type: "port", Protocol: "tcp", DstPort: "443", Strategy: "accept", Description: "legacy tls", - } - if err := db.Table("firewalls").Create(&record).Error; err != nil { - t.Fatal(err) - } - domainRules, err := legacyHostFirewallRules(record, filter.ProviderIptables) - if err != nil { - t.Fatal(err) - } - existing, err := hostFirewallRuleModel(domainRules[0]) - if err != nil { - t.Fatal(err) - } - existing.Description = "" - existing.Origin = constant.FirewallRuleOriginCreated - if err := db.Create(&existing).Error; err != nil { - t.Fatal(err) - } - - if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil { - t.Fatal(err) - } - rules := loadTransferredHostFirewallRules(t, db) - if len(rules) != 1 || rules[0].Description != "legacy tls" || rules[0].Origin != constant.FirewallRuleOriginCreated { - t.Fatalf("unexpected existing rule after transfer: %#v", rules) - } -} - -func TestTransferHostFirewallReadsDeprecatedColumnsOnDirectUpgrade(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - records := []legacyHostFirewallRecord{ - {Type: "port", Port: "8080", Address: "Anywhere", Protocol: "tcp", Strategy: "accept", Description: "legacy port"}, - {Type: "address", Address: "192.0.2.25", Strategy: "drop", Description: "legacy address"}, - { - Type: "port", Port: "9000", Address: "192.0.2.90", DstPort: "9001", SrcIP: "192.0.2.91", - Protocol: "udp", Strategy: "accept", Description: "new columns win", - }, - } - if err := db.Table("firewalls").Create(&records).Error; err != nil { - t.Fatal(err) - } - - if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil { - t.Fatal(err) - } - rules := loadTransferredHostFirewallRules(t, db) - if len(rules) != 3 { - t.Fatalf("transferred %d direct-upgrade rules, want 3", len(rules)) - } - for _, rule := range rules { - switch rule.Description { - case "legacy port": - if rule.DestinationPort != "8080" || rule.SourceAddress != "" { - t.Fatalf("deprecated port columns were not migrated: %#v", rule) - } - case "legacy address": - if rule.SourceAddress != "192.0.2.25/32" || rule.Action != string(filter.ActionDrop) { - t.Fatalf("deprecated address column was not migrated: %#v", rule) - } - case "new columns win": - if rule.DestinationPort != "9001" || rule.SourceAddress != "192.0.2.91/32" { - t.Fatalf("deprecated columns overwrote normalized columns: %#v", rule) - } - } - } -} - -func TestTransferHostFirewallCoalescesDuplicatesAndSkipsInvalidRows(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - records := []legacyHostFirewallRecord{ - {Type: "port", Protocol: "tcp", DstPort: "443", Strategy: "accept", Description: "first"}, - {Type: "port", Protocol: "tcp", DstPort: "443", Strategy: "accept", Description: "latest"}, - {Type: "port", Protocol: "tcp", DstPort: "invalid", Strategy: "accept", Description: "invalid"}, - } - if err := db.Table("firewalls").Create(&records).Error; err != nil { - t.Fatal(err) - } - - if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil { - t.Fatal(err) - } - rules := loadTransferredHostFirewallRules(t, db) - if len(rules) != 1 || rules[0].Description != "latest" || rules[0].DestinationPort != "443" { - t.Fatalf("unexpected duplicate/invalid migration result: %#v", rules) - } - assertHostFirewallTransferCompleted(t, db) -} - -func TestTransferHostFirewallFailureDoesNotMarkCompletion(t *testing.T) { - t.Run("unsupported provider", func(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - err := transferHostFirewall(context.Background(), db, filter.Provider("unknown")) - if err == nil { - t.Fatal("unsupported provider was accepted") - } - assertHostFirewallTransferNotCompleted(t, db) - }) - - t.Run("cancelled context", func(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - record := legacyHostFirewallRecord{Type: "port", Protocol: "tcp", DstPort: "80", Strategy: "accept"} - if err := db.Table("firewalls").Create(&record).Error; err != nil { - t.Fatal(err) - } - ctx, cancel := context.WithCancel(context.Background()) - cancel() - err := transferHostFirewall(ctx, db, filter.ProviderIptables) - if !errors.Is(err, context.Canceled) { - t.Fatalf("cancelled transfer error = %v", err) - } - assertHostFirewallTransferNotCompleted(t, db) - }) -} - -func TestTransferHostFirewallPreservesExistingDescription(t *testing.T) { - db := newHostFirewallTransferTestDB(t, true) - record := legacyHostFirewallRecord{ - Type: "port", Protocol: "tcp", DstPort: "443", Strategy: "accept", Description: "legacy description", - } - if err := db.Table("firewalls").Create(&record).Error; err != nil { - t.Fatal(err) - } - domainRules, err := legacyHostFirewallRules(record, filter.ProviderIptables) - if err != nil { - t.Fatal(err) - } - existing, err := hostFirewallRuleModel(domainRules[0]) - if err != nil { - t.Fatal(err) - } - existing.Description = "user description" - if err := db.Create(&existing).Error; err != nil { - t.Fatal(err) - } - - if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil { - t.Fatal(err) - } - rules := loadTransferredHostFirewallRules(t, db) - if len(rules) != 1 || rules[0].Description != "user description" { - t.Fatalf("migration overwrote existing description: %#v", rules) - } -} - -func newHostFirewallTransferTestDB(t *testing.T, withLegacyTable bool) *gorm.DB { - t.Helper() - db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "migration.db")), &gorm.Config{ - Logger: logger.Default.LogMode(logger.Silent), - }) - if err != nil { - t.Fatal(err) - } - if err := db.AutoMigrate(&model.FirewallRule{}); err != nil { - t.Fatal(err) - } - if err := db.Exec("CREATE TABLE migrations (id VARCHAR(255) PRIMARY KEY)").Error; err != nil { - t.Fatal(err) - } - if withLegacyTable { - if err := db.Table("firewalls").AutoMigrate(&legacyHostFirewallRecord{}); err != nil { - t.Fatal(err) - } - } - return db -} - -func loadTransferredHostFirewallRules(t *testing.T, db *gorm.DB) []model.FirewallRule { - t.Helper() - var rules []model.FirewallRule - if err := db.Order("family ASC, protocol ASC, destination_port ASC, uuid ASC").Find(&rules).Error; err != nil { - t.Fatal(err) - } - return rules -} - -func assertTransferredRuleDefaults(t *testing.T, rules []model.FirewallRule) { - t.Helper() - for _, rule := range rules { - if rule.UUID == "" || rule.Revision != 1 || rule.CompatibilityError != "" || - rule.Priority != nil || rule.Sequence != nil || - rule.Origin != constant.FirewallRuleOriginAdopted || rule.Owner != constant.FirewallRuleSourceUser { - t.Fatalf("unexpected transferred defaults: %#v", rule) - } - } -} - -func assertHostFirewallTransferCompleted(t *testing.T, db *gorm.DB) { - t.Helper() - completed, err := migrationRecordExists(db, hostFirewallTransferMigrationID) - if err != nil { - t.Fatal(err) - } - if !completed { - t.Fatal("host firewall transfer was not marked complete") - } -} - -func assertHostFirewallTransferNotCompleted(t *testing.T, db *gorm.DB) { - t.Helper() - completed, err := migrationRecordExists(db, hostFirewallTransferMigrationID) - if err != nil { - t.Fatal(err) - } - if completed { - t.Fatal("failed host firewall transfer was marked complete") - } -} diff --git a/agent/utils/docker/docker.go b/agent/utils/docker/docker.go index 7647f8633893..92f95f3b9295 100644 --- a/agent/utils/docker/docker.go +++ b/agent/utils/docker/docker.go @@ -3,6 +3,7 @@ package docker import ( "context" "encoding/json" + "errors" "fmt" "io" "os" @@ -22,6 +23,8 @@ import ( "github.com/docker/docker/client" ) +var ErrUnavailable = errors.New("Docker is unavailable") + func NewDockerClient() (*client.Client, error) { var settingItem model.Setting _ = global.DB.Where("key = ?", "DockerSockPath").First(&settingItem).Error diff --git a/agent/utils/firewall/docker_guard/manager_test.go b/agent/utils/firewall/docker_guard/manager_test.go deleted file mode 100644 index 6c828e27ebf0..000000000000 --- a/agent/utils/firewall/docker_guard/manager_test.go +++ /dev/null @@ -1,255 +0,0 @@ -package docker_guard - -import ( - "errors" - "reflect" - "strings" - "testing" -) - -type restoreCall struct { - executable string - input string - args []string -} - -type recordingRunner struct { - restoreCalls []restoreCall - exists map[string]bool - chains string - dockerRules string - guardRules string - runErr error -} - -func (r *recordingRunner) Run(_ string, args ...string) (string, error) { - if r.runErr != nil { - return "", r.runErr - } - if reflect.DeepEqual(args, []string{"-w", "-t", "filter", "-S"}) { - if r.chains != "" { - return r.chains, nil - } - return "-N DOCKER-USER\n-N 1PANEL_DOCKER\n", nil - } - if reflect.DeepEqual(args, []string{"-w", "-t", "filter", "-S", DockerChain}) { - if r.dockerRules != "" { - return r.dockerRules, nil - } - return "-A DOCKER-USER -j 1PANEL_DOCKER\n", nil - } - if reflect.DeepEqual(args, []string{"-w", "-t", "filter", "-S", Chain}) { - return r.guardRules, nil - } - return "", nil -} - -func (r *recordingRunner) RunInput(executable, input string, args ...string) (string, error) { - r.restoreCalls = append(r.restoreCalls, restoreCall{executable: executable, input: input, args: args}) - return "", nil -} - -func (r *recordingRunner) Exists(executable string) bool { - if r.exists != nil { - return r.exists[executable] - } - return executable == "iptables" || executable == "iptables-restore" -} - -func TestCompilePolicyUsesOriginalDestination(t *testing.T) { - rules := compilePolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "192.0.2.1", HostPort: 8080, Protocol: "tcp", Mode: ModeSources, Sources: []string{"203.0.113.0/24"}}) - got := strings.Join(rules[0], " ") - for _, want := range []string{"--ctorigdst 192.0.2.1", "--ctorigdstport 8080", "-s 203.0.113.0/24", "-j DROP"} { - if !strings.Contains(got, want) { - t.Fatalf("compiled rule %q does not contain %q", got, want) - } - } -} - -func TestCompileWildcardDoesNotMatchUnroutableWildcardAddress(t *testing.T) { - rules := compilePolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 53, Protocol: "udp", Mode: ModeAll}) - got := strings.Join(rules[0], " ") - if strings.Contains(got, "--ctorigdst ") { - t.Fatalf("wildcard binding must not compile an original destination address: %s", got) - } - if !strings.Contains(got, "--ctorigdstport 53") { - t.Fatalf("original destination port missing: %s", got) - } -} - -func TestCompileAllowSourcesReturnsAllowedAndDropsOthers(t *testing.T) { - rules := compilePolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow, Sources: []string{"203.0.113.10/32", "192.0.2.0/24"}}) - if len(rules) != 3 { - t.Fatalf("compiled %d rules, want 3", len(rules)) - } - for i, source := range []string{"203.0.113.10/32", "192.0.2.0/24"} { - got := strings.Join(rules[i], " ") - if !strings.Contains(got, "-s "+source) || !strings.HasSuffix(got, "-j RETURN") { - t.Fatalf("allow rule = %q", got) - } - } - if got := strings.Join(rules[2], " "); strings.Contains(got, " -s ") || !strings.HasSuffix(got, "-j DROP") { - t.Fatalf("fallback rule = %q", got) - } -} - -func TestCompileEmptyAllowSourcesDropsAll(t *testing.T) { - rules := compilePolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow}) - if len(rules) != 1 || !strings.HasSuffix(strings.Join(rules[0], " "), "-j DROP") { - t.Fatalf("rules = %#v", rules) - } -} - -func TestParseIptablesDockerGuardPolicies(t *testing.T) { - output := strings.Join([]string{ - `-A 1PANEL_DOCKER -p tcp -m conntrack --ctorigdstport 8080 -m comment --comment "1panel-docker:deny" -j DROP`, - `-A 1PANEL_DOCKER -p tcp -m conntrack --ctorigdst 192.0.2.10/32 --ctorigdstport 5432 -s 203.0.113.1/32 -m comment --comment "1panel-docker:allow" -j RETURN`, - `-A 1PANEL_DOCKER -p tcp -m conntrack --ctorigdst 192.0.2.10/32 --ctorigdstport 5432 -m comment --comment "1panel-docker:allow" -j DROP`, - }, "\n") - policies, err := parseDockerGuardPolicies(output, FamilyIPv4) - if err != nil { - t.Fatal(err) - } - if len(policies) != 2 || policies[0].UUID != "deny" || policies[0].Mode != ModeAll || - policies[1].UUID != "allow" || policies[1].HostIP != "192.0.2.10" || policies[1].Mode != ModeAllow || - !reflect.DeepEqual(policies[1].Sources, []string{"203.0.113.1/32"}) { - t.Fatalf("policies = %#v", policies) - } -} - -func TestEffectiveJumpMustBeFirstAndUnique(t *testing.T) { - if !hasFirstUniqueJump("-A DOCKER-USER -j 1PANEL_DOCKER\n-A DOCKER-USER -j RETURN\n") { - t.Fatal("expected first unique jump to be effective") - } - if hasFirstUniqueJump("-A DOCKER-USER -j OTHER\n-A DOCKER-USER -j 1PANEL_DOCKER\n") { - t.Fatal("jump after another rule must not be reported effective") - } - if hasFirstUniqueJump("-A DOCKER-USER -j 1PANEL_DOCKER\n-A DOCKER-USER -j 1PANEL_DOCKER\n") { - t.Fatal("duplicate jumps must not be reported effective") - } -} - -func TestReconcileUsesSingleAtomicRestorePerFamily(t *testing.T) { - runner := &recordingRunner{} - manager := NewManagerWithRunner(runner) - policies := []Policy{ - {UUID: "first", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, Protocol: "tcp", Mode: ModeAll}, - {UUID: "second", Family: FamilyIPv4, HostIP: "192.0.2.10", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow, Sources: []string{"203.0.113.1/32"}}, - } - if err := manager.Reconcile(policies); err != nil { - t.Fatal(err) - } - if len(runner.restoreCalls) != 1 { - t.Fatalf("restore calls = %d, want 1", len(runner.restoreCalls)) - } - call := runner.restoreCalls[0] - if call.executable != "iptables-restore" || !reflect.DeepEqual(call.args, []string{"--noflush", "--wait"}) { - t.Fatalf("restore call = %#v", call) - } - for _, want := range []string{ - "*filter\n", - "-F 1PANEL_DOCKER\n", - "-A 1PANEL_DOCKER -m conntrack --ctstate RELATED,ESTABLISHED -j RETURN\n", - "--ctorigdstport 8080", - "--ctorigdst 192.0.2.10 --ctorigdstport 5432", - "-s 203.0.113.1/32", - "-A 1PANEL_DOCKER -j RETURN\nCOMMIT\n", - } { - if !strings.Contains(call.input, want) { - t.Fatalf("restore input does not contain %q:\n%s", want, call.input) - } - } - if strings.Contains(call.input, "-F DOCKER-USER") { - t.Fatalf("restore input must not flush Docker chains:\n%s", call.input) - } -} - -func TestDockerGuardLifecycleRulesBatchCreateAndRebind(t *testing.T) { - output := strings.Join([]string{ - "-N " + DockerChain, - "-A " + DockerChain + " -j " + Chain, - "-A " + DockerChain + " -j " + Chain, - }, "\n") - rules := dockerGuardLifecycleRules(output, true, true) - script, err := buildRestoreScript(rules) - if err != nil { - t.Fatal(err) - } - for _, line := range []string{ - "-N " + Chain, - "-D " + DockerChain + " -j " + Chain, - "-I " + DockerChain + " 1 -j " + Chain, - } { - if !strings.Contains(script, line+"\n") { - t.Fatalf("lifecycle restore is missing %q:\n%s", line, script) - } - } - if strings.Count(script, "-D "+DockerChain+" -j "+Chain+"\n") != 2 || strings.Count(script, "COMMIT\n") != 1 { - t.Fatalf("duplicate jumps were not removed in one transaction:\n%s", script) - } -} - -func TestCleanupRemovesExistingChainWithoutRecreatingIt(t *testing.T) { - runner := &recordingRunner{} - manager := NewManagerWithRunner(runner) - if err := manager.Cleanup(); err != nil { - t.Fatal(err) - } - if len(runner.restoreCalls) != 1 { - t.Fatalf("restore calls = %d, want 1", len(runner.restoreCalls)) - } - script := runner.restoreCalls[0].input - if strings.Contains(script, "-N "+Chain+"\n") { - t.Fatalf("cleanup must not recreate the existing chain:\n%s", script) - } - for _, want := range []string{"-F " + Chain + "\n", "-X " + Chain + "\n"} { - if !strings.Contains(script, want) { - t.Fatalf("cleanup restore is missing %q:\n%s", want, script) - } - } -} - -func TestReconcileReturnsChainInspectionError(t *testing.T) { - manager := NewManagerWithRunner(&recordingRunner{runErr: errors.New("inspect failed")}) - err := manager.Reconcile(nil) - if err == nil { - t.Fatal("expected chain inspection error") - } - var familyErr *FamilyError - if !errors.As(err, &familyErr) || familyErr.Family != FamilyIPv4 { - t.Fatalf("error = %#v, want IPv4 FamilyError", err) - } -} - -func TestBuildRestoreScriptRejectsUnsafeTokens(t *testing.T) { - if _, err := buildRestoreScript([][]string{{"-A", Chain, "--comment", "unsafe value"}}); err == nil { - t.Fatal("expected unsafe token to be rejected") - } -} - -func TestFamilyStatusExplainsIncompleteStep(t *testing.T) { - tests := []struct { - name string - runner *recordingRunner - wantState string - wantReason string - initialized bool - bound bool - }{ - {name: "command missing", runner: &recordingRunner{exists: map[string]bool{}}, wantState: StatusDisabled, wantReason: ReasonCommandMissing}, - {name: "Docker chain missing", runner: &recordingRunner{chains: "-N 1PANEL_DOCKER\n"}, wantState: StatusDisabled, wantReason: ReasonDockerChainMissing}, - {name: "guard chain missing", runner: &recordingRunner{chains: "-N DOCKER-USER\n"}, wantState: StatusDisabled, wantReason: ReasonGuardChainMissing}, - {name: "jump missing", runner: &recordingRunner{dockerRules: "-A DOCKER-USER -j RETURN\n"}, wantState: StatusNotEffective, wantReason: ReasonJumpMissing, initialized: true}, - {name: "jump not first", runner: &recordingRunner{dockerRules: "-A DOCKER-USER -j OTHER\n-A DOCKER-USER -j 1PANEL_DOCKER\n"}, wantState: StatusNotEffective, wantReason: ReasonJumpNotFirst, initialized: true}, - {name: "jump duplicate", runner: &recordingRunner{dockerRules: "-A DOCKER-USER -j 1PANEL_DOCKER\n-A DOCKER-USER -j 1PANEL_DOCKER\n"}, wantState: StatusNotEffective, wantReason: ReasonJumpDuplicate, initialized: true}, - {name: "effective", runner: &recordingRunner{}, wantState: StatusEffective, initialized: true, bound: true}, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - status := NewManagerWithRunner(test.runner).Status(FamilyIPv4) - if status.State != test.wantState || status.Reason != test.wantReason || status.Initialized != test.initialized || status.Bound != test.bound { - t.Fatalf("status = %#v", status) - } - }) - } -} diff --git a/agent/utils/firewall/docker_guard/nftables_test.go b/agent/utils/firewall/docker_guard/nftables_test.go deleted file mode 100644 index 1071534ad8b7..000000000000 --- a/agent/utils/firewall/docker_guard/nftables_test.go +++ /dev/null @@ -1,210 +0,0 @@ -package docker_guard - -import ( - "errors" - "reflect" - "strings" - "testing" -) - -type nftRecordingRunner struct { - objects map[string]bool - baseRules map[string]string - policyRules map[string]string - runCalls [][]string - inputCalls []restoreCall -} - -func newNftRecordingRunner() *nftRecordingRunner { - return &nftRecordingRunner{objects: map[string]bool{}, baseRules: map[string]string{}, policyRules: map[string]string{}} -} - -func (r *nftRecordingRunner) Run(executable string, args ...string) (string, error) { - r.runCalls = append(r.runCalls, append([]string{executable}, args...)) - if executable != "nft" { - return "", errors.New("unexpected executable") - } - plain := args - if len(plain) > 0 && plain[0] == "-a" { - plain = plain[1:] - } - if len(plain) >= 3 && plain[0] == "list" { - key := strings.Join(plain[1:], "|") - if !r.objects[key] { - return "", errors.New("not found") - } - if plain[1] == "chain" && len(plain) == 5 && plain[4] == NftBaseChain { - return r.baseRules[plain[2]], nil - } - if plain[1] == "chain" && len(plain) == 5 && plain[4] == NftChain { - return r.policyRules[plain[2]], nil - } - return "", nil - } - if len(plain) >= 4 && plain[0] == "add" && (plain[1] == "table" || plain[1] == "chain") { - nameLen := 3 - if plain[1] == "chain" { - nameLen = 4 - } - r.objects[strings.Join(plain[1:1+nameLen], "|")] = true - return "", nil - } - if len(plain) >= 7 && plain[0] == "insert" && plain[1] == "rule" { - r.baseRules[plain[2]] = "jump " + NftChain + " # handle 1\n" - return "", nil - } - return "", nil -} - -func (r *nftRecordingRunner) RunInput(executable, input string, args ...string) (string, error) { - r.inputCalls = append(r.inputCalls, restoreCall{executable: executable, input: input, args: args}) - for _, line := range strings.Split(input, "\n") { - fields := strings.Fields(line) - if len(fields) >= 4 && fields[0] == "add" && fields[1] == "table" { - r.objects["table|"+fields[2]+"|"+fields[3]] = true - } - if len(fields) >= 5 && fields[0] == "add" && fields[1] == "chain" { - r.objects["chain|"+fields[2]+"|"+fields[3]+"|"+fields[4]] = true - } - if len(fields) >= 7 && fields[0] == "insert" && fields[1] == "rule" { - r.baseRules[fields[2]] = "jump " + NftChain + " # handle 1\n" - } - } - return "", nil -} - -func (r *nftRecordingRunner) Exists(executable string) bool { return executable == "nft" } - -func (r *nftRecordingRunner) addTable(family, table string) { - r.objects["table|"+family+"|"+table] = true -} - -func (r *nftRecordingRunner) addChain(family, table, chain string) { - r.objects["chain|"+family+"|"+table+"|"+chain] = true -} - -func TestCompileNftPolicyUsesOriginalDestination(t *testing.T) { - rules := compileNftPolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "192.0.2.1", HostPort: 8080, Protocol: "tcp", Mode: ModeSources, Sources: []string{"203.0.113.0/24"}}) - got := strings.Join(rules[0], " ") - for _, want := range []string{"meta l4proto tcp", "ct original ip daddr 192.0.2.1", "ct original proto-dst 8080", "ip saddr 203.0.113.0/24", "drop"} { - if !strings.Contains(got, want) { - t.Fatalf("compiled rule %q does not contain %q", got, want) - } - } -} - -func TestCompileNftWildcardDoesNotMatchWildcardAddress(t *testing.T) { - rules := compileNftPolicy(Policy{UUID: "id", Family: FamilyIPv6, HostIP: "::", HostPort: 53, Protocol: "udp", Mode: ModeAll}) - got := strings.Join(rules[0], " ") - if strings.Contains(got, "ct original ip6 daddr") { - t.Fatalf("wildcard binding must not compile an original destination address: %s", got) - } - if !strings.Contains(got, "ct original proto-dst 53") { - t.Fatalf("original destination port missing: %s", got) - } -} - -func TestCompileNftEmptyAllowSourcesDropsAll(t *testing.T) { - rules := compileNftPolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow}) - if len(rules) != 1 || !strings.HasSuffix(strings.Join(rules[0], " "), "drop") { - t.Fatalf("rules = %#v", rules) - } -} - -func TestParseNftablesDockerGuardPolicies(t *testing.T) { - output := strings.Join([]string{ - `meta l4proto udp ct original proto-dst 53 comment "1panel-docker:deny" drop # handle 3`, - `meta l4proto tcp ct original ip daddr 192.0.2.10 ct original proto-dst 5432 ip saddr 203.0.113.1/32 comment "1panel-docker:allow" return # handle 4`, - `meta l4proto tcp ct original ip daddr 192.0.2.10 ct original proto-dst 5432 comment "1panel-docker:allow" drop # handle 5`, - }, "\n") - policies, err := parseDockerGuardPolicies(output, FamilyIPv4) - if err != nil { - t.Fatal(err) - } - if len(policies) != 2 || policies[0].UUID != "deny" || policies[0].Mode != ModeAll || - policies[1].UUID != "allow" || policies[1].HostIP != "192.0.2.10" || policies[1].Mode != ModeAllow || - !reflect.DeepEqual(policies[1].Sources, []string{"203.0.113.1/32"}) { - t.Fatalf("policies = %#v", policies) - } -} - -func TestNftInitializeCreatesOwnedChainsBeforeDockerForwardRules(t *testing.T) { - runner := newNftRecordingRunner() - runner.addTable("ip", dockerNftTable) - manager := NewNftablesManagerWithRunner(runner) - if err := manager.Initialize(nil); err != nil { - t.Fatal(err) - } - if len(runner.inputCalls) != 2 { - t.Fatalf("batch calls = %d, want lifecycle plus rule restore", len(runner.inputCalls)) - } - all := runner.inputCalls[0].input - for _, want := range []string{ - "add table ip " + NftTable, - "add chain ip " + NftTable + " " + NftBaseChain + " { type filter hook forward priority filter - 1 ; policy accept ; }", - "add chain ip " + NftTable + " " + NftChain, - "insert rule ip " + NftTable + " " + NftBaseChain + " jump " + NftChain, - } { - if !strings.Contains(all, want) { - t.Fatalf("nft lifecycle batch does not contain %q:\n%s", want, all) - } - } - if !strings.Contains(runner.inputCalls[1].input, "flush chain ip "+NftTable+" "+NftChain) { - t.Fatalf("policy restore batch was not executed:\n%s", runner.inputCalls[1].input) - } -} - -func TestNftReconcileUsesSingleAtomicScriptPerFamily(t *testing.T) { - runner := newNftRecordingRunner() - runner.addChain("ip", NftTable, NftChain) - manager := NewNftablesManagerWithRunner(runner) - policies := []Policy{ - {UUID: "first", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, Protocol: "tcp", Mode: ModeAll}, - {UUID: "second", Family: FamilyIPv4, HostIP: "192.0.2.10", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow, Sources: []string{"203.0.113.1/32"}}, - } - if err := manager.Reconcile(policies); err != nil { - t.Fatal(err) - } - if len(runner.inputCalls) != 1 { - t.Fatalf("restore calls = %d, want 1", len(runner.inputCalls)) - } - call := runner.inputCalls[0] - if call.executable != "nft" || !reflect.DeepEqual(call.args, []string{"-f", "-"}) { - t.Fatalf("restore call = %#v", call) - } - for _, want := range []string{ - "flush chain ip " + NftTable + " " + NftChain, - "ct state { established,related } return", - "ct original proto-dst 8080 comment \"1panel-docker:first\" drop", - "ct original ip daddr 192.0.2.10 ct original proto-dst 5432 ip saddr 203.0.113.1/32", - "add rule ip " + NftTable + " " + NftChain + " return", - } { - if !strings.Contains(call.input, want) { - t.Fatalf("restore input does not contain %q:\n%s", want, call.input) - } - } -} - -func TestNftFamilyStatusExplainsIncompleteStep(t *testing.T) { - runner := newNftRecordingRunner() - runner.addTable("ip", dockerNftTable) - runner.addTable("ip", NftTable) - runner.addChain("ip", NftTable, NftBaseChain) - runner.addChain("ip", NftTable, NftChain) - runner.baseRules["ip"] = "counter packets 0 bytes 0 # handle 1\njump " + NftChain + " # handle 2\n" - status := NewNftablesManagerWithRunner(runner).Status(FamilyIPv4) - if status.State != StatusNotEffective || status.Reason != ReasonJumpNotFirst || !status.Initialized { - t.Fatalf("status = %#v", status) - } - runner.baseRules["ip"] = "jump " + NftChain + " # handle 2\n" - status = NewNftablesManagerWithRunner(runner).Status(FamilyIPv4) - if status.State != StatusEffective || !status.Initialized || !status.Bound || !status.Effective { - t.Fatalf("status = %#v", status) - } -} - -func TestBuildNftScriptRejectsUnsafeTokens(t *testing.T) { - if _, err := buildNftScript([][]string{{"add", "rule", "unsafe value"}}); err == nil { - t.Fatal("expected unsafe token to be rejected") - } -} diff --git a/agent/utils/firewall/docker_guard/policy.go b/agent/utils/firewall/docker_guard/policy.go new file mode 100644 index 000000000000..22195886782d --- /dev/null +++ b/agent/utils/firewall/docker_guard/policy.go @@ -0,0 +1,131 @@ +package docker_guard + +import ( + "encoding/json" + "errors" + "fmt" + "net/netip" + "sort" + "strconv" + "strings" +) + +var ErrInvalidPolicy = errors.New("invalid Docker port guard request") + +func NormalizePolicy(policy Policy) (Policy, error) { + policy.Family = strings.ToLower(strings.TrimSpace(policy.Family)) + policy.HostIP = strings.TrimSpace(policy.HostIP) + policy.Protocol = strings.ToLower(strings.TrimSpace(policy.Protocol)) + policy.Mode = strings.ToLower(strings.TrimSpace(policy.Mode)) + if policy.HostPort == 0 || + (policy.Protocol != "tcp" && policy.Protocol != "udp") || + (policy.Family != FamilyIPv4 && policy.Family != FamilyIPv6) || + (policy.Mode != ModeAll && policy.Mode != ModeSources && policy.Mode != ModeAllow) { + return Policy{}, fmt.Errorf("%w: invalid policy fields", ErrInvalidPolicy) + } + address, err := netip.ParseAddr(policy.HostIP) + if err != nil || (policy.Family == FamilyIPv4) != address.Is4() { + return Policy{}, fmt.Errorf("%w: host IP does not match address family", ErrInvalidPolicy) + } + normalizedSources := make([]string, 0, len(policy.Sources)) + seen := make(map[string]struct{}, len(policy.Sources)) + for _, source := range policy.Sources { + source = strings.TrimSpace(source) + if source == "" { + continue + } + prefix, err := netip.ParsePrefix(source) + if err != nil { + if sourceAddress, addressErr := netip.ParseAddr(source); addressErr == nil { + bits := 128 + if sourceAddress.Is4() { + bits = 32 + } + prefix = netip.PrefixFrom(sourceAddress, bits) + } else { + return Policy{}, fmt.Errorf("%w: invalid source address %q", ErrInvalidPolicy, source) + } + } + if (policy.Family == FamilyIPv4) != prefix.Addr().Is4() { + return Policy{}, fmt.Errorf("%w: source %q does not match address family", ErrInvalidPolicy, source) + } + canonical := prefix.Masked().String() + if _, exists := seen[canonical]; !exists { + seen[canonical] = struct{}{} + normalizedSources = append(normalizedSources, canonical) + } + } + if policy.Mode == ModeSources && len(normalizedSources) == 0 { + return Policy{}, fmt.Errorf("%w: deny_sources requires at least one source", ErrInvalidPolicy) + } + if policy.Mode == ModeAll { + normalizedSources = []string{} + } + sort.Strings(normalizedSources) + policy.Sources = normalizedSources + return policy, nil +} + +func NormalizePolicyUUIDs(values []string) ([]string, error) { + uuids := make([]string, 0, len(values)) + seen := make(map[string]struct{}, len(values)) + for _, policyUUID := range values { + policyUUID = strings.TrimSpace(policyUUID) + if policyUUID == "" { + return nil, fmt.Errorf("%w: policy UUID cannot be empty", ErrInvalidPolicy) + } + if _, exists := seen[policyUUID]; exists { + continue + } + seen[policyUUID] = struct{}{} + uuids = append(uuids, policyUUID) + } + if len(uuids) == 0 { + return nil, fmt.Errorf("%w: policy UUIDs cannot be empty", ErrInvalidPolicy) + } + return uuids, nil +} + +func PolicySyncKey(policy Policy) string { + mode := policy.Mode + if mode == ModeAllow && len(policy.Sources) == 0 { + mode = ModeAll + } + sources := append([]string(nil), policy.Sources...) + sort.Strings(sources) + return strings.Join([]string{ + policy.UUID, policy.Family, CanonicalHost(policy.HostIP), strconv.Itoa(int(policy.HostPort)), + policy.Protocol, mode, strings.Join(sources, ","), + }, "\x00") +} + +func PolicyStatesEqual(left, right []Policy) bool { + if len(left) != len(right) { + return false + } + counts := make(map[string]int, len(left)) + for _, policy := range left { + counts[PolicySyncKey(policy)]++ + } + for _, policy := range right { + key := PolicySyncKey(policy) + if counts[key] == 0 { + return false + } + counts[key]-- + } + return true +} + +func CanonicalHost(value string) string { + if address, err := netip.ParseAddr(value); err == nil { + return address.String() + } + return value +} + +func DecodeSources(value string) []string { + result := []string{} + _ = json.Unmarshal([]byte(value), &result) + return result +} diff --git a/agent/utils/firewall/docker_guard/runtime.go b/agent/utils/firewall/docker_guard/runtime.go new file mode 100644 index 000000000000..2d8dc49ca3a1 --- /dev/null +++ b/agent/utils/firewall/docker_guard/runtime.go @@ -0,0 +1,83 @@ +package docker_guard + +import ( + "fmt" + + "github.com/1Panel-dev/1Panel/agent/constant" +) + +type Runtime interface { + Initialize([]Policy) error + Bind() error + Reconcile([]Policy) error + Unbind() error + Cleanup() error + Initialized(string) (bool, error) + Status(string) FamilyStatus + ListPolicies() ([]Policy, error) +} + +func NewRuntime(provider string) Runtime { + if provider == constant.FirewallProviderNftables { + return NewNftablesManager() + } + return NewManager() +} + +func Verify(runtime Runtime, desired []Policy) error { + actual, err := runtime.ListPolicies() + if err != nil { + return fmt.Errorf("verify synchronized Docker firewall policies: %w", err) + } + if !PolicyStatesEqual(actual, desired) { + return fmt.Errorf("verify synchronized Docker firewall policies: target policies do not match the database") + } + return nil +} + +func ReconcileTarget(backend string, policies []Policy, runtime Runtime) error { + families := make(map[string]struct{}, len(policies)) + needsInitialize, needsBind := false, false + for _, policy := range policies { + families[policy.Family] = struct{}{} + } + if len(families) == 0 { + initialized := false + for _, family := range []string{FamilyIPv4, FamilyIPv6} { + status := runtime.Status(family) + if status.Reason == ReasonInspectFailed { + return fmt.Errorf("inspect Docker firewall target %s for %s failed", backend, family) + } + initialized = initialized || status.Initialized + } + if initialized { + return runtime.Reconcile(nil) + } + return nil + } + for family := range families { + status := runtime.Status(family) + needsInitialize = needsInitialize || !status.Initialized + needsBind = needsBind || !status.Bound || !status.Effective + } + var err error + if needsInitialize { + err = runtime.Initialize(policies) + } else { + if needsBind { + err = runtime.Bind() + } + if err == nil { + err = runtime.Reconcile(policies) + } + } + if err != nil { + return err + } + for family := range families { + if !runtime.Status(family).Effective { + return fmt.Errorf("Docker firewall target %s is not effective for %s", backend, family) + } + } + return nil +} diff --git a/agent/utils/firewall/filter/check_flag.go b/agent/utils/firewall/filter/check_flag.go new file mode 100644 index 000000000000..2eee4b2d1a08 --- /dev/null +++ b/agent/utils/firewall/filter/check_flag.go @@ -0,0 +1,170 @@ +package filter + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "encoding/json" + "fmt" + "strings" +) + +type CheckFlagCodec struct { + secret []byte + version int +} + +type checkFlagClaims struct { + Version int `json:"version"` + Provider Provider `json:"provider"` + ScopeKey string `json:"scopeKey"` + RuleDigest string `json:"ruleDigest"` + SnapshotRevision string `json:"snapshotRevision"` + ManagedRevision string `json:"managedRevision"` + Decision CheckDecision `json:"decision"` + Classification CheckClassification `json:"classification"` + AllowedActions []CheckAction `json:"allowedActions"` + AdoptionCandidates []checkFlagAdoptionCandidate `json:"adoptionCandidates,omitempty"` +} + +type checkFlagAdoptionCandidate struct { + InstanceKey string `json:"instanceKey"` + Locator Locator `json:"locator"` +} + +type CreateAuthorization struct { + Operation ChangeOperation + Locator *Locator +} + +func NewCheckFlagCodec(secret []byte, version int) *CheckFlagCodec { + return &CheckFlagCodec{secret: append([]byte(nil), secret...), version: version} +} + +func (c *CheckFlagCodec) Sign(result RuleCheckResult, snapshot Snapshot, managedRevision string) (string, error) { + ruleDigest, err := ruleDigest(result.RequestedRule) + if err != nil { + return "", err + } + claims := checkFlagClaims{ + Version: c.version, + Provider: result.RequestedRule.Scope.Provider, + ScopeKey: result.RequestedRule.Scope.Key(), + RuleDigest: ruleDigest, + SnapshotRevision: snapshot.Revision, + ManagedRevision: managedRevision, + Decision: result.Decision, + Classification: result.Classification, + AllowedActions: result.AllowedActions, + } + if result.Classification == CheckClassificationExactExternal { + claims.AdoptionCandidates = make([]checkFlagAdoptionCandidate, 0, len(result.Candidates)) + for _, candidate := range result.Candidates { + claims.AdoptionCandidates = append(claims.AdoptionCandidates, checkFlagAdoptionCandidate{ + InstanceKey: candidate.InstanceKey, + Locator: candidate.Locator, + }) + } + } + payload, err := json.Marshal(claims) + if err != nil { + return "", err + } + signature := c.signature(payload) + return base64.RawURLEncoding.EncodeToString(payload) + "." + base64.RawURLEncoding.EncodeToString(signature), nil +} + +func (c *CheckFlagCodec) Authorize( + checkFlag string, + action CheckAction, + adoptInstanceKey string, + rule FirewallRule, + snapshot Snapshot, + managedRevision string, +) (CreateAuthorization, error) { + claims, err := c.parse(checkFlag) + if err != nil { + return CreateAuthorization{}, err + } + ruleDigest, err := ruleDigest(rule) + if err != nil { + return CreateAuthorization{}, err + } + if claims.Version != c.version || + claims.Provider != rule.Scope.Provider || + claims.ScopeKey != rule.Scope.Key() || + claims.RuleDigest != ruleDigest || + claims.SnapshotRevision != snapshot.Revision || + claims.ManagedRevision != managedRevision { + return CreateAuthorization{}, fmt.Errorf("%w: firewall or managed rules changed", ErrRuleCheckRequired) + } + if claims.Decision != CheckDecisionReady && claims.Decision != CheckDecisionConfirmationRequired { + return CreateAuthorization{}, ErrRuleOperation + } + if !containsCheckAction(claims.AllowedActions, action) { + return CreateAuthorization{}, ErrRuleOperation + } + + switch action { + case CheckActionCreate, CheckActionCreateAnyway: + if strings.TrimSpace(adoptInstanceKey) != "" { + return CreateAuthorization{}, ErrRuleOperation + } + return CreateAuthorization{Operation: ChangeCreate}, nil + case CheckActionAdopt, CheckActionSelectAdopt: + for _, candidate := range claims.AdoptionCandidates { + if candidate.InstanceKey == adoptInstanceKey && adoptInstanceKey != "" { + locator := candidate.Locator + return CreateAuthorization{Operation: ChangeAdopt, Locator: &locator}, nil + } + } + return CreateAuthorization{}, ErrRuleOperation + default: + return CreateAuthorization{}, ErrRuleOperation + } +} + +func (c *CheckFlagCodec) parse(checkFlag string) (checkFlagClaims, error) { + parts := strings.Split(strings.TrimSpace(checkFlag), ".") + if len(parts) != 2 || parts[0] == "" || parts[1] == "" { + return checkFlagClaims{}, ErrRuleCheckRequired + } + payload, err := base64.RawURLEncoding.DecodeString(parts[0]) + if err != nil { + return checkFlagClaims{}, ErrRuleCheckRequired + } + signature, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil || !hmac.Equal(signature, c.signature(payload)) { + return checkFlagClaims{}, ErrRuleCheckRequired + } + var claims checkFlagClaims + if err := json.Unmarshal(payload, &claims); err != nil { + return checkFlagClaims{}, ErrRuleCheckRequired + } + return claims, nil +} + +func (c *CheckFlagCodec) signature(payload []byte) []byte { + mac := hmac.New(sha256.New, c.secret) + _, _ = mac.Write(payload) + return mac.Sum(nil) +} + +func ruleDigest(rule FirewallRule) (string, error) { + payload, err := json.Marshal(rule) + if err != nil { + return "", err + } + sum := sha256.Sum256(payload) + return hex.EncodeToString(sum[:]), nil +} + +func containsCheckAction(actions []CheckAction, expected CheckAction) bool { + for _, action := range actions { + if action == expected { + return true + } + } + return false +} diff --git a/agent/utils/firewall/filter/check_test.go b/agent/utils/firewall/filter/check_test.go deleted file mode 100644 index 4d0723e83916..000000000000 --- a/agent/utils/firewall/filter/check_test.go +++ /dev/null @@ -1,366 +0,0 @@ -package filter - -import "testing" - -func TestCheckCreateRequestsAdoptionForEquivalentExternalRule(t *testing.T) { - rule := checkAddressRule("172.16.10.111", ActionDrop) - snapshot := checkSnapshot(t, rule) - plan, err := CheckCreate(snapshot, rule, nil, "") - if err != nil { - t.Fatalf("plan create: %v", err) - } - if plan.Decision != CheckDecisionConfirmationRequired || plan.Classification != CheckClassificationExactExternal || plan.Reason != "equivalent_external_rule" { - t.Fatalf("unexpected adoption check: %#v", plan) - } - if len(plan.Candidates) != 1 || len(plan.AllowedActions) != 2 || plan.AllowedActions[0] != CheckActionAdopt || plan.Candidates[0].InstanceKey == "" { - t.Fatalf("adoption details missing: %#v", plan) - } -} - -func TestCheckCreateTreatsManagedDuplicateAsIdempotent(t *testing.T) { - rule := checkPortRule("22", ActionAccept) - marker := "onepanel:created:ssh" - observed := checkObservedRule(rule, marker, 1) - snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{observed}) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - ruleKey, _ := RuleKey(rule) - desired := DesiredRule{ - UUID: "managed", Rule: rule, RuleKey: ruleKey, Origin: RuleOriginCreated, Marker: marker, - } - - plan, err := CheckCreate(snapshot, rule, []DesiredRule{desired}, "") - if err != nil { - t.Fatalf("plan managed duplicate: %v", err) - } - if plan.Decision != CheckDecisionNoChange || plan.Classification != CheckClassificationExactManaged || plan.ExistingRuleUUID != desired.UUID { - t.Fatalf("managed duplicate was not idempotent: %#v", plan) - } -} - -func TestCheckCreateBlocksMissingManagedRuleInsteadOfCreatingDuplicate(t *testing.T) { - rule := checkPortRule("22", ActionAccept) - snapshot, err := NewSnapshot(rule.Scope, nil) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - ruleKey, _ := RuleKey(rule) - desired := DesiredRule{UUID: "managed", Rule: rule, RuleKey: ruleKey, Origin: RuleOriginCreated} - plan, err := CheckCreate(snapshot, rule, []DesiredRule{desired}, "") - if err != nil { - t.Fatalf("plan missing managed rule: %v", err) - } - if plan.Decision != CheckDecisionBlocked || plan.Reason != "managed_rule_drifted" { - t.Fatalf("missing managed rule was recreated: %#v", plan) - } -} - -func TestCheckCreateRequiresCandidateSelectionForDuplicates(t *testing.T) { - rule := checkPortRule("3306", ActionAccept) - first := checkObservedRule(rule, "", 1) - second := checkObservedRule(rule, "", 2) - snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{first, second}) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - plan, err := CheckCreate(snapshot, rule, nil, "") - if err != nil { - t.Fatalf("plan duplicate candidates: %v", err) - } - if plan.Decision != CheckDecisionConfirmationRequired || plan.Reason != "multiple_equivalent_external_rules" || len(plan.Candidates) != 2 || plan.AllowedActions[0] != CheckActionSelectAdopt { - t.Fatalf("check guessed an equivalent candidate: %#v", plan) - } - if plan.Candidates[0].InstanceKey == "" || plan.Candidates[1].InstanceKey == "" || plan.Candidates[0].InstanceKey == plan.Candidates[1].InstanceKey { - t.Fatalf("candidate instance keys must be present and unique: %+v", plan.Candidates) - } - if _, err := FindCandidate(plan.Candidates, ""); err == nil { - t.Fatalf("expected candidate selection to be required, got %v", err) - } - secondKey, _ := InstanceKey(second) - candidate, err := FindCandidate(plan.Candidates, secondKey) - if err != nil || candidate.Locator.Position == nil || *candidate.Locator.Position != 2 { - t.Fatalf("selected candidate was not found: candidate=%#v err=%v", candidate, err) - } -} - -func TestCheckCreateBlocksAdoptionOfProtectedRule(t *testing.T) { - rule := checkPortRule("22", ActionAccept) - observed := checkObservedRule(rule, "", 1) - observed.Protected = true - snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{observed}) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - - plan, err := CheckCreate(snapshot, rule, nil, "") - if err != nil { - t.Fatalf("plan protected rule: %v", err) - } - if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationProtected || plan.Reason != "protected_rule" || len(plan.AllowedActions) != 0 { - t.Fatalf("protected rule could be adopted: %#v", plan) - } -} - -func TestCheckCreateBlocksRuntimePermanentMismatch(t *testing.T) { - rule := checkPortRule("443", ActionAccept) - observed := checkObservedRule(rule, "", 1) - observed.Persistence = PersistenceStatusRuntimeOnly - snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{observed}) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - plan, err := CheckCreate(snapshot, rule, nil, "") - if err != nil { - t.Fatalf("plan persistence drift: %v", err) - } - if plan.Decision != CheckDecisionBlocked || plan.Reason != "runtime_permanent_mismatch" { - t.Fatalf("runtime-only rule could be adopted: %#v", plan) - } -} - -func TestCheckCreateDoesNotLetOpaqueFirewalldServiceBlockUnrelatedRule(t *testing.T) { - rule := FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput}, - NativeKind: NativeKindZonePort, Protocol: "tcp", DestinationPort: "8080", Action: ActionAccept, - } - service := ObservedRule{ - Rule: FirewallRule{Scope: rule.Scope, NativeKind: NativeKindZoneService}, - Locator: Locator{Provider: ProviderFirewalld, ScopeKey: rule.Scope.Key(), Canonical: "service:ssh"}, - ParseStatus: ParseStatusOpaque, Raw: "ssh", Persistence: PersistenceStatusConverged, - } - snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{service}) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - plan, err := CheckCreate(snapshot, rule, nil, "") - if err != nil { - t.Fatalf("plan with opaque service: %v", err) - } - if plan.Decision != CheckDecisionReady || plan.Classification != CheckClassificationNone { - t.Fatalf("opaque service blocked unrelated native port: %#v", plan) - } - opaqueRich := service - opaqueRich.Rule.NativeKind = NativeKindRichRule - opaqueRich.Locator.Canonical = `rich:rule log prefix="audit" accept` - opaqueRich.Raw = `rule log prefix="audit" accept` - snapshot, err = NewSnapshot(rule.Scope, []ObservedRule{opaqueRich}) - if err != nil { - t.Fatalf("opaque rich snapshot: %v", err) - } - plan, err = CheckCreate(snapshot, rule, nil, "") - if err != nil { - t.Fatalf("plan with opaque rich rule: %v", err) - } - if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationUnsupported { - t.Fatalf("opaque rich rule did not block an unsafe plan: %#v", plan) - } -} - -func TestCheckCreateClassifiesCoverageAndConflict(t *testing.T) { - requested := checkAddressRule("172.16.10.111", ActionDrop) - covering := checkAddressRule("172.16.10.0/24", ActionDrop) - coveredSnapshot := checkSnapshot(t, covering) - plan, err := CheckCreate(coveredSnapshot, requested, nil, "") - if err != nil { - t.Fatalf("plan covered rule: %v", err) - } - if plan.Classification != CheckClassificationCovered || plan.Decision != CheckDecisionConfirmationRequired || plan.AllowedActions[0] != CheckActionCreateAnyway { - t.Fatalf("unexpected covered plan: %#v", plan) - } - - conflicting := checkAddressRule("172.16.10.0/24", ActionAccept) - conflictSnapshot := checkSnapshot(t, conflicting) - plan, err = CheckCreate(conflictSnapshot, requested, nil, "") - if err != nil { - t.Fatalf("plan conflicting rule: %v", err) - } - if plan.Classification != CheckClassificationConflict || plan.Decision != CheckDecisionBlocked { - t.Fatalf("unexpected conflict plan: %#v", plan) - } -} - -func TestCheckCreateAllowsPartialOverlapWithOppositeActionAfterConfirmation(t *testing.T) { - scope := Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: FirewalldInputZone, Direction: DirectionInput} - existingRule := FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: FirewalldInputZone, Direction: DirectionInput}, - NativeKind: NativeKindRichRule, Protocol: "all", SourceAddress: "1.1.1.1", Action: ActionDrop, - } - requested := FirewallRule{ - Scope: scope, NativeKind: NativeKindZonePort, Protocol: "tcp", DestinationPort: "8080", Action: ActionAccept, - } - snapshot := checkSnapshot(t, existingRule) - plan, err := CheckCreate(snapshot, requested, nil, "") - if err != nil { - t.Fatalf("plan partially overlapping rule: %v", err) - } - if plan.Decision != CheckDecisionConfirmationRequired || plan.Classification != CheckClassificationConflict || - plan.Reason != "partially_overlapping_rule_with_different_action" || - len(plan.AllowedActions) == 0 || plan.AllowedActions[0] != CheckActionCreateAnyway { - t.Fatalf("partial overlap was not confirmable: %#v", plan) - } -} - -func TestRuleCoverageAndOverlapRespectAddressFamilies(t *testing.T) { - base := Scope{Provider: ProviderFirewalld, Zone: FirewalldInputZone, Direction: DirectionInput} - ipv4 := FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: base.Zone, Direction: base.Direction}, - NativeKind: NativeKindRichRule, Protocol: "all", SourceAddress: "1.1.1.1", Action: ActionDrop, - } - ipv6 := FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv6, Zone: base.Zone, Direction: base.Direction}, - NativeKind: NativeKindRichRule, Protocol: "tcp", DestinationPort: "8080", Action: ActionAccept, - } - if RulesOverlap(ipv4, ipv6) || RuleCovers(ipv4, ipv6) { - t.Fatal("disjoint IPv4 and IPv6 rules must not overlap") - } - inet := ipv6 - inet.Scope.Family = FamilyInet - inet.SourceAddress = "" - if !RulesOverlap(ipv4, inet) { - t.Fatal("inet rule should overlap an IPv4 rule") - } - if RuleCovers(ipv4, inet) { - t.Fatal("IPv4 rule must not cover a dual-stack inet rule") - } -} - -func TestCheckCreateAllowsOrderedFirewalldDenyBeforeNativePort(t *testing.T) { - scope := Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: "public", Direction: DirectionInput} - existingRule := FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput}, - NativeKind: NativeKindZonePort, Protocol: "tcp", DestinationPort: "3306", Action: ActionAccept, - } - priority := -100 - requested := FirewallRule{ - Scope: scope, NativeKind: NativeKindRichRule, Protocol: "tcp", SourceAddress: "172.16.10.111", - DestinationPort: "3306", Action: ActionDrop, Priority: &priority, - } - existing := checkObservedRule(existingRule, "", 1) - existing.Persistence = PersistenceStatusConverged - snapshot, err := NewSnapshot(scope, []ObservedRule{existing}) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - plan, err := CheckCreate(snapshot, requested, nil, "") - if err != nil { - t.Fatalf("plan ordered deny: %v", err) - } - if plan.Decision != CheckDecisionReady || plan.Classification != CheckClassificationNone { - t.Fatalf("negative-priority deny was blocked: %#v", plan) - } - - existing.Protected = true - protectedSnapshot, _ := NewSnapshot(scope, []ObservedRule{existing}) - plan, err = CheckCreate(protectedSnapshot, requested, nil, "") - if err != nil { - t.Fatalf("plan protected overlap: %v", err) - } - if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationConflict { - t.Fatalf("protected port overlap was allowed: %#v", plan) - } -} - -func TestRuleCoverageAndOverlapUseNormalizedRanges(t *testing.T) { - existing := checkPortRule("8000-9000", ActionAccept) - requested := checkPortRule("8080", ActionAccept) - if !RuleCovers(existing, requested) || !RulesOverlap(existing, requested) { - t.Fatal("expected port range to cover and overlap contained port") - } - differentProtocol := requested - differentProtocol.Protocol = "udp" - if RulesOverlap(existing, differentProtocol) { - t.Fatal("different transport protocols should not overlap") - } -} - -func TestCheckCreateBlocksRuleTargetingCurrentManagementClient(t *testing.T) { - rule := checkAddressRule("203.0.113.9", ActionDrop) - snapshot, err := NewSnapshot(rule.Scope, nil) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - plan, err := CheckCreate(snapshot, rule, nil, "203.0.113.9") - if err != nil { - t.Fatalf("plan create: %v", err) - } - if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationProtected || plan.Reason != "current_management_connection" { - t.Fatalf("unexpected current connection decision: %#v", plan) - } -} - -func TestCheckCreateBlocksProtectedManagementPortWithoutObservedAllowRule(t *testing.T) { - rule := checkPortRule("22", ActionDrop) - snapshot, err := NewSnapshot(rule.Scope, nil) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - plan, err := CheckCreate(snapshot, rule, nil, "203.0.113.9", PortWhitelist{ - Family: "ipv4", Port: "22", Protocol: "tcp", - }) - if err != nil { - t.Fatalf("plan create: %v", err) - } - if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationProtected || plan.Reason != "current_management_connection" { - t.Fatalf("protected management port was not blocked: %#v", plan) - } -} - -func TestCheckCreateAllowsProtectedPortDenyForUnrelatedSource(t *testing.T) { - rule := checkPortRule("22", ActionDrop) - rule.SourceAddress = "198.51.100.8/32" - snapshot, err := NewSnapshot(rule.Scope, nil) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - plan, err := CheckCreate(snapshot, rule, nil, "203.0.113.9", PortWhitelist{ - Family: "ipv4", Port: "22", Protocol: "tcp", - }) - if err != nil { - t.Fatalf("plan create: %v", err) - } - if plan.Decision != CheckDecisionReady { - t.Fatalf("unrelated source was incorrectly blocked: %#v", plan) - } -} - -func checkSnapshot(t *testing.T, rules ...FirewallRule) Snapshot { - t.Helper() - observed := make([]ObservedRule, 0, len(rules)) - for index, rule := range rules { - observed = append(observed, checkObservedRule(rule, "", index+1)) - } - snapshot, err := NewSnapshot(rules[0].Scope, observed) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - return snapshot -} - -func checkObservedRule(rule FirewallRule, marker string, position int) ObservedRule { - return ObservedRule{ - Rule: rule, - Locator: Locator{Provider: rule.Scope.Provider, ScopeKey: rule.Scope.Key(), Position: &position}, - Marker: marker, ParseStatus: ParseStatusSupported, - } -} - -func checkAddressRule(address string, action Action) FirewallRule { - return FirewallRule{ - Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput}, - NativeKind: NativeKindRule, - Protocol: "all", - SourceAddress: address, - Action: action, - } -} - -func checkPortRule(port string, action Action) FirewallRule { - return FirewallRule{ - Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput}, - NativeKind: NativeKindRule, - Protocol: "tcp", - DestinationPort: port, - Action: action, - } -} diff --git a/agent/utils/firewall/filter/identity_test.go b/agent/utils/firewall/filter/identity_test.go deleted file mode 100644 index ad1484839462..000000000000 --- a/agent/utils/firewall/filter/identity_test.go +++ /dev/null @@ -1,195 +0,0 @@ -package filter - -import ( - "errors" - "testing" -) - -func TestRuleKeyUsesSemanticFields(t *testing.T) { - priority := 10 - base := FirewallRule{ - UUID: "first", - Scope: Scope{ - Provider: ProviderFirewalld, - Family: FamilyIPv4, - Zone: "public", - Direction: DirectionInput, - }, - NativeKind: NativeKindRichRule, - Protocol: "TCP", - SourceAddress: "172.16.10.111", - DestinationPort: "3306:3306", - Action: "ALLOW", - Priority: &priority, - Description: "first description", - } - other := base - other.UUID = "second" - other.Description = "description is not identity" - other.SourceAddress = "172.16.10.111/32" - other.DestinationPort = "3306" - - firstKey, err := RuleKey(base) - if err != nil { - t.Fatalf("first rule key: %v", err) - } - secondKey, err := RuleKey(other) - if err != nil { - t.Fatalf("second rule key: %v", err) - } - if firstKey != secondKey { - t.Fatalf("equivalent rules produced different keys:\n%s\n%s", firstKey, secondKey) - } - - changedPriority := 11 - other.Priority = &changedPriority - changedKey, err := RuleKey(other) - if err != nil { - t.Fatalf("changed rule key: %v", err) - } - if changedKey == firstKey { - t.Fatal("priority change did not change rule key") - } -} - -func TestRuleKeyKeepsFamilyInsideSharedFirewalldScope(t *testing.T) { - ipv4 := FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: "public", Direction: DirectionInput}, - NativeKind: NativeKindRichRule, Protocol: "tcp", Action: ActionAccept, - } - ipv6 := ipv4 - ipv6.Scope.Family = FamilyIPv6 - first, err := RuleKey(ipv4) - if err != nil { - t.Fatalf("IPv4 key: %v", err) - } - second, err := RuleKey(ipv6) - if err != nil { - t.Fatalf("IPv6 key: %v", err) - } - if first == second { - t.Fatal("firewalld family-specific rich rules collapsed to one identity") - } -} - -func TestInstanceKeyIncludesLocator(t *testing.T) { - positionOne := 1 - positionTwo := 2 - rule := ObservedRule{ - Rule: FirewallRule{ - Scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput}, - Protocol: "tcp", - DestinationPort: "22", - Action: ActionAccept, - }, - Marker: "1panel-rule:test", - ParseStatus: ParseStatusSupported, - Locator: Locator{Position: &positionOne}, - } - first, err := InstanceKey(rule) - if err != nil { - t.Fatalf("first instance key: %v", err) - } - rule.Locator.Position = &positionTwo - second, err := InstanceKey(rule) - if err != nil { - t.Fatalf("second instance key: %v", err) - } - if first == second { - t.Fatal("position change did not change instance key") - } - rule.Locator.Position = &positionOne - rule.Persistence = PersistenceStatusRuntimeOnly - runtimeOnly, err := InstanceKey(rule) - if err != nil { - t.Fatalf("runtime-only instance key: %v", err) - } - if runtimeOnly == first { - t.Fatal("runtime/permanent presence did not change instance key") - } -} - -func TestSnapshotRevisionIsInputOrderIndependent(t *testing.T) { - scope := Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput} - positionOne := 1 - positionTwo := 2 - rules := []ObservedRule{ - { - Rule: FirewallRule{Scope: scope, Protocol: "tcp", DestinationPort: "22", Action: ActionAccept}, - Locator: Locator{Position: &positionOne}, - ParseStatus: ParseStatusSupported, - }, - { - Rule: FirewallRule{Scope: scope, Protocol: "tcp", DestinationPort: "80", Action: ActionAccept}, - Locator: Locator{Position: &positionTwo}, - ParseStatus: ParseStatusSupported, - }, - } - - first, err := SnapshotRevision(scope, rules) - if err != nil { - t.Fatalf("first snapshot revision: %v", err) - } - reversed := []ObservedRule{rules[1], rules[0]} - second, err := SnapshotRevision(scope, reversed) - if err != nil { - t.Fatalf("second snapshot revision: %v", err) - } - if first != second { - t.Fatalf("slice order changed snapshot revision:\n%s\n%s", first, second) - } - - rules[0].Locator.Position = &positionTwo - changed, err := SnapshotRevision(scope, rules) - if err != nil { - t.Fatalf("changed snapshot revision: %v", err) - } - if changed == first { - t.Fatal("locator position change did not change snapshot revision") - } -} - -func TestSnapshotRevisionSupportsOpaqueRules(t *testing.T) { - scope := Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput} - position := 1 - rules := []ObservedRule{ - { - Rule: FirewallRule{Scope: scope}, - Locator: Locator{Position: &position, Canonical: "service:ssh"}, - ParseStatus: ParseStatusOpaque, - Raw: "services: ssh", - }, - } - revision, err := SnapshotRevision(scope, rules) - if err != nil { - t.Fatalf("opaque snapshot revision: %v", err) - } - if revision == "" { - t.Fatal("expected non-empty opaque snapshot revision") - } - instanceKey, err := InstanceKey(rules[0]) - if err != nil || instanceKey == "" { - t.Fatalf("opaque instance key: key=%q err=%v", instanceKey, err) - } - rules[0].Persistence = PersistenceStatusPermanentOnly - changed, err := SnapshotRevision(scope, rules) - if err != nil { - t.Fatalf("opaque persistence revision: %v", err) - } - if changed == revision { - t.Fatal("opaque runtime/permanent presence did not change snapshot revision") - } -} - -func TestSnapshotRevisionRequiresPositionForOrderedProvider(t *testing.T) { - scope := Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput} - _, err := SnapshotRevision(scope, []ObservedRule{ - { - Rule: FirewallRule{Scope: scope, Protocol: "tcp", DestinationPort: "22", Action: ActionAccept}, - ParseStatus: ParseStatusSupported, - }, - }) - if !errors.Is(err, ErrInvalidRule) { - t.Fatalf("expected missing position error, got %v", err) - } -} diff --git a/agent/utils/firewall/filter/inventory_test.go b/agent/utils/firewall/filter/inventory_test.go deleted file mode 100644 index 51a9d1be5d27..000000000000 --- a/agent/utils/firewall/filter/inventory_test.go +++ /dev/null @@ -1,256 +0,0 @@ -package filter - -import "testing" - -func TestMergeInventoryPreservesObservedOrderAndOwnership(t *testing.T) { - ssh := inventoryTestRule("22") - http := inventoryTestRule("80") - https := inventoryTestRule("443") - sshObserved := inventoryObservedRule(t, ssh, "onepanel:created:ssh", 1) - httpObserved := inventoryObservedRule(t, http, "", 2) - protectedKey, _ := RuleKey(http) - - items, err := MergeInventory(InventoryMergeInput{ - Observed: []ObservedRule{httpObserved, sshObserved}, - Desired: []DesiredRule{ - {UUID: "ssh", Rule: ssh, Origin: RuleOriginCreated, Marker: "onepanel:created:ssh"}, - {UUID: "https", Rule: https, Origin: RuleOriginAdopted}, - }, - ProtectedObservedKeys: map[string]struct{}{protectedKey: {}}, - }) - if err != nil { - t.Fatalf("merge inventory: %v", err) - } - if len(items) != 3 { - t.Fatalf("unexpected item count: %#v", items) - } - if items[0].Rule.DestinationPort != "80" || items[0].State != InventoryStateProtected || items[0].Desired != nil { - t.Fatalf("observed order/protected classification changed: %#v", items[0]) - } - if items[1].Rule.DestinationPort != "22" || items[1].State != InventoryStateManaged || items[1].Match != InventoryMatchExact { - t.Fatalf("managed rule was not matched: %#v", items[1]) - } - if items[2].Rule.DestinationPort != "443" || items[2].State != InventoryStateDrifted || items[2].Match != InventoryMatchMissing || items[2].Observed != nil { - t.Fatalf("missing desired rule was not appended as drifted: %#v", items[2]) - } -} - -func TestMergeInventoryUsesDesiredDescriptionForManagedRule(t *testing.T) { - rule := inventoryTestRule("8080") - desired := rule - desired.Description = "user description" - observed := inventoryObservedRule(t, rule, "1panel-rule:web", 1) - - items, err := MergeInventory(InventoryMergeInput{ - Observed: []ObservedRule{observed}, - Desired: []DesiredRule{{ - UUID: "web", Rule: desired, Origin: RuleOriginCreated, Marker: observed.Marker, - }}, - }) - if err != nil { - t.Fatalf("merge inventory: %v", err) - } - if len(items) != 1 || items[0].Rule.Description != desired.Description { - t.Fatalf("managed description was not taken from desired state: %#v", items) - } - if items[0].Observed == nil || items[0].Observed.Rule.Description != "" { - t.Fatalf("observed runtime rule was modified: %#v", items[0].Observed) - } -} - -func TestMergeInventoryUsesInstanceBeforeSemanticIdentity(t *testing.T) { - rule := inventoryTestRule("22") - first := inventoryObservedRule(t, rule, "", 1) - second := inventoryObservedRule(t, rule, "", 2) - secondKey, err := InstanceKey(second) - if err != nil { - t.Fatalf("instance key: %v", err) - } - - items, err := MergeInventory(InventoryMergeInput{ - Observed: []ObservedRule{first, second}, - Desired: []DesiredRule{{ - UUID: "adopted", Rule: rule, Origin: RuleOriginAdopted, ObservedInstanceKey: secondKey, - }}, - }) - if err != nil { - t.Fatalf("merge inventory: %v", err) - } - if items[0].State != InventoryStateExternal || items[1].State != InventoryStateAdopted || items[1].Observed.Locator.Position == nil || *items[1].Observed.Locator.Position != 2 { - t.Fatalf("instance locator did not select the intended candidate: %#v", items) - } -} - -func TestMergeInventoryDoesNotGuessBetweenEquivalentRules(t *testing.T) { - rule := inventoryTestRule("3306") - items, err := MergeInventory(InventoryMergeInput{ - Observed: []ObservedRule{ - inventoryObservedRule(t, rule, "", 1), - inventoryObservedRule(t, rule, "", 2), - }, - Desired: []DesiredRule{{UUID: "managed", Rule: rule, Origin: RuleOriginCreated}}, - }) - if err != nil { - t.Fatalf("merge inventory: %v", err) - } - if len(items) != 3 || items[0].State != InventoryStateExternal || items[1].State != InventoryStateExternal || items[2].State != InventoryStateDrifted || items[2].Match != InventoryMatchAmbiguous { - t.Fatalf("ambiguous candidates were guessed: %#v", items) - } -} - -func TestMergeInventoryUsesMarkerToReportChangedManagedRule(t *testing.T) { - desiredRule := inventoryTestRule("80") - changedRule := inventoryTestRule("8080") - marker := "onepanel:created:web" - items, err := MergeInventory(InventoryMergeInput{ - Observed: []ObservedRule{inventoryObservedRule(t, changedRule, marker, 1)}, - Desired: []DesiredRule{{ - UUID: "web", Rule: desiredRule, Origin: RuleOriginCreated, Marker: marker, - }}, - }) - if err != nil { - t.Fatalf("merge inventory: %v", err) - } - if len(items) != 1 || items[0].State != InventoryStateDrifted || items[0].Match != InventoryMatchChanged || items[0].Observed == nil || items[0].Desired == nil { - t.Fatalf("marker-owned semantic drift was not retained: %#v", items) - } -} - -func TestMergeInventoryMarkerSurvivesTransientLocatorChange(t *testing.T) { - rule := inventoryTestRule("443") - marker := "1panel-rule:https" - previous := inventoryObservedRule(t, rule, marker, 2) - previousKey, err := InstanceKey(previous) - if err != nil { - t.Fatalf("previous instance key: %v", err) - } - current := inventoryObservedRule(t, rule, marker, 7) - items, err := MergeInventory(InventoryMergeInput{ - Observed: []ObservedRule{current}, - Desired: []DesiredRule{{ - UUID: "https", Rule: rule, Origin: RuleOriginCreated, - Marker: marker, ObservedInstanceKey: previousKey, - }}, - }) - if err != nil { - t.Fatalf("merge locator drift: %v", err) - } - if len(items) != 1 || items[0].State != InventoryStateManaged || items[0].Match != InventoryMatchExact || - items[0].Observed == nil || items[0].Observed.Locator.Position == nil || *items[0].Observed.Locator.Position != 7 { - t.Fatalf("stable marker did not survive transient locator change: %#v", items) - } -} - -func TestMergeInventoryKeepsOpaqueRulesExternal(t *testing.T) { - opaque := ObservedRule{ - Rule: FirewallRule{Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput}}, - Locator: Locator{Provider: ProviderFirewalld, ScopeKey: "firewalld:public:input", NativeID: "service:ssh"}, - ParseStatus: ParseStatusOpaque, - Raw: "ssh service", - } - items, err := MergeInventory(InventoryMergeInput{Observed: []ObservedRule{opaque}}) - if err != nil { - t.Fatalf("merge opaque inventory: %v", err) - } - if len(items) != 1 || items[0].State != InventoryStateExternal || items[0].Match != InventoryMatchOpaque || items[0].Observed.Raw != opaque.Raw { - t.Fatalf("opaque rule was not preserved: %#v", items) - } -} - -func TestMergeInventoryUsesAdapterProtectedClassification(t *testing.T) { - rule := inventoryTestRule("22") - observed := inventoryObservedRule(t, rule, "", 1) - observed.Protected = true - - items, err := MergeInventory(InventoryMergeInput{Observed: []ObservedRule{observed}}) - if err != nil { - t.Fatalf("merge protected inventory: %v", err) - } - if len(items) != 1 || items[0].State != InventoryStateProtected || items[0].Observed == nil || !items[0].Observed.Protected { - t.Fatalf("adapter protected classification was lost: %#v", items) - } -} - -func TestMergeInventoryReportsRuntimePermanentDrift(t *testing.T) { - rule := inventoryTestRule("443") - observed := inventoryObservedRule(t, rule, "onepanel:created:https", 1) - observed.Persistence = PersistenceStatusPermanentOnly - items, err := MergeInventory(InventoryMergeInput{ - Observed: []ObservedRule{observed}, - Desired: []DesiredRule{{ - UUID: "https", Rule: rule, Origin: RuleOriginCreated, Marker: observed.Marker, - }}, - }) - if err != nil { - t.Fatalf("merge persistence drift: %v", err) - } - if len(items) != 1 || items[0].State != InventoryStateDrifted || items[0].Match != InventoryMatchExact { - t.Fatalf("runtime/permanent drift was hidden: %#v", items) - } -} - -func TestAttachRuntimeUsageDoesNotChangeOwnership(t *testing.T) { - rule := inventoryTestRule("8080") - items := []InventoryItem{{Rule: rule, State: InventoryStateAdopted, Match: InventoryMatchExact}} - usage := map[string]RuntimeUsage{ - RuntimeUsageKey(rule): {UsedBy: []string{"nginx", "", "nginx", "1panel"}, Reason: "listener"}, - } - - result := AttachRuntimeUsage(items, usage) - if result[0].State != InventoryStateAdopted || result[0].Usage == nil || !result[0].Usage.Used { - t.Fatalf("usage changed ownership or was not attached: %#v", result[0]) - } - if len(result[0].Usage.UsedBy) != 2 || result[0].Usage.UsedBy[0] != "1panel" || result[0].Usage.UsedBy[1] != "nginx" { - t.Fatalf("usage owners were not normalized: %#v", result[0].Usage) - } - if items[0].Usage != nil { - t.Fatal("AttachRuntimeUsage mutated its input") - } -} - -func TestAttachRuntimeUsageAggregatesPortRanges(t *testing.T) { - rule := inventoryTestRule("8000-8010") - items := []InventoryItem{{Rule: rule}} - usage := map[string]RuntimeUsage{ - "tcp\x008001": {Used: true, UsedBy: []string{"app"}, Reason: "application"}, - "tcp\x008009": {Used: true, UsedBy: []string{"worker"}, Reason: "listener"}, - "udp\x008005": {Used: true, UsedBy: []string{"ignored"}}, - } - result := AttachRuntimeUsage(items, usage) - if result[0].Usage == nil || len(result[0].Usage.UsedBy) != 2 || result[0].Usage.Reason != "multiple" { - t.Fatalf("range usage was not aggregated: %#v", result[0].Usage) - } -} - -func TestMergeInventoryRejectsStoredRuleKeyMismatch(t *testing.T) { - _, err := MergeInventory(InventoryMergeInput{Desired: []DesiredRule{{ - UUID: "broken", Rule: inventoryTestRule("53"), RuleKey: "sha256:stale", Origin: RuleOriginCreated, - }}}) - if err == nil { - t.Fatal("expected stored rule key mismatch to fail") - } -} - -func inventoryTestRule(port string) FirewallRule { - return FirewallRule{ - Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput}, - NativeKind: NativeKindRule, - Protocol: "tcp", - DestinationPort: port, - Action: ActionAccept, - } -} - -func inventoryObservedRule(t *testing.T, rule FirewallRule, marker string, position int) ObservedRule { - t.Helper() - return ObservedRule{ - Rule: rule, - Locator: Locator{ - Provider: rule.Scope.Provider, - ScopeKey: rule.Scope.Key(), - Position: &position, - }, - Marker: marker, - ParseStatus: ParseStatusSupported, - } -} diff --git a/agent/utils/firewall/filter/model.go b/agent/utils/firewall/filter/model.go index 9ff3d67e3324..f4b8dfc6e585 100644 --- a/agent/utils/firewall/filter/model.go +++ b/agent/utils/firewall/filter/model.go @@ -125,6 +125,35 @@ type Scope struct { Direction Direction `json:"direction"` } +func ManagedInputScopes(provider Provider) []Scope { + base := Scope{Provider: provider, Direction: DirectionInput} + switch provider { + case ProviderIptables, ProviderNftables: + result := make([]Scope, 0, 6) + for _, family := range []Family{FamilyIPv4, FamilyIPv6} { + for _, chain := range []string{BasicBeforeChain, IptablesInputChain, BasicAfterChain} { + scope := base + scope.Family, scope.Table, scope.Chain = family, "filter", chain + result = append(result, scope) + } + } + return result + case ProviderFirewalld: + base.Family, base.Zone = FamilyInet, FirewalldInputZone + return []Scope{base} + case ProviderUFW: + result := make([]Scope, 0, 2) + for _, family := range []Family{FamilyIPv4, FamilyIPv6} { + scope := base + scope.Family, scope.Chain = family, UFWInputChain + result = append(result, scope) + } + return result + default: + return nil + } +} + func (s Scope) Normalize() Scope { s.Provider = Provider(strings.ToLower(strings.TrimSpace(string(s.Provider)))) s.Family = Family(strings.ToLower(strings.TrimSpace(string(s.Family)))) diff --git a/agent/utils/firewall/filter/model_test.go b/agent/utils/firewall/filter/model_test.go deleted file mode 100644 index 8c348c6e8a72..000000000000 --- a/agent/utils/firewall/filter/model_test.go +++ /dev/null @@ -1,85 +0,0 @@ -package filter - -import ( - "errors" - "testing" -) - -func TestScopeValidateMVP(t *testing.T) { - tests := []struct { - name string - scope Scope - key string - wantErr error - }{ - { - name: "iptables basic chain", - scope: Scope{Provider: "IPTABLES", Family: FamilyIPv4, Table: "FILTER", Chain: "1panel_basic", Direction: DirectionInput}, - key: "iptables:ipv4:filter:1PANEL_BASIC:input", - }, - { - name: "firewalld public", - scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "PUBLIC", Direction: DirectionInput}, - key: "firewalld:public:input", - }, - { - name: "ufw incoming default chain", - scope: Scope{Provider: ProviderUFW, Family: FamilyIPv6, Direction: DirectionInput}, - key: "ufw:incoming:ipv6", - }, - { - name: "firewalld private unsupported", - scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "private", Direction: DirectionInput}, - wantErr: ErrUnsupportedScope, - }, - { - name: "iptables external chain unsupported", - scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "DOCKER", Direction: DirectionInput}, - wantErr: ErrUnsupportedScope, - }, - { - name: "ufw output unsupported", - scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: Direction("output")}, - wantErr: ErrInvalidScope, - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - err := test.scope.ValidateMVP() - if test.wantErr != nil { - if !errors.Is(err, test.wantErr) { - t.Fatalf("expected %v, got %v", test.wantErr, err) - } - return - } - if err != nil { - t.Fatalf("validate scope: %v", err) - } - if got := test.scope.Key(); got != test.key { - t.Fatalf("expected key %q, got %q", test.key, got) - } - }) - } -} - -func TestCapabilitiesSupportsMVPScope(t *testing.T) { - capabilities := Capabilities{Scopes: MVPScopePatterns()} - if !capabilities.SupportsScope(Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput}) { - t.Fatal("expected UFW incoming IPv4 scope to be supported") - } - if capabilities.SupportsScope(Scope{Provider: ProviderUFW, Family: FamilyIPv6, Chain: "outgoing", Direction: Direction("output")}) { - t.Fatal("did not expect UFW outgoing scope to be supported") - } - if capabilities.SupportsScope(Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "private", Direction: DirectionInput}) { - t.Fatal("did not expect firewalld private zone to be supported") - } -} - -func TestFirewalldFamiliesSharePublicExecutionScope(t *testing.T) { - ipv4 := Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: "public", Direction: DirectionInput} - ipv6 := Scope{Provider: ProviderFirewalld, Family: FamilyIPv6, Zone: "public", Direction: DirectionInput} - if ipv4.Key() != ipv6.Key() || ipv4.Key() != "firewalld:public:input" { - t.Fatalf("firewalld public pipeline was split by family: ipv4=%q ipv6=%q", ipv4.Key(), ipv6.Key()) - } -} diff --git a/agent/utils/firewall/filter/normalize_test.go b/agent/utils/firewall/filter/normalize_test.go deleted file mode 100644 index e18ac17fb845..000000000000 --- a/agent/utils/firewall/filter/normalize_test.go +++ /dev/null @@ -1,270 +0,0 @@ -package filter - -import ( - "errors" - "fmt" - "strings" - "testing" -) - -func TestNormalizeRule(t *testing.T) { - rule, err := NormalizeRule(FirewallRule{ - Scope: Scope{ - Provider: ProviderUFW, - Family: FamilyIPv4, - Direction: DirectionInput, - }, - Protocol: " TCP ", - SourceAddress: "172.16.10.111", - DestinationAddress: "0.0.0.0/0", - DestinationPort: "080:080", - Interface: " * ", - Action: "ALLOW", - ConnectionStates: []string{"NEW", "established", "new"}, - }) - if err != nil { - t.Fatalf("normalize rule: %v", err) - } - - if rule.Scope.Key() != "ufw:incoming:ipv4" { - t.Fatalf("unexpected scope key %q", rule.Scope.Key()) - } - if rule.NativeKind != NativeKindUFWRule { - t.Fatalf("unexpected native kind %q", rule.NativeKind) - } - if rule.Protocol != "tcp" || rule.Action != ActionAccept { - t.Fatalf("unexpected protocol/action: %s/%s", rule.Protocol, rule.Action) - } - if rule.SourceAddress != "172.16.10.111/32" || rule.DestinationAddress != "" { - t.Fatalf("unexpected normalized addresses: %q -> %q", rule.SourceAddress, rule.DestinationAddress) - } - if rule.DestinationPort != "80" { - t.Fatalf("unexpected destination port %q", rule.DestinationPort) - } - if rule.Interface != "" { - t.Fatalf("wildcard interface was not normalized: %q", rule.Interface) - } - if len(rule.ConnectionStates) != 2 || rule.ConnectionStates[0] != "established" || rule.ConnectionStates[1] != "new" { - t.Fatalf("unexpected states: %#v", rule.ConnectionStates) - } -} - -func TestNormalizeRuleRejectsCompositeAndFamilyMismatch(t *testing.T) { - base := FirewallRule{ - Scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput}, - Protocol: "tcp/udp", - DestinationPort: "80", - Action: ActionAccept, - } - if _, err := NormalizeRule(base); !errors.Is(err, ErrCompositeRule) { - t.Fatalf("expected composite rule error, got %v", err) - } - - base.Protocol = "tcp" - base.SourceAddress = "2001:db8::1" - if _, err := NormalizeRule(base); !errors.Is(err, ErrInvalidRule) { - t.Fatalf("expected family validation error, got %v", err) - } -} - -func TestNormalizeRuleRejectsInvalidConnectionState(t *testing.T) { - rule := FirewallRule{ - Scope: Scope{ - Provider: ProviderNftables, Family: FamilyIPv4, Table: "filter", - Chain: IptablesInputChain, Direction: DirectionInput, - }, - Protocol: "tcp", Action: ActionAccept, - ConnectionStates: []string{"established } accept\nflush ruleset"}, - } - if _, err := NormalizeRule(rule); !errors.Is(err, ErrInvalidRule) { - t.Fatalf("expected invalid connection state error, got %v", err) - } -} - -func TestNormalizeNativeDestinationPortSets(t *testing.T) { - for _, provider := range []Provider{ProviderIptables, ProviderUFW} { - scope := Scope{Provider: provider, Family: FamilyIPv4, Direction: DirectionInput} - if provider == ProviderIptables { - scope.Table = "filter" - scope.Chain = IptablesInputChain - } - rule, err := NormalizeRule(FirewallRule{ - Scope: scope, Protocol: "tcp", DestinationPort: "080,443,8080:8090,443", Action: ActionAccept, - }) - if err != nil { - t.Fatalf("normalize %s port set: %v", provider, err) - } - if rule.DestinationPort != "80,443,8080-8090" { - t.Fatalf("unexpected %s port set: %q", provider, rule.DestinationPort) - } - } - - _, err := NormalizeRule(FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: FirewalldInputZone, Direction: DirectionInput}, - Protocol: "tcp", DestinationPort: "80,443", Action: ActionAccept, - }) - if !errors.Is(err, ErrCompositeRule) { - t.Fatalf("expected firewalld port set expansion, got %v", err) - } - - tooMany := "1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16" - _, err = NormalizeRule(FirewallRule{ - Scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput}, - Protocol: "tcp", DestinationPort: tooMany, Action: ActionAccept, - }) - if !errors.Is(err, ErrInvalidRule) { - t.Fatalf("expected native port-set limit error, got %v", err) - } -} - -func TestNormalizeUFWAllowsAllProtocolsForDestinationPort(t *testing.T) { - rule, err := NormalizeRule(FirewallRule{ - Scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput}, - Protocol: "all", DestinationPort: "53", Action: ActionAccept, - }) - if err != nil { - t.Fatalf("normalize UFW all-protocol port rule: %v", err) - } - if rule.Protocol != "all" || rule.DestinationPort != "53" { - t.Fatalf("unexpected normalized rule: %#v", rule) - } - - for _, provider := range []Provider{ProviderIptables, ProviderFirewalld} { - testRule := rule - testRule.Scope.Provider = provider - switch provider { - case ProviderIptables: - testRule.Scope.Table = "filter" - testRule.Scope.Chain = IptablesInputChain - case ProviderFirewalld: - testRule.Scope.Family = FamilyInet - testRule.Scope.Zone = FirewalldInputZone - testRule.Scope.Chain = "" - } - if _, err = NormalizeRule(testRule); !errors.Is(err, ErrInvalidRule) { - t.Fatalf("expected %s all-protocol port rejection, got %v", provider, err) - } - } -} - -func TestExpandAtomicRules(t *testing.T) { - rules, err := ExpandAtomicRules(FirewallRule{ - Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput}, - Protocol: "tcp/udp", - SourceAddress: "172.16.10.111, 172.16.10.112", - DestinationPort: "80,443,80", - Action: ActionAccept, - }) - if err != nil { - t.Fatalf("expand atomic rules: %v", err) - } - if len(rules) != 4 { - t.Fatalf("expected 4 rules with native iptables port sets, got %d", len(rules)) - } - for _, rule := range rules { - if rule.Protocol == "tcp/udp" || rule.SourceAddress == "" || rule.DestinationPort != "80,443" { - t.Fatalf("rule was not expanded correctly: %#v", rule) - } - } -} - -func TestExpandAtomicRulesKeepsNativePortSetsAndSplitsFirewalld(t *testing.T) { - iptablesRules, err := ExpandAtomicRules(FirewallRule{ - Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput}, - Protocol: "tcp", DestinationPort: "80,443", Action: ActionAccept, - }) - if err != nil || len(iptablesRules) != 1 || iptablesRules[0].DestinationPort != "80,443" { - t.Fatalf("iptables port set was expanded: rules=%#v err=%v", iptablesRules, err) - } - firewalldRules, err := ExpandAtomicRules(FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: FirewalldInputZone, Direction: DirectionInput}, - Protocol: "tcp", DestinationPort: "80,443", Action: ActionDrop, - }) - if err != nil || len(firewalldRules) != 2 || firewalldRules[0].DestinationPort != "80" || firewalldRules[1].DestinationPort != "443" { - t.Fatalf("firewalld port set was not expanded: rules=%#v err=%v", firewalldRules, err) - } -} - -func TestExpandUFWInetByFamily(t *testing.T) { - rules, err := ExpandAtomicRules(FirewallRule{ - Scope: Scope{Provider: ProviderUFW, Family: FamilyInet, Direction: DirectionInput}, - Protocol: "tcp", - DestinationPort: "22", - Action: ActionAccept, - }) - if err != nil { - t.Fatalf("expand UFW families: %v", err) - } - if len(rules) != 2 || rules[0].Scope.Family != FamilyIPv4 || rules[1].Scope.Family != FamilyIPv6 { - t.Fatalf("unexpected UFW family expansion: %#v", rules) - } - - rules, err = ExpandAtomicRules(FirewallRule{ - Scope: Scope{Provider: ProviderUFW, Family: FamilyInet, Direction: DirectionInput}, - Protocol: "all", - SourceAddress: "172.16.10.111", - Action: ActionDrop, - }) - if err != nil { - t.Fatalf("expand family-specific UFW rule: %v", err) - } - if len(rules) != 1 || rules[0].Scope.Family != FamilyIPv4 { - t.Fatalf("expected one IPv4 rule, got %#v", rules) - } - - rules, err = ExpandAtomicRules(FirewallRule{ - Scope: Scope{Provider: ProviderUFW, Family: FamilyInet, Direction: DirectionInput}, - Protocol: "all", - SourceAddress: "::/0", - Action: ActionDrop, - }) - if err != nil { - t.Fatalf("expand IPv6-any UFW rule: %v", err) - } - if len(rules) != 1 || rules[0].Scope.Family != FamilyIPv6 { - t.Fatalf("expected one IPv6 rule, got %#v", rules) - } -} - -func TestNormalizeFirewalldNativeKindSetsExecutionBucket(t *testing.T) { - priority := -100 - rich, err := NormalizeRule(FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: "public", Direction: DirectionInput}, - NativeKind: NativeKindRichRule, Protocol: "tcp", DestinationPort: "3306", Action: ActionDrop, Priority: &priority, - }) - if err != nil { - t.Fatalf("normalize firewalld rich rule: %v", err) - } - if rich.OrderBucket != OrderBucketRichPre || rich.Priority == nil || *rich.Priority != -100 { - t.Fatalf("unexpected rich rule placement: %#v", rich) - } - zonePort, err := NormalizeRule(FirewallRule{ - Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput}, - NativeKind: NativeKindZonePort, Protocol: "tcp", DestinationPort: "3306", Action: ActionAccept, Priority: &priority, - }) - if err != nil { - t.Fatalf("normalize firewalld zone port: %v", err) - } - if zonePort.OrderBucket != OrderBucketZonePrimitiveAllow || zonePort.Priority != nil { - t.Fatalf("native port exposed fake priority: %#v", zonePort) - } -} - -func TestExpandAtomicRulesLimit(t *testing.T) { - var addresses strings.Builder - for i := 1; i <= 65; i++ { - if i > 1 { - addresses.WriteString(",") - } - addresses.WriteString(fmt.Sprintf("10.0.0.%d", i)) - } - _, err := ExpandAtomicRules(FirewallRule{ - Scope: Scope{Provider: ProviderUFW, Family: FamilyInet, Direction: DirectionInput}, - Protocol: "tcp/udp", - SourceAddress: addresses.String(), - Action: ActionAccept, - }) - if !errors.Is(err, ErrExpansionLimit) { - t.Fatalf("expected expansion limit error, got %v", err) - } -} diff --git a/agent/utils/firewall/filter/providers/firewalld/adapter_test.go b/agent/utils/firewall/filter/providers/firewalld/adapter_test.go deleted file mode 100644 index 9557dbf213d7..000000000000 --- a/agent/utils/firewall/filter/providers/firewalld/adapter_test.go +++ /dev/null @@ -1,579 +0,0 @@ -package firewalld - -import ( - "context" - "errors" - "reflect" - "sort" - "strings" - "sync" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -func TestPrepareRuleChoosesNativePortOrRichRule(t *testing.T) { - adapter := NewAdapterWithReader(newFakeCommandReader()) - zonePort, err := adapter.PrepareRule(filter.FirewallRule{ - Scope: testScope(filter.FamilyInet), Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, - }) - if err != nil { - t.Fatalf("prepare native port: %v", err) - } - if zonePort.NativeKind != filter.NativeKindZonePort || zonePort.OrderBucket != filter.OrderBucketZonePrimitiveAllow || zonePort.Priority != nil { - t.Fatalf("simple allow did not remain a native port: %#v", zonePort) - } - priority := -100 - rich, err := adapter.PrepareRule(filter.FirewallRule{ - Scope: testScope(filter.FamilyIPv4), Protocol: "tcp", SourceAddress: "172.16.10.111", DestinationPort: "443", - Action: filter.ActionDrop, Priority: &priority, - }) - if err != nil { - t.Fatalf("prepare rich rule: %v", err) - } - if rich.NativeKind != filter.NativeKindRichRule || rich.OrderBucket != filter.OrderBucketRichPre || rich.SourceAddress != "172.16.10.111/32" { - t.Fatalf("address deny did not become a rich rule: %#v", rich) - } -} - -func TestLegacyFirewalldRejectsOnlyExplicitRichRulePriority(t *testing.T) { - reader := newFakeCommandReader() - reader.outputs["--version"] = "0.6.3\n" - adapter := NewAdapterWithReader(reader) - priority := -100 - rule := filter.FirewallRule{ - Scope: testScope(filter.FamilyIPv4), NativeKind: filter.NativeKindRichRule, - Protocol: "tcp", DestinationPort: "443", Action: filter.ActionDrop, Priority: &priority, - } - if err := adapter.CheckRule(context.Background(), rule); !errors.Is(err, filter.ErrUnsupportedScope) { - t.Fatalf("legacy firewalld accepted explicit priority: %v", err) - } - priority = 0 - if err := adapter.CheckRule(context.Background(), rule); err != nil { - t.Fatalf("legacy firewalld rejected priority-zero compatibility rule: %v", err) - } - rule.UUID = "legacy-rich-rule" - snapshot, err := filter.NewSnapshot(testScope(filter.FamilyIPv4), nil) - if err != nil { - t.Fatalf("build legacy snapshot: %v", err) - } - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile legacy compatibility rule: %v", err) - } - if option := plan.Rules[0].Commands[0].Args[1]; strings.Contains(option, "priority=") { - t.Fatalf("legacy compatibility command included priority: %s", option) - } - capabilities, err := adapter.Capabilities(context.Background()) - if err != nil || capabilities.ExplicitPriority { - t.Fatalf("legacy firewalld exposed explicit priority: %#v err=%v", capabilities, err) - } -} - -func TestModernFirewalldSupportsExplicitRichRulePriority(t *testing.T) { - reader := newFakeCommandReader() - reader.outputs["--version"] = "1.3.4\n" - adapter := NewAdapterWithReader(reader) - priority := 100 - rule := filter.FirewallRule{ - Scope: testScope(filter.FamilyIPv4), NativeKind: filter.NativeKindRichRule, - Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, Priority: &priority, - } - if err := adapter.CheckRule(context.Background(), rule); err != nil { - t.Fatalf("modern firewalld rejected explicit priority: %v", err) - } - capabilities, err := adapter.Capabilities(context.Background()) - if err != nil || !capabilities.ExplicitPriority { - t.Fatalf("modern firewalld hid explicit priority: %#v err=%v", capabilities, err) - } -} - -func TestParseFirewalldVersion(t *testing.T) { - for _, test := range []struct { - version string - major, minor int - }{ - {version: "0.6.3", major: 0, minor: 6}, - {version: "0.7.0-1.el7", major: 0, minor: 7}, - {version: "2.1.0", major: 2, minor: 1}, - } { - major, minor, err := parseFirewalldVersion(test.version) - if err != nil || major != test.major || minor != test.minor { - t.Fatalf("parse %q: got %d.%d err=%v", test.version, major, minor, err) - } - } -} - -func TestCompileCreateUsesExplicitPublicRuntimeAndPermanentCommands(t *testing.T) { - adapter := NewAdapterWithReader(newFakeCommandReader()) - snapshot, _ := filter.NewSnapshot(testScope(filter.FamilyInet), nil) - rule := filter.FirewallRule{ - UUID: "https", Scope: testScope(filter.FamilyInet), Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, - } - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile native port: %v", err) - } - want := [][]string{ - {"--zone=public", "--add-port=443/tcp"}, - {"--permanent", "--zone=public", "--add-port=443/tcp"}, - } - if len(plan.Rules) != 1 || len(plan.Rules[0].Commands) != 2 || - !reflect.DeepEqual(plan.Rules[0].Commands[0].Args, want[0]) || !reflect.DeepEqual(plan.Rules[0].Commands[1].Args, want[1]) { - t.Fatalf("unexpected native port plan: %#v", plan) - } - if plan.Rules[0].Expected.Rule.NativeKind != filter.NativeKindZonePort || plan.Rules[0].Expected.Marker != "" { - t.Fatalf("firewalld plan invented marker or representation: %#v", plan.Rules[0].Expected) - } - - richSnapshot, _ := filter.NewSnapshot(testScope(filter.FamilyIPv4), nil) - priority := -100 - rich := filter.FirewallRule{ - UUID: "blocked-ip", Scope: testScope(filter.FamilyIPv4), Protocol: "tcp", SourceAddress: "172.16.10.111", - DestinationPort: "3306", Action: filter.ActionDrop, Priority: &priority, - } - plan, err = adapter.Compile(richSnapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rich}}) - if err != nil { - t.Fatalf("compile rich rule: %v", err) - } - option := `--add-rich-rule=rule family="ipv4" priority="-100" source address="172.16.10.111/32" port port="3306" protocol="tcp" drop` - if plan.Rules[0].Commands[0].Args[1] != option || plan.Rules[0].Expected.Rule.NativeKind != filter.NativeKindRichRule { - t.Fatalf("unexpected rich rule plan: %#v", plan.Rules[0]) - } -} - -func TestCompileAdoptsCanonicalExternalRuleWithoutSystemMutation(t *testing.T) { - reader := newFakeCommandReader() - reader.set(false, "--list-ports", "8080/tcp\n") - reader.set(true, "--list-ports", "8080/tcp\n") - adapter := NewAdapterWithReader(reader) - snapshot, err := adapter.Observe(context.Background(), testScope(filter.FamilyInet)) - if err != nil { - t.Fatalf("observe external port: %v", err) - } - rule := snapshot.Rules[0].Rule - rule.UUID = "adopted" - locator := snapshot.Rules[0].Locator - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeAdopt, After: &rule, Locator: &locator}}) - if err != nil { - t.Fatalf("compile adoption: %v", err) - } - if len(plan.Rules[0].Commands) != 0 || plan.Rules[0].Previous == nil || plan.Rules[0].Expected.Locator.Canonical != "port:8080/tcp" { - t.Fatalf("canonical adoption mutated firewalld: %#v", plan.Rules[0]) - } - if _, err := adapter.Apply(context.Background(), plan); err != nil { - t.Fatalf("no-op canonical adoption required a writer: %v", err) - } -} - -func TestCompileUpdateAndDeleteValidateManagedCanonicalTarget(t *testing.T) { - reader := newFakeCommandReader() - reader.set(false, "--list-ports", "8080/tcp\n") - reader.set(true, "--list-ports", "8080/tcp\n") - adapter := NewAdapterWithReader(reader) - snapshot, _ := adapter.Observe(context.Background(), testScope(filter.FamilyInet)) - before := snapshot.Rules[0].Rule - before.UUID = "owned" - after := before - after.DestinationPort = "9090" - locator := snapshot.Rules[0].Locator - - update, err := adapter.Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeUpdate, Before: &before, After: &after, Locator: &locator, - }}) - if err != nil { - t.Fatalf("compile update: %v", err) - } - if len(update.Rules[0].Commands) != 4 || len(update.Rules[0].RollbackCommands) != 4 || - update.Rules[0].Commands[0].Args[1] != "--remove-port=8080/tcp" || update.Rules[0].Commands[2].Args[1] != "--add-port=9090/tcp" { - t.Fatalf("unexpected update plan: %#v", update.Rules[0]) - } - deletePlan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeDelete, Before: &before, Locator: &locator}}) - if err != nil { - t.Fatalf("compile delete: %v", err) - } - if len(deletePlan.Rules[0].Commands) != 2 || deletePlan.Rules[0].Previous == nil || deletePlan.Rules[0].Commands[1].Args[2] != "--remove-port=8080/tcp" { - t.Fatalf("unexpected delete plan: %#v", deletePlan.Rules[0]) - } - verified, err := adapter.Verify(context.Background(), deletePlan) - if err != nil || verified.Matched { - t.Fatalf("delete verified while canonical target remained: result=%#v err=%v", verified, err) - } - reader.set(false, "--list-ports", "") - reader.set(true, "--list-ports", "") - verified, err = adapter.Verify(context.Background(), deletePlan) - if err != nil || !verified.Matched { - t.Fatalf("deleted canonical target did not verify: result=%#v err=%v", verified, err) - } - - protected := snapshot - protected.Rules = append([]filter.ObservedRule(nil), snapshot.Rules...) - protected.Rules[0].Protected = true - if _, err := adapter.Compile(protected, []filter.DesiredChange{{Operation: filter.ChangeDelete, Before: &before, Locator: &locator}}); !errors.Is(err, filter.ErrProtectedRule) { - t.Fatalf("expected protected delete rejection, got %v", err) - } -} - -func TestApplyCompensatesOnlySuccessfulFirewalldSteps(t *testing.T) { - reader := newFakeCommandReader() - writer := &fakeCommandWriter{failAt: 2, err: errors.New("permanent write failed")} - adapter := NewAdapterWithBackend(reader, writer) - snapshot, _ := adapter.Observe(context.Background(), testScope(filter.FamilyInet)) - rule := filter.FirewallRule{UUID: "web", Scope: testScope(filter.FamilyInet), Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - plan, _ := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - - if _, err := adapter.Apply(context.Background(), plan); err == nil || !strings.Contains(err.Error(), "permanent write failed") { - t.Fatalf("expected apply failure, got %v", err) - } - if len(writer.commands) != 3 || writer.commands[2].Args[1] != "--remove-port=80/tcp" { - t.Fatalf("runtime write was not compensated precisely: %#v", writer.commands) - } -} - -func TestRollbackReversesFullyAppliedFirewalldPlan(t *testing.T) { - reader := newFakeCommandReader() - writer := &fakeCommandWriter{} - adapter := NewAdapterWithBackend(reader, writer) - snapshot, _ := adapter.Observe(context.Background(), testScope(filter.FamilyInet)) - rule := filter.FirewallRule{UUID: "rollback", Scope: testScope(filter.FamilyInet), Protocol: "tcp", DestinationPort: "8080", Action: filter.ActionAccept} - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile rollback plan: %v", err) - } - if err := adapter.Rollback(context.Background(), plan); err != nil { - t.Fatalf("rollback applied plan: %v", err) - } - if len(writer.commands) != 2 || writer.commands[0].Args[0] != "--permanent" || - writer.commands[0].Args[2] != "--remove-port=8080/tcp" || writer.commands[1].Args[1] != "--remove-port=8080/tcp" { - t.Fatalf("unexpected rollback writes: %#v", writer.commands) - } -} - -func TestVerifyRequiresConvergedCanonicalRule(t *testing.T) { - reader := newFakeCommandReader() - adapter := NewAdapterWithReader(reader) - snapshot, _ := adapter.Observe(context.Background(), testScope(filter.FamilyInet)) - rule := filter.FirewallRule{UUID: "dns", Scope: testScope(filter.FamilyInet), Protocol: "udp", DestinationPort: "53", Action: filter.ActionAccept} - plan, _ := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - reader.set(false, "--list-ports", "53/udp\n") - - verified, err := adapter.Verify(context.Background(), plan) - if err != nil || verified.Matched { - t.Fatalf("runtime-only rule verified as converged: result=%#v err=%v", verified, err) - } - reader.set(true, "--list-ports", "53/udp\n") - verified, err = adapter.Verify(context.Background(), plan) - if err != nil || !verified.Matched { - t.Fatalf("converged rule did not verify: result=%#v err=%v", verified, err) - } -} - -func TestCompileRejectsBroadDenyAndPersistenceDrift(t *testing.T) { - adapter := NewAdapterWithReader(newFakeCommandReader()) - scope := testScope(filter.FamilyInet) - snapshot, _ := filter.NewSnapshot(scope, nil) - rule := filter.FirewallRule{UUID: "deny-all", Scope: scope, Protocol: "all", Action: filter.ActionDrop} - if _, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}); !errors.Is(err, filter.ErrLockoutRisk) { - t.Fatalf("expected broad deny rejection, got %v", err) - } - port := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindZonePort, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - observed := observedForRule(port) - observed.Persistence = filter.PersistenceStatusRuntimeOnly - drifted, _ := filter.NewSnapshot(scope, []filter.ObservedRule{observed}) - create := filter.FirewallRule{UUID: "https", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept} - if _, err := adapter.Compile(drifted, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &create}}); !errors.Is(err, filter.ErrRuleStale) { - t.Fatalf("expected runtime/permanent drift rejection, got %v", err) - } -} - -func TestObservePublicInetMergesNativeObjectsAndReportsScopeNotices(t *testing.T) { - reader := newFakeCommandReader() - reader.set(false, "--list-ports", "22/tcp 53/udp 123/sctp\n") - reader.set(true, "--list-ports", "22/tcp 80/tcp 123/sctp\n") - reader.set(false, "--list-rich-rules", `rule port port="8080" protocol="tcp" accept`+"\n") - reader.set(true, "--list-rich-rules", `rule port port="8080" protocol="tcp" accept`+"\n") - reader.set(false, "--list-services", "ssh dhcpv6-client\n") - reader.set(true, "--list-services", "dhcpv6-client ssh\n") - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), testScope(filter.FamilyInet)) - if err != nil { - t.Fatalf("observe public inet: %v", err) - } - if len(snapshot.Rules) != 7 { - t.Fatalf("unexpected firewalld object count: %#v", snapshot.Rules) - } - assertPresence(t, snapshot.Rules, "port:22/tcp", filter.PersistenceStatusConverged) - assertPresence(t, snapshot.Rules, "port:53/udp", filter.PersistenceStatusRuntimeOnly) - assertPresence(t, snapshot.Rules, "port:80/tcp", filter.PersistenceStatusPermanentOnly) - sctp := findObserved(snapshot.Rules, "port:123/sctp") - if sctp == nil || sctp.ParseStatus != filter.ParseStatusOpaque { - t.Fatalf("unsupported native port was guessed: %#v", sctp) - } - rich := findObserved(snapshot.Rules, `rich:rule port port="8080" protocol="tcp" accept`) - if rich == nil || rich.ParseStatus != filter.ParseStatusSupported || rich.Rule.NativeKind != filter.NativeKindRichRule || rich.Rule.OrderBucket != filter.OrderBucketRichZeroAllow { - t.Fatalf("family-neutral rich rule was not normalized: %#v", rich) - } - service := findObserved(snapshot.Rules, "service:ssh") - if service == nil || service.ParseStatus != filter.ParseStatusOpaque || service.Rule.NativeKind != filter.NativeKindZoneService || - service.Rule.Protocol != "" || service.Rule.DestinationPort != "" || service.Rule.Description != "ssh" || service.Raw != "ssh" { - t.Fatalf("service object was not preserved as opaque: %#v", service) - } - dhcpv6Client := findObserved(snapshot.Rules, "service:dhcpv6-client") - if dhcpv6Client == nil || dhcpv6Client.Rule.Protocol != "" || dhcpv6Client.Rule.DestinationPort != "" || - dhcpv6Client.Rule.Description != "dhcpv6-client" || dhcpv6Client.Raw != "dhcpv6-client" { - t.Fatalf("dhcpv6 service was exposed as an allow-all rule: %#v", dhcpv6Client) - } - if !hasNotice(snapshot.Notices, filter.ScopeNoticeRuntimePermanentMismatch, "ports") { - t.Fatalf("scope notices missing: %#v", snapshot.Notices) - } -} - -func TestObserveReadsRuntimeAndPermanentWithListAll(t *testing.T) { - reader := newFakeCommandReader() - reader.set(false, "--list-ports", "22/tcp\n") - reader.set(true, "--list-ports", "22/tcp 80/tcp\n") - if _, err := NewAdapterWithReader(reader).Observe(context.Background(), testScope(filter.FamilyInet)); err != nil { - t.Fatalf("observe public zone: %v", err) - } - want := []string{ - "--zone=public\x00--list-all", - "--permanent\x00--zone=public\x00--list-all", - } - got := reader.readCalls() - sort.Strings(got) - sort.Strings(want) - if !reflect.DeepEqual(got, want) { - t.Fatalf("unexpected firewall-cmd reads: got %#v, want %#v", got, want) - } -} - -func TestParseZoneOutput(t *testing.T) { - output := `public (active) - target: default - interfaces: eth0 - services: cockpit ssh - ports: 22/tcp 53/udp - protocols: - forward: yes - masquerade: no - forward-ports: - source-ports: - icmp-blocks: - rich rules: - rule family="ipv4" source address="10.0.0.1" accept - rule port port="8080" protocol="tcp" accept` - got := parseZoneOutput(output) - want := zoneOutput{ - ports: "22/tcp 53/udp", - services: "cockpit ssh", - active: true, - rich: strings.Join([]string{ - `rule family="ipv4" source address="10.0.0.1" accept`, - `rule port port="8080" protocol="tcp" accept`, - }, "\n"), - } - if !reflect.DeepEqual(got, want) { - t.Fatalf("unexpected parsed zone output: got %#v, want %#v", got, want) - } -} - -func TestObservePublicPipelineKeepsRuleFamiliesInOneZoneScope(t *testing.T) { - reader := newFakeCommandReader() - rich := strings.Join([]string{ - `rule family="ipv4" priority="-100" source address="172.16.10.111" port port="3306" protocol="tcp" drop`, - `rule family="ipv4" destination address="10.0.0.1" reject`, - `rule family="ipv4" source address="10.0.0.0/8" protocol value="tcp" accept`, - `rule family="ipv4" log prefix="audit" accept`, - `rule family="ipv6" source address="2001:db8::1" accept`, - }, "\n") + "\n" - reader.set(false, "--list-rich-rules", rich) - reader.set(true, "--list-rich-rules", rich) - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), testScope(filter.FamilyIPv4)) - if err != nil { - t.Fatalf("observe public IPv4: %v", err) - } - if snapshot.Scope.Family != filter.FamilyInet || len(snapshot.Rules) != 5 { - t.Fatalf("unexpected public pipeline: %#v", snapshot) - } - first := snapshot.Rules[0] - if first.ParseStatus != filter.ParseStatusSupported || first.Rule.SourceAddress != "172.16.10.111/32" || first.Rule.DestinationPort != "3306" || - first.Rule.Priority == nil || *first.Rule.Priority != -100 || first.Rule.OrderBucket != filter.OrderBucketRichPre { - t.Fatalf("priority rich rule was not normalized: %#v", first) - } - opaque := findObserved(snapshot.Rules, `rich:rule family="ipv4" log prefix="audit" accept`) - if opaque == nil || opaque.ParseStatus != filter.ParseStatusOpaque || opaque.Raw == "" { - t.Fatalf("unsupported rich rule was not kept opaque: %#v", opaque) - } - protocolRule := findObserved(snapshot.Rules, `rich:rule family="ipv4" source address="10.0.0.0/8" protocol value="tcp" accept`) - if protocolRule == nil || protocolRule.ParseStatus != filter.ParseStatusSupported || protocolRule.Rule.Protocol != "tcp" || protocolRule.Rule.DestinationPort != "" { - t.Fatalf("protocol-only rich rule was not parsed: %#v", protocolRule) - } - ipv6 := findObserved(snapshot.Rules, `rich:rule family="ipv6" source address="2001:db8::1/128" accept`) - if ipv6 == nil || ipv6.Rule.Scope.Family != filter.FamilyIPv6 { - t.Fatalf("IPv6 rich rule was not retained in the public execution scope: %#v", ipv6) - } -} - -func TestNativeDetailRunsDedicatedFirewalldCommand(t *testing.T) { - reader := newFakeCommandReader() - info := "ssh\n ports: 22/tcp\n protocols:\n source-ports:\n helpers:\n destination:" - reader.setServiceInfo(false, "ssh", info) - reader.setServiceInfo(true, "ssh", info) - adapter := NewAdapterWithReader(reader) - - runtimeInfo, err := adapter.NativeDetail(context.Background(), "ssh", false) - if err != nil { - t.Fatalf("read runtime service info: %v", err) - } - permanentInfo, err := adapter.NativeDetail(context.Background(), "ssh", true) - if err != nil { - t.Fatalf("read permanent service info: %v", err) - } - if runtimeInfo != info || permanentInfo != info { - t.Fatalf("service info output was not preserved: runtime=%q permanent=%q", runtimeInfo, permanentInfo) - } -} - -func TestNativeDetailRejectsInvalidServiceName(t *testing.T) { - adapter := NewAdapterWithReader(newFakeCommandReader()) - if _, err := adapter.NativeDetail(context.Background(), "../ssh", false); !errors.Is(err, filter.ErrInvalidRule) { - t.Fatalf("expected invalid service rejection, got %v", err) - } -} - -func TestObserveRejectsZoneOutsidePublic(t *testing.T) { - scope := testScope(filter.FamilyInet) - scope.Zone = "work" - _, err := NewAdapterWithReader(newFakeCommandReader()).Observe(context.Background(), scope) - if !errors.Is(err, filter.ErrUnsupportedScope) { - t.Fatalf("expected unsupported scope, got %v", err) - } -} - -func TestZoneNoticesReportInactivePublic(t *testing.T) { - notices := publicZoneNotices(zoneOutput{}, zoneOutput{}) - if !hasNotice(notices, filter.ScopeNoticeManagedScopeInactive, "") { - t.Fatalf("inactive public notices missing: %#v", notices) - } -} - -type fakeCommandReader struct { - mu sync.Mutex - outputs map[string]string - errors map[string]error - calls []string - zones map[bool]zoneOutput -} - -type fakeCommandWriter struct { - commands []filter.NativeCommand - failAt int - err error -} - -func (f *fakeCommandWriter) Run(_ context.Context, command filter.NativeCommand) error { - f.commands = append(f.commands, command) - if f.failAt > 0 && len(f.commands) == f.failAt { - return f.err - } - return nil -} - -func newFakeCommandReader() *fakeCommandReader { - return &fakeCommandReader{ - outputs: make(map[string]string), errors: make(map[string]error), zones: map[bool]zoneOutput{false: {active: true}}, - } -} - -func (f *fakeCommandReader) set(permanent bool, option, output string) { - zone := f.zones[permanent] - switch option { - case "--list-ports": - zone.ports = strings.TrimSpace(output) - case "--list-rich-rules": - zone.rich = strings.TrimSpace(output) - case "--list-services": - zone.services = strings.TrimSpace(output) - } - f.zones[permanent] = zone -} - -func (f *fakeCommandReader) setServiceInfo(permanent bool, service, output string) { - args := make([]string, 0, 2) - if permanent { - args = append(args, "--permanent") - } - args = append(args, "--info-service="+service) - f.outputs[strings.Join(args, "\x00")] = output -} - -func (f *fakeCommandReader) Read(_ context.Context, args ...string) (string, error) { - key := strings.Join(args, "\x00") - f.mu.Lock() - f.calls = append(f.calls, key) - f.mu.Unlock() - if err := f.errors[key]; err != nil { - return "", err - } - if len(args) >= 2 && args[len(args)-2] == "--zone=public" && args[len(args)-1] == "--list-all" { - permanent := args[0] == "--permanent" - zone := f.zones[permanent] - header := "public" - if zone.active { - header += " (active)" - } - output := header + "\n services: " + zone.services + "\n ports: " + zone.ports + "\n rich rules:" - if zone.rich != "" { - output += "\n " + strings.ReplaceAll(zone.rich, "\n", "\n ") - } - return output + "\n", nil - } - output, exists := f.outputs[key] - if !exists { - return "", errors.New("unexpected firewall-cmd call: " + strings.Join(args, " ")) - } - return output, nil -} - -func (f *fakeCommandReader) readCalls() []string { - f.mu.Lock() - defer f.mu.Unlock() - return append([]string(nil), f.calls...) -} - -func testScope(family filter.Family) filter.Scope { - return filter.Scope{Provider: filter.ProviderFirewalld, Family: family, Zone: "public", Direction: filter.DirectionInput} -} - -func findObserved(rules []filter.ObservedRule, canonical string) *filter.ObservedRule { - for index := range rules { - if rules[index].Locator.Canonical == canonical { - return &rules[index] - } - } - return nil -} - -func assertPresence(t *testing.T, rules []filter.ObservedRule, canonical string, expected filter.PersistenceStatus) { - t.Helper() - rule := findObserved(rules, canonical) - if rule == nil || rule.Persistence != expected { - t.Fatalf("unexpected presence for %s: %#v", canonical, rule) - } -} - -func hasNotice(notices []filter.ScopeNotice, code filter.ScopeNoticeCode, value string) bool { - for _, notice := range notices { - if notice.Code != code { - continue - } - if value == "" { - return true - } - for _, candidate := range notice.Values { - if candidate == value { - return true - } - } - } - return false -} diff --git a/agent/utils/firewall/filter/providers/iptables/adapter_test.go b/agent/utils/firewall/filter/providers/iptables/adapter_test.go deleted file mode 100644 index 2a043af677bf..000000000000 --- a/agent/utils/firewall/filter/providers/iptables/adapter_test.go +++ /dev/null @@ -1,753 +0,0 @@ -package iptables - -import ( - "context" - "errors" - "fmt" - "reflect" - "strings" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -func TestCompileCreateAndAdoptCommands(t *testing.T) { - scope := testScope("1PANEL_BASIC") - snapshot, err := filter.NewSnapshot(scope, nil) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - rule := filter.FirewallRule{ - UUID: "rule-1", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", - SourceAddress: "10.0.0.0/8", DestinationPort: "443", Action: filter.ActionAccept, - } - adapter := NewAdapterWithReader(&fakeRuleReader{}) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile create: %v", err) - } - command := plan.Rules[0].Commands[0] - if len(plan.Rules) != 1 || command.Executable != "iptables-restore" || - !strings.Contains(command.Stdin, "-A 1PANEL_BASIC -p tcp -s 10.0.0.0/8 --dport 443 -m comment --comment 1panel-rule:rule-1 -j ACCEPT") || - plan.Rules[0].Expected.Marker != "1panel-rule:rule-1" { - t.Fatalf("unexpected create plan: %#v", plan) - } - - position := 1 - adoptSnapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{executorObserved(rule, position)}) - plan, err = adapter.Compile(adoptSnapshot, []filter.DesiredChange{{Operation: filter.ChangeAdopt, After: &rule, Locator: &filter.Locator{Position: &position}}}) - if err != nil { - t.Fatalf("compile adopt: %v", err) - } - if plan.Rules[0].Commands[0].Executable != "iptables-restore" || - !strings.Contains(plan.Rules[0].Commands[0].Stdin, "1panel-rule:rule-1") { - t.Fatalf("adoption did not replace the selected position: %#v", plan.Rules[0]) - } -} - -func TestIPv6ObserveCompileAndCapabilities(t *testing.T) { - scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6) - reader := &fakeRuleReader{output: "-A 1PANEL_BASIC -p ipv6-icmp -s 2001:db8::/64 -m comment --comment \"1panel-rule:ping6\" -j ACCEPT"} - adapter := NewAdapterWithReader(reader) - snapshot, err := adapter.Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe IPv6 rules: %v", err) - } - if len(snapshot.Rules) != 1 || snapshot.Rules[0].Rule.Protocol != "icmpv6" || - snapshot.Rules[0].Rule.SourceAddress != "2001:db8::/64" { - t.Fatalf("unexpected IPv6 snapshot: %#v", snapshot) - } - rule := filter.FirewallRule{ - UUID: "ping6", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "icmpv6", - SourceAddress: "2001:db8::/64", Action: filter.ActionAccept, - } - empty, _ := filter.NewSnapshot(scope, nil) - plan, err := adapter.Compile(empty, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile IPv6 rule: %v", err) - } - command := plan.Rules[0].Commands[0] - if command.Executable != "ip6tables-restore" || !strings.Contains(command.Stdin, "-p ipv6-icmp -s 2001:db8::/64") { - t.Fatalf("unexpected IPv6 command: %#v", command) - } - capabilities, err := adapter.Capabilities(context.Background()) - if err != nil || !capabilities.SupportsScope(scope) || !capabilities.SupportsScope(testScope("1PANEL_BASIC")) { - t.Fatalf("IPv4/IPv6 capabilities are incomplete: %#v err=%v", capabilities, err) - } -} - -func TestContainsChainDeclarationRejectsMissingManagedChain(t *testing.T) { - if !containsChainDeclaration("-N 1PANEL_BASIC\n-A 1PANEL_BASIC -j ACCEPT\n", "1PANEL_BASIC") { - t.Fatal("managed chain declaration was not detected") - } - if containsChainDeclaration("", "1PANEL_BASIC") || - containsChainDeclaration("-N 1PANEL_BASIC_OTHER\n", "1PANEL_BASIC") { - t.Fatal("missing managed chain was treated as initialized") - } -} - -func TestObserveReportsMissingManagedChainWithoutHidingOtherFamilies(t *testing.T) { - scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6) - adapter := NewAdapterWithReader(&fakeRuleReader{err: fmt.Errorf("%w: IPv6 chain is missing", filter.ErrProviderUnavailable)}) - snapshot, err := adapter.Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe missing managed chain: %v", err) - } - if len(snapshot.Rules) != 0 || len(snapshot.Notices) != 1 || - snapshot.Notices[0].Code != filter.ScopeNoticeManagedScopeMissing { - t.Fatalf("unexpected missing-chain snapshot: %#v", snapshot) - } -} - -func TestCompileRejectsICMPFamilyMismatch(t *testing.T) { - for _, test := range []struct { - family filter.Family - protocol string - }{ - {family: filter.FamilyIPv4, protocol: "icmpv6"}, - {family: filter.FamilyIPv6, protocol: "icmp"}, - } { - scope := testScopeFamily("1PANEL_BASIC", test.family) - snapshot, _ := filter.NewSnapshot(scope, nil) - rule := filter.FirewallRule{ - UUID: "icmp", Scope: scope, NativeKind: filter.NativeKindRule, - Protocol: test.protocol, Action: filter.ActionAccept, - } - _, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if !errors.Is(err, filter.ErrInvalidRule) { - t.Fatalf("expected %s/%s to be rejected, got %v", test.family, test.protocol, err) - } - } -} - -func TestApplyRejectsCrossFamilyExecutable(t *testing.T) { - scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6) - snapshot, _ := filter.NewSnapshot(scope, nil) - rule := filter.FirewallRule{ - UUID: "web6", Scope: scope, NativeKind: filter.NativeKindRule, - Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, - } - writer := &fakeRuleWriter{} - adapter := NewAdapterWithBackend(&fakeRuleReader{}, writer) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile IPv6 rule: %v", err) - } - plan.Rules[0].Commands[0].Executable = "iptables" - if _, err := adapter.Apply(context.Background(), plan); !errors.Is(err, filter.ErrInvalidRule) { - t.Fatalf("expected cross-family command rejection, got %v", err) - } - if len(writer.commands) != 0 { - t.Fatalf("cross-family command reached the writer: %#v", writer.commands) - } -} - -func TestCompileRejectsBroadIPv4AndIPv6Deny(t *testing.T) { - for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} { - scope := testScopeFamily("1PANEL_BASIC", family) - snapshot, _ := filter.NewSnapshot(scope, nil) - rule := filter.FirewallRule{ - UUID: "deny-all", Scope: scope, NativeKind: filter.NativeKindRule, - Protocol: "all", Action: filter.ActionDrop, - } - _, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if !errors.Is(err, filter.ErrLockoutRisk) { - t.Fatalf("expected broad %s deny to be rejected, got %v", family, err) - } - } -} - -func TestCompileInsertsAllowBeforeTerminalDrop(t *testing.T) { - scope := testScope("1PANEL_BASIC_AFTER") - dropTCP := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", Action: filter.ActionDrop} - dropUDP := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "udp", Action: filter.ActionDrop} - snapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{executorObserved(dropTCP, 1), executorObserved(dropUDP, 2)}) - allow := filter.FirewallRule{UUID: "dns", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "udp", DestinationPort: "53", Action: filter.ActionAccept} - plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &allow}}) - if err != nil { - t.Fatalf("compile terminal insert: %v", err) - } - script := plan.Rules[0].Commands[0].Stdin - if strings.Index(script, "1panel-rule:dns") > strings.Index(script, "-p tcp -j DROP") { - t.Fatalf("allow rule was inserted after terminal drop: %#v", plan.Rules[0].Commands[0]) - } -} - -func TestCompileCreateUsesRequestedPosition(t *testing.T) { - scope := testScope("1PANEL_BASIC") - first := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - second := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "81", Action: filter.ActionAccept} - snapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{executorObserved(first, 1), executorObserved(second, 2)}) - position := int64(2) - rule := filter.FirewallRule{ - UUID: "inserted", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", - DestinationPort: "443", Action: filter.ActionAccept, OrderIndex: &position, - } - plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile positioned create: %v", err) - } - script := plan.Rules[0].Commands[0].Stdin - if strings.Index(script, "1panel-rule:inserted") < strings.Index(script, "--dport 80") || - strings.Index(script, "1panel-rule:inserted") > strings.Index(script, "--dport 81") { - t.Fatalf("create ignored requested position: %#v", plan.Rules[0].Commands[0]) - } -} - -func TestMultiportCheckCompileAndObserve(t *testing.T) { - scope := testScope("1PANEL_BASIC") - reader := &fakeRuleReader{} - adapter := NewAdapterWithReader(reader) - rule := filter.FirewallRule{ - UUID: "web", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", - DestinationPort: "80,443,8080-8090", Action: filter.ActionAccept, - } - if err := adapter.CheckRule(context.Background(), rule); err != nil { - t.Fatalf("check multiport: %v", err) - } - if !reflect.DeepEqual(reader.checks, []filter.Family{filter.FamilyIPv4}) { - t.Fatalf("unexpected multiport checks: %#v", reader.checks) - } - if err := adapter.CheckRule(context.Background(), rule); err != nil || len(reader.checks) != 1 { - t.Fatalf("successful multiport capability was not cached: checks=%#v err=%v", reader.checks, err) - } - snapshot, _ := filter.NewSnapshot(scope, nil) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile multiport: %v", err) - } - script := plan.Rules[0].Commands[0].Stdin - if !strings.Contains(script, "-m multiport --dports 80,443,8080:8090") { - t.Fatalf("unexpected multiport command: %s", script) - } - - reader.output = `-A 1PANEL_BASIC -p tcp -m multiport --dports 80,443,8080:8090 -j ACCEPT -m comment --comment "web"` - observed, err := adapter.Observe(context.Background(), scope) - if err != nil || len(observed.Rules) != 1 || observed.Rules[0].ParseStatus != filter.ParseStatusSupported || - observed.Rules[0].Rule.DestinationPort != "80,443,8080-8090" { - t.Fatalf("multiport was not observed: snapshot=%#v err=%v", observed, err) - } - - failingReader := &fakeRuleReader{multiportErr: errors.New("extension missing")} - if err := NewAdapterWithReader(failingReader).CheckRule(context.Background(), rule); !errors.Is(err, filter.ErrUnsupportedScope) { - t.Fatalf("expected unavailable multiport error, got %v", err) - } -} - -func TestCompileRejectsPositionFromOtherFamilySnapshot(t *testing.T) { - snapshotScope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv4) - snapshot, _ := filter.NewSnapshot(snapshotScope, nil) - position := int64(1) - rule := filter.FirewallRule{ - UUID: "ipv6-rule", - Scope: testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6), NativeKind: filter.NativeKindRule, - Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, OrderIndex: &position, - } - _, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeCreate, - After: &rule, - }}) - if !errors.Is(err, filter.ErrUnsupportedScope) { - t.Fatalf("expected cross-family position to be rejected, got %v", err) - } -} - -func TestCompileReordersManagedRuleWithinChain(t *testing.T) { - scope := testScope("1PANEL_BASIC") - first := filter.FirewallRule{UUID: "first", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - second := filter.FirewallRule{UUID: "second", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "81", Action: filter.ActionAccept} - third := filter.FirewallRule{UUID: "third", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "82", Action: filter.ActionAccept} - rules := []filter.ObservedRule{executorObserved(first, 1), executorObserved(second, 2), executorObserved(third, 3)} - for index := range rules { - rules[index].Marker = "1panel-rule:" + rules[index].Rule.UUID - } - snapshot, _ := filter.NewSnapshot(scope, rules) - target := int64(3) - after := first - after.OrderIndex = &target - plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeReorder, Before: &first, After: &after, Locator: &rules[0].Locator, - }}) - if err != nil { - t.Fatalf("compile reorder: %v", err) - } - rulePlan := plan.Rules[0] - if len(rulePlan.Commands) != 1 || rulePlan.Commands[0].Executable != "iptables-restore" || - strings.Index(rulePlan.Commands[0].Stdin, "1panel-rule:first") < strings.Index(rulePlan.Commands[0].Stdin, "1panel-rule:third") || - rulePlan.Expected.Locator.Position == nil || *rulePlan.Expected.Locator.Position != 3 { - t.Fatalf("unexpected reorder plan: %#v", rulePlan) - } -} - -func TestCompileUpdateMovesAndChangesManagedRule(t *testing.T) { - scope := testScope("1PANEL_BASIC") - first := filter.FirewallRule{UUID: "first", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - second := filter.FirewallRule{UUID: "second", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "81", Action: filter.ActionAccept} - third := filter.FirewallRule{UUID: "third", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "82", Action: filter.ActionAccept} - rules := []filter.ObservedRule{executorObserved(first, 1), executorObserved(second, 2), executorObserved(third, 3)} - for index := range rules { - rules[index].Marker = "1panel-rule:" + rules[index].Rule.UUID - } - snapshot, _ := filter.NewSnapshot(scope, rules) - target := int64(3) - after := first - after.DestinationPort = "443" - after.OrderIndex = &target - plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeUpdate, Before: &first, After: &after, Locator: &rules[0].Locator, - }}) - if err != nil { - t.Fatalf("compile positioned update: %v", err) - } - rulePlan := plan.Rules[0] - if len(rulePlan.Commands) != 1 || rulePlan.Commands[0].Executable != "iptables-restore" || - !strings.Contains(rulePlan.Commands[0].Stdin, "--dport 443") || rulePlan.Expected.Locator.Position == nil || - *rulePlan.Expected.Locator.Position != 3 { - t.Fatalf("unexpected positioned update plan: %#v", rulePlan) - } -} - -func TestCompileBlocksReorderAcrossExternalOrProtectedRule(t *testing.T) { - scope := testScope("1PANEL_BASIC") - first := filter.FirewallRule{UUID: "first", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - middle := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "81", Action: filter.ActionAccept} - last := filter.FirewallRule{UUID: "last", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "82", Action: filter.ActionAccept} - rules := []filter.ObservedRule{executorObserved(first, 1), executorObserved(middle, 2), executorObserved(last, 3)} - rules[0].Marker = "1panel-rule:first" - rules[2].Marker = "1panel-rule:last" - snapshot, _ := filter.NewSnapshot(scope, rules) - target := int64(3) - after := first - after.OrderIndex = &target - change := []filter.DesiredChange{{Operation: filter.ChangeReorder, Before: &first, After: &after, Locator: &rules[0].Locator}} - if _, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, change); !errors.Is(err, filter.ErrUnsupportedScope) { - t.Fatalf("expected external boundary rejection, got %v", err) - } - rules[1].Marker = "1panel-rule:middle" - rules[1].Protected = true - snapshot, _ = filter.NewSnapshot(scope, rules) - change[0].Locator = &snapshot.Rules[0].Locator - if _, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, change); !errors.Is(err, filter.ErrProtectedRule) { - t.Fatalf("expected protected boundary rejection, got %v", err) - } -} - -func TestApplyUsesCompiledSnapshotWithoutSecondRead(t *testing.T) { - scope := testScope("1PANEL_BASIC") - initial, _ := filter.NewSnapshot(scope, nil) - rule := filter.FirewallRule{UUID: "web", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - reader := &fakeRuleReader{} - writer := &fakeRuleWriter{} - adapter := NewAdapterWithBackend(reader, writer) - plan, err := adapter.Compile(initial, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile create: %v", err) - } - reader.output = "-A 1PANEL_BASIC -p tcp --dport 22 -j ACCEPT" - if _, err := adapter.Apply(context.Background(), plan); err != nil { - t.Fatalf("apply compiled plan: %v", err) - } - if len(writer.commands) != 1 || !writer.saved { - t.Fatalf("compiled plan was not applied: %#v", writer) - } -} - -func TestBatchCreateUsesSingleRestoreForOwnedChain(t *testing.T) { - scope := testScope("1PANEL_BASIC") - reader := &fakeRuleReader{output: `-A 1PANEL_BASIC -p tcp --dport 22 -m comment --comment external-ssh -j ACCEPT`} - adapter := NewAdapterWithReader(reader) - snapshot, err := adapter.Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe initial chain: %v", err) - } - first := filter.FirewallRule{ - UUID: "web", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", - DestinationPort: "80", Action: filter.ActionAccept, - } - second := filter.FirewallRule{ - UUID: "tls", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", - DestinationPort: "443", Action: filter.ActionAccept, - } - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{ - {Operation: filter.ChangeCreate, After: &first}, - {Operation: filter.ChangeCreate, After: &second}, - }) - if err != nil { - t.Fatalf("compile batch: %v", err) - } - if len(plan.Rules) != 2 || len(plan.Rules[0].Commands) != 1 || len(plan.Rules[1].Commands) != 0 { - t.Fatalf("batch did not compile to one restore command: %#v", plan) - } - command := plan.Rules[0].Commands[0] - if command.Executable != "iptables-restore" || !reflect.DeepEqual(command.Args, []string{"--noflush", "--wait"}) { - t.Fatalf("unexpected restore command: %#v", command) - } - for _, want := range []string{ - "*filter\n-F 1PANEL_BASIC\n", - "-A 1PANEL_BASIC -p tcp --dport 22 -m comment --comment external-ssh -j ACCEPT\n", - "--dport 80 -m comment --comment 1panel-rule:web -j ACCEPT\n", - "--dport 443 -m comment --comment 1panel-rule:tls -j ACCEPT\n", - "COMMIT\n", - } { - if !strings.Contains(command.Stdin, want) { - t.Fatalf("restore input does not contain %q:\n%s", want, command.Stdin) - } - } - if strings.Contains(command.Stdin, "-F INPUT") { - t.Fatalf("restore input flushes an unmanaged chain:\n%s", command.Stdin) - } - - writer := &fakeRuleWriter{} - adapter = NewAdapterWithBackend(reader, writer) - if _, err = adapter.Apply(context.Background(), plan); err != nil { - t.Fatalf("apply batch: %v", err) - } - if len(writer.commands) != 1 || writer.saveCalls != 1 { - t.Fatalf("batch was not restored and persisted once: %#v", writer) - } - if err = adapter.Rollback(context.Background(), plan); err != nil { - t.Fatalf("rollback batch: %v", err) - } - if len(writer.commands) != 2 || !strings.Contains(writer.commands[1].Stdin, "external-ssh") || - strings.Contains(writer.commands[1].Stdin, "1panel-rule:web") || writer.saveCalls != 2 { - t.Fatalf("rollback did not restore the original chain: %#v", writer) - } -} - -func TestIPv6BatchCreateUsesIPv6Restore(t *testing.T) { - scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6) - snapshot, _ := filter.NewSnapshot(scope, nil) - first := filter.FirewallRule{UUID: "one", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - second := filter.FirewallRule{UUID: "two", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept} - plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{ - {Operation: filter.ChangeCreate, After: &first}, - {Operation: filter.ChangeCreate, After: &second}, - }) - if err != nil { - t.Fatalf("compile IPv6 batch: %v", err) - } - if got := plan.Rules[0].Commands[0].Executable; got != "ip6tables-restore" { - t.Fatalf("IPv6 batch executable = %q", got) - } -} - -func TestBatchDeleteUsesSingleRestoreAndPreservesExternalRules(t *testing.T) { - scope := testScope("1PANEL_BASIC") - reader := &fakeRuleReader{output: strings.Join([]string{ - `-A 1PANEL_BASIC -p tcp --dport 80 -m comment --comment 1panel-rule:web -j ACCEPT`, - `-A 1PANEL_BASIC -p tcp --dport 22 -m comment --comment external-ssh -j ACCEPT`, - `-A 1PANEL_BASIC -p tcp --dport 443 -m comment --comment 1panel-rule:tls -j ACCEPT`, - }, "\n")} - adapter := NewAdapterWithReader(reader) - snapshot, err := adapter.Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe initial chain: %v", err) - } - web := snapshot.Rules[0].Rule - web.UUID = "web" - tls := snapshot.Rules[2].Rule - tls.UUID = "tls" - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{ - {Operation: filter.ChangeDelete, Before: &tls, Locator: &snapshot.Rules[2].Locator}, - {Operation: filter.ChangeDelete, Before: &web, Locator: &snapshot.Rules[0].Locator}, - }) - if err != nil { - t.Fatalf("compile batch delete: %v", err) - } - command := plan.Rules[0].Commands[0] - if command.Executable != "iptables-restore" || strings.Contains(command.Stdin, "1panel-rule:web") || - strings.Contains(command.Stdin, "1panel-rule:tls") || !strings.Contains(command.Stdin, "external-ssh") { - t.Fatalf("unexpected batch delete restore input:\n%s", command.Stdin) - } - rollback := plan.Rules[0].RollbackCommands[0].Stdin - if !strings.Contains(rollback, "1panel-rule:web") || !strings.Contains(rollback, "1panel-rule:tls") || - !strings.Contains(rollback, "external-ssh") { - t.Fatalf("batch delete rollback does not contain the original chain:\n%s", rollback) - } -} - -func TestApplyAndVerifyMarker(t *testing.T) { - scope := testScope("1PANEL_BASIC") - reader := &fakeRuleReader{} - writer := &fakeRuleWriter{} - adapter := NewAdapterWithBackend(reader, writer) - snapshot, _ := filter.NewSnapshot(scope, nil) - rule := filter.FirewallRule{UUID: "ssh", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "22", Action: filter.ActionAccept} - plan, _ := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if _, err := adapter.Apply(context.Background(), plan); err != nil { - t.Fatalf("apply plan: %v", err) - } - if len(writer.commands) != 1 || !writer.saved { - t.Fatalf("command was not executed and persisted: %#v", writer) - } - reader.output = `-A 1PANEL_BASIC -p tcp --dport 22 -m comment --comment "1panel-rule:ssh" -j ACCEPT` - verified, err := adapter.Verify(context.Background(), plan) - if err != nil || !verified.Matched { - t.Fatalf("verify marker: result=%#v err=%v", verified, err) - } -} - -func TestApplyCompensatesWhenPersistenceFails(t *testing.T) { - scope := testScope("1PANEL_BASIC") - reader := &fakeRuleReader{} - writer := &fakeRuleWriter{saveErrors: []error{errors.New("disk full"), nil}} - adapter := NewAdapterWithBackend(reader, writer) - snapshot, _ := filter.NewSnapshot(scope, nil) - rule := filter.FirewallRule{UUID: "ssh", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "22", Action: filter.ActionAccept} - plan, _ := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - - if _, err := adapter.Apply(context.Background(), plan); err == nil || err.Error() != "disk full" { - t.Fatalf("expected original persistence error, got %v", err) - } - if len(writer.commands) != 2 || writer.commands[0].Executable != "iptables-restore" || - writer.commands[1].Executable != "iptables-restore" || - !strings.Contains(writer.commands[0].Stdin, "1panel-rule:ssh") || - strings.Contains(writer.commands[1].Stdin, "1panel-rule:ssh") || writer.saveCalls != 2 { - t.Fatalf("failed write was not compensated and persisted: %#v", writer) - } -} - -func TestRollbackReversesFullyAppliedIptablesPlan(t *testing.T) { - scope := testScope("1PANEL_BASIC") - writer := &fakeRuleWriter{} - adapter := NewAdapterWithBackend(&fakeRuleReader{}, writer) - snapshot, _ := filter.NewSnapshot(scope, nil) - rule := filter.FirewallRule{UUID: "rollback", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "8080", Action: filter.ActionAccept} - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile rollback plan: %v", err) - } - if err := adapter.Rollback(context.Background(), plan); err != nil { - t.Fatalf("rollback applied plan: %v", err) - } - if len(writer.commands) != 1 || writer.commands[0].Executable != "iptables-restore" || - strings.Contains(writer.commands[0].Stdin, "1panel-rule:rollback") || writer.saveCalls != 1 { - t.Fatalf("unexpected rollback writes: %#v", writer) - } -} - -func TestApplyCompensatesFailedBatchReorder(t *testing.T) { - scope := testScope("1PANEL_BASIC") - reader := &fakeRuleReader{output: "-A 1PANEL_BASIC -p tcp --dport 80 -m comment --comment 1panel-rule:first -j ACCEPT\n" + - "-A 1PANEL_BASIC -p tcp --dport 81 -m comment --comment 1panel-rule:second -j ACCEPT"} - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe reorder snapshot: %v", err) - } - first := snapshot.Rules[0].Rule - first.UUID = "first" - target := int64(2) - after := first - after.OrderIndex = &target - plan, err := NewAdapterWithReader(reader).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeReorder, Before: &first, After: &after, Locator: &snapshot.Rules[0].Locator, - }}) - if err != nil { - t.Fatalf("compile reorder: %v", err) - } - writer := &fakeRuleWriter{runErrors: []error{errors.New("restore failed"), nil}} - adapter := NewAdapterWithBackend(reader, writer) - if _, err := adapter.Apply(context.Background(), plan); err == nil || err.Error() != "restore failed" { - t.Fatalf("expected reorder restore failure, got %v", err) - } - if len(writer.commands) != 1 || writer.commands[0].Executable != "iptables-restore" || writer.saveCalls != 1 { - t.Fatalf("failed batch reorder issued partial commands: %#v", writer) - } -} - -func TestVerifyDeleteRequiresMarkerToDisappear(t *testing.T) { - scope := testScope("1PANEL_BASIC") - rule := filter.FirewallRule{UUID: "ssh", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "22", Action: filter.ActionAccept} - position := 1 - observed := executorObserved(rule, position) - observed.Marker = "1panel-rule:ssh" - snapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{observed}) - plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{ - {Operation: filter.ChangeDelete, Before: &rule, Locator: &observed.Locator}, - }) - if err != nil { - t.Fatalf("compile delete: %v", err) - } - reader := &fakeRuleReader{output: `-A 1PANEL_BASIC -p tcp --dport 2200 -m comment --comment "1panel-rule:ssh" -j ACCEPT`} - verified, err := NewAdapterWithReader(reader).Verify(context.Background(), plan) - if err != nil || verified.Matched { - t.Fatalf("delete verification ignored a surviving marker: result=%#v err=%v", verified, err) - } -} - -func TestCompileRejectsProtectedMutation(t *testing.T) { - scope := testScope("1PANEL_BASIC_BEFORE") - rule := filter.FirewallRule{UUID: "loopback", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "all", Interface: "lo", Action: filter.ActionAccept} - position := 1 - observed := executorObserved(rule, position) - observed.Protected = true - snapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{observed}) - - _, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{ - {Operation: filter.ChangeDelete, Before: &rule, Locator: &observed.Locator}, - }) - if !errors.Is(err, filter.ErrProtectedRule) { - t.Fatalf("expected protected mutation rejection, got %v", err) - } -} - -func TestObserveMergesIPPortAndCombinedRulesInNativeOrder(t *testing.T) { - scope := testScope("1PANEL_BASIC") - reader := &fakeRuleReader{output: `-N 1PANEL_BASIC --A 1PANEL_BASIC -s 172.16.10.111/32 -j DROP --A 1PANEL_BASIC -p tcp -m tcp --dport 22 -j ACCEPT -m comment --comment "ssh" --A 1PANEL_BASIC -p tcp -m tcp -s 10.0.0.0/8 --dport 443 -j ACCEPT -m comment --comment "1panel-rule:managed" -`} - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe rules: %v", err) - } - if len(snapshot.Rules) != 3 || snapshot.Revision == "" { - t.Fatalf("unexpected snapshot: %#v", snapshot) - } - if snapshot.Rules[0].Rule.SourceAddress != "172.16.10.111/32" || snapshot.Rules[0].Rule.DestinationPort != "" { - t.Fatalf("IP rule was not preserved: %#v", snapshot.Rules[0]) - } - if snapshot.Rules[1].Rule.DestinationPort != "22" || snapshot.Rules[1].Rule.Description != "ssh" { - t.Fatalf("port rule was not preserved: %#v", snapshot.Rules[1]) - } - if snapshot.Rules[2].Rule.SourceAddress != "10.0.0.0/8" || snapshot.Rules[2].Rule.DestinationPort != "443" || snapshot.Rules[2].Marker != "1panel-rule:managed" { - t.Fatalf("combined managed rule was not preserved: %#v", snapshot.Rules[2]) - } - for index, observed := range snapshot.Rules { - if observed.Locator.Position == nil || *observed.Locator.Position != index+1 { - t.Fatalf("native position was not preserved: %#v", observed.Locator) - } - } -} - -func TestObserveKeepsUnsupportedRulesOpaque(t *testing.T) { - scope := testScope("1PANEL_BASIC_BEFORE") - reader := &fakeRuleReader{output: `-A 1PANEL_BASIC_BEFORE -m limit --limit 5/min -j ACCEPT --A 1PANEL_BASIC_BEFORE -p tcp -m multiport --dports 80,443 -j ACCEPT -`} - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe opaque rules: %v", err) - } - if len(snapshot.Rules) != 2 || snapshot.Rules[0].ParseStatus != filter.ParseStatusOpaque || snapshot.Rules[1].ParseStatus != filter.ParseStatusSupported || - snapshot.Rules[1].Rule.DestinationPort != "80,443" { - t.Fatalf("unsupported rules were guessed: %#v", snapshot.Rules) - } - if snapshot.Rules[0].Raw == "" || snapshot.Rules[0].Locator.Canonical == "" { - t.Fatalf("opaque diagnostics were not retained: %#v", snapshot.Rules[0]) - } -} - -func TestObserveAcceptsDefaultIPv6RejectRepresentation(t *testing.T) { - scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6) - reader := &fakeRuleReader{output: "-A 1PANEL_BASIC -s 2001:db8::/64 -j REJECT --reject-with icmp6-port-unreachable"} - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe IPv6 reject: %v", err) - } - if len(snapshot.Rules) != 1 || snapshot.Rules[0].ParseStatus != filter.ParseStatusSupported || - snapshot.Rules[0].Rule.Action != filter.ActionReject { - t.Fatalf("default IPv6 reject became opaque: %#v", snapshot.Rules) - } - reader.output = "-A 1PANEL_BASIC -s 2001:db8::/64 -j REJECT --reject-with icmp6-adm-prohibited" - snapshot, err = NewAdapterWithReader(reader).Observe(context.Background(), scope) - if err != nil || len(snapshot.Rules) != 1 || snapshot.Rules[0].ParseStatus != filter.ParseStatusOpaque { - t.Fatalf("non-default IPv6 reject semantics were guessed: snapshot=%#v err=%v", snapshot, err) - } -} - -func TestObserveNormalizesConnectionStateSafetyRule(t *testing.T) { - scope := testScope("1PANEL_BASIC_BEFORE") - reader := &fakeRuleReader{output: `-A 1PANEL_BASIC_BEFORE -m conntrack --ctstate RELATED,ESTABLISHED -j ACCEPT -m comment --comment "ESTABLISHED Whitelist"`} - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe state rule: %v", err) - } - states := snapshot.Rules[0].Rule.ConnectionStates - if len(states) != 2 || states[0] != "established" || states[1] != "related" { - t.Fatalf("connection states were not normalized: %#v", states) - } - if !snapshot.Rules[0].Protected { - t.Fatalf("established whitelist was not protected: %#v", snapshot.Rules[0]) - } -} - -func TestObserveProtectsSystemPresetChains(t *testing.T) { - before := testScope("1PANEL_BASIC_BEFORE") - beforeSnapshot, err := NewAdapterWithReader(&fakeRuleReader{output: `-A 1PANEL_BASIC_BEFORE -p tcp --dport 8080 -j ACCEPT`}).Observe(context.Background(), before) - if err != nil || len(beforeSnapshot.Rules) != 1 || !beforeSnapshot.Rules[0].Protected { - t.Fatalf("BEFORE preset rule was not protected: snapshot=%#v err=%v", beforeSnapshot, err) - } - after := testScope("1PANEL_BASIC_AFTER") - afterSnapshot, err := NewAdapterWithReader(&fakeRuleReader{output: `-A 1PANEL_BASIC_AFTER -p tcp --dport 8080 -j ACCEPT`}).Observe(context.Background(), after) - if err != nil || len(afterSnapshot.Rules) != 1 || !afterSnapshot.Rules[0].Protected { - t.Fatalf("AFTER preset rule was not protected: snapshot=%#v err=%v", afterSnapshot, err) - } -} - -func TestObserveRejectsScopeOutsideOwnedChains(t *testing.T) { - scope := testScope("INPUT") - _, err := NewAdapterWithReader(&fakeRuleReader{}).Observe(context.Background(), scope) - if err == nil { - t.Fatal("expected INPUT scope to be rejected") - } -} - -type fakeRuleReader struct { - output string - err error - multiportErr error - checks []filter.Family -} - -type fakeRuleWriter struct { - commands []filter.NativeCommand - saved bool - saveCalls int - saveErrors []error - runErrors []error -} - -func (f *fakeRuleWriter) Run(_ context.Context, command filter.NativeCommand) error { - f.commands = append(f.commands, command) - if len(f.runErrors) >= len(f.commands) { - return f.runErrors[len(f.commands)-1] - } - return nil -} - -func (f *fakeRuleWriter) Save(context.Context, filter.Scope) error { - f.saveCalls++ - f.saved = true - if len(f.saveErrors) >= f.saveCalls { - return f.saveErrors[f.saveCalls-1] - } - return nil -} - -func executorObserved(rule filter.FirewallRule, position int) filter.ObservedRule { - return filter.ObservedRule{ - Rule: rule, ParseStatus: filter.ParseStatusSupported, - Locator: filter.Locator{Provider: filter.ProviderIptables, ScopeKey: rule.Scope.Key(), Position: &position}, - } -} - -func (f *fakeRuleReader) ListChain(context.Context, filter.Scope) (string, error) { - return f.output, f.err -} - -func (f *fakeRuleReader) CheckMultiport(_ context.Context, family filter.Family) error { - f.checks = append(f.checks, family) - return f.multiportErr -} - -func testScope(chain string) filter.Scope { - return testScopeFamily(chain, filter.FamilyIPv4) -} - -func testScopeFamily(chain string, family filter.Family) filter.Scope { - return filter.Scope{ - Provider: filter.ProviderIptables, Family: family, Table: "filter", Chain: chain, Direction: filter.DirectionInput, - } -} diff --git a/agent/utils/firewall/filter/providers/nftables/adapter_test.go b/agent/utils/firewall/filter/providers/nftables/adapter_test.go deleted file mode 100644 index f4290a9b4813..000000000000 --- a/agent/utils/firewall/filter/providers/nftables/adapter_test.go +++ /dev/null @@ -1,243 +0,0 @@ -package nftables - -import ( - "context" - "errors" - "strings" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -type fakeBackend struct { - output string - commands []filter.NativeCommand - saves int - failAt int -} - -func (f *fakeBackend) ListChain(context.Context, filter.Scope) (string, error) { return f.output, nil } -func (f *fakeBackend) Run(_ context.Context, command filter.NativeCommand) error { - f.commands = append(f.commands, command) - if f.failAt != 0 && len(f.commands) == f.failAt { - return errors.New("run failed") - } - return nil -} -func (f *fakeBackend) Save(context.Context) error { f.saves++; return nil } - -func TestObserveNativeNftablesRules(t *testing.T) { - scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC") - backend := &fakeBackend{output: `table ip nft_1panel_filter { - chain NFT_1PANEL_BASIC { - meta l4proto tcp ip saddr 10.0.0.0/8 tcp dport 443 accept comment "1panel-rule:web" # handle 12 - meta l4proto udp udp dport 53 drop # handle 13 - meta l4proto tcp ct state established,related tcp dport 8443 accept comment "1panel-rule:stateful" # handle 14 - } -}`} - snapshot, err := NewAdapterWithBackend(backend).Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe: %v", err) - } - if len(snapshot.Rules) != 3 || snapshot.Rules[0].Marker != "1panel-rule:web" || snapshot.Rules[0].Locator.NativeID != "12" { - t.Fatalf("unexpected snapshot: %#v", snapshot) - } - if snapshot.Rules[0].Rule.SourceAddress != "10.0.0.0/8" || snapshot.Rules[0].Rule.DestinationPort != "443" || snapshot.Rules[1].Rule.Action != filter.ActionDrop { - t.Fatalf("unexpected parsed rules: %#v", snapshot.Rules) - } - if len(snapshot.Rules[2].Rule.ConnectionStates) != 2 || snapshot.Rules[2].Rule.ConnectionStates[0] != "established" { - t.Fatalf("connection states were not parsed: %#v", snapshot.Rules[2]) - } -} - -func TestObserveIgnoresChainDeclarationHandle(t *testing.T) { - scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC") - backend := &fakeBackend{output: `table ip nft_1panel_filter { - chain NFT_1PANEL_BASIC { # handle 3 - } -}`} - snapshot, err := NewAdapterWithBackend(backend).Observe(context.Background(), scope) - if err != nil { - t.Fatalf("observe empty chain: %v", err) - } - if len(snapshot.Rules) != 0 { - t.Fatalf("chain declaration was parsed as a rule: %#v", snapshot.Rules) - } -} - -func TestApplyCompensatesFailedRulesetTransaction(t *testing.T) { - scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC") - snapshot, err := filter.NewSnapshot(scope, nil) - if err != nil { - t.Fatal(err) - } - rule := filter.FirewallRule{UUID: "web", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept} - backend := &fakeBackend{failAt: 1} - adapter := NewAdapterWithBackend(backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile: %v", err) - } - if _, err := adapter.Apply(context.Background(), plan); err == nil { - t.Fatal("expected ruleset transaction failure") - } - if len(backend.commands) != 2 { - t.Fatalf("expected failed transaction and rollback transaction; got %#v", backend.commands) - } - if backend.saves != 1 { - t.Fatalf("rollback was not persisted: saves=%d", backend.saves) - } -} - -func TestCompileApplyAndRollbackRebuildOwnedChain(t *testing.T) { - scope := testScope(filter.FamilyIPv6, "1PANEL_BASIC") - snapshot, err := filter.NewSnapshot(scope, nil) - if err != nil { - t.Fatal(err) - } - rule := filter.FirewallRule{UUID: "dns6", Scope: scope, Protocol: "udp", DestinationPort: "53", Action: filter.ActionAccept} - backend := &fakeBackend{} - adapter := NewAdapterWithBackend(backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile: %v", err) - } - commands := plan.Rules[0].Commands - if len(commands) != 1 || commands[0].Executable != "nft" || strings.Join(commands[0].Args, " ") != "-f -" { - t.Fatalf("unexpected commands: %#v", commands) - } - joined := commands[0].Stdin - for _, want := range []string{"add rule ip6 nft_1panel_filter NFT_1PANEL_BASIC", "udp dport 53", `comment "1panel-rule:dns6"`} { - if !strings.Contains(joined, want) { - t.Fatalf("command %q does not contain %q", joined, want) - } - } - if _, err := adapter.Apply(context.Background(), plan); err != nil { - t.Fatalf("apply: %v", err) - } - if len(backend.commands) != 1 || backend.saves != 1 { - t.Fatalf("unexpected apply calls: commands=%#v saves=%d", backend.commands, backend.saves) - } - if err := adapter.Rollback(context.Background(), plan); err != nil { - t.Fatalf("rollback: %v", err) - } - if len(backend.commands) != 2 || backend.commands[1].Stdin != "flush chain ip6 nft_1panel_filter NFT_1PANEL_BASIC\n" || backend.saves != 2 { - t.Fatalf("unexpected rollback calls: commands=%#v saves=%d", backend.commands, backend.saves) - } -} - -func TestCompileAndApplyBatchCreateUsesOneRulesetTransaction(t *testing.T) { - scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC") - snapshot, err := filter.NewSnapshot(scope, nil) - if err != nil { - t.Fatal(err) - } - web := filter.FirewallRule{UUID: "web", Scope: scope, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - tls := filter.FirewallRule{UUID: "tls", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept} - backend := &fakeBackend{} - adapter := NewAdapterWithBackend(backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{ - {Operation: filter.ChangeCreate, After: &web}, - {Operation: filter.ChangeCreate, After: &tls}, - }) - if err != nil { - t.Fatalf("compile batch create: %v", err) - } - if len(plan.Rules) != 2 || len(plan.Rules[0].Commands) != 1 || len(plan.Rules[1].Commands) != 0 { - t.Fatalf("batch was not compiled into one transaction: %#v", plan.Rules) - } - script := plan.Rules[0].Commands[0].Stdin - if strings.Count(script, "flush chain ") != 1 || strings.Count(script, "add rule ") != 2 || - !strings.Contains(script, "1panel-rule:web") || !strings.Contains(script, "1panel-rule:tls") { - t.Fatalf("unexpected batch ruleset:\n%s", script) - } - result, err := adapter.Apply(context.Background(), plan) - if err != nil { - t.Fatalf("apply batch create: %v", err) - } - if len(result.Applied) != 2 || len(backend.commands) != 1 || backend.saves != 1 { - t.Fatalf("batch create was not applied once: result=%#v commands=%#v saves=%d", result, backend.commands, backend.saves) - } -} - -func TestCompileBatchDeletePreservesExternalRules(t *testing.T) { - scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC") - positionOne, positionTwo, positionThree := 1, 2, 3 - web := filter.FirewallRule{UUID: "web", Scope: scope, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - tls := filter.FirewallRule{UUID: "tls", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept} - external := filter.ObservedRule{ - Rule: filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindOpaque}, Raw: "fib saddr . iif oif missing drop", - ParseStatus: filter.ParseStatusOpaque, Locator: filter.Locator{Provider: filter.ProviderNftables, ScopeKey: scope.Key(), Position: &positionOne}, - } - webObserved := observedRule(web, "1panel-rule:web", positionTwo, "") - tlsObserved := observedRule(tls, "1panel-rule:tls", positionThree, "") - snapshot, err := filter.NewSnapshot(scope, []filter.ObservedRule{external, webObserved, tlsObserved}) - if err != nil { - t.Fatal(err) - } - adapter := NewAdapterWithBackend(&fakeBackend{}) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{ - {Operation: filter.ChangeDelete, Before: &tls, Locator: &tlsObserved.Locator}, - {Operation: filter.ChangeDelete, Before: &web, Locator: &webObserved.Locator}, - }) - if err != nil { - t.Fatalf("compile batch delete: %v", err) - } - applyScript := plan.Rules[0].Commands[0].Stdin - if strings.Count(applyScript, "flush chain ") != 1 || strings.Count(applyScript, "add rule ") != 1 || - !strings.Contains(applyScript, external.Raw) || strings.Contains(applyScript, "1panel-rule:web") || strings.Contains(applyScript, "1panel-rule:tls") { - t.Fatalf("unexpected batch delete ruleset:\n%s", applyScript) - } - rollbackScript := plan.Rules[0].RollbackCommands[0].Stdin - if !strings.Contains(rollbackScript, "1panel-rule:web") || !strings.Contains(rollbackScript, "1panel-rule:tls") || - !strings.Contains(rollbackScript, external.Raw) { - t.Fatalf("rollback does not restore the original ruleset:\n%s", rollbackScript) - } -} - -func TestVerifyBatchCreateChecksEveryRule(t *testing.T) { - scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC") - snapshot, err := filter.NewSnapshot(scope, nil) - if err != nil { - t.Fatal(err) - } - web := filter.FirewallRule{UUID: "web", Scope: scope, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept} - tls := filter.FirewallRule{UUID: "tls", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept} - backend := &fakeBackend{output: strings.Join([]string{ - `meta l4proto tcp tcp dport 80 accept comment "1panel-rule:web" # handle 10`, - `meta l4proto tcp tcp dport 443 accept comment "1panel-rule:tls" # handle 11`, - }, "\n")} - adapter := NewAdapterWithBackend(backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{ - {Operation: filter.ChangeCreate, After: &web}, - {Operation: filter.ChangeCreate, After: &tls}, - }) - if err != nil { - t.Fatalf("compile batch create: %v", err) - } - verified, err := adapter.Verify(context.Background(), plan) - if err != nil || !verified.Matched { - t.Fatalf("verify complete batch: result=%#v err=%v", verified, err) - } - backend.output = `meta l4proto tcp tcp dport 80 accept comment "1panel-rule:web" # handle 10` - verified, err = adapter.Verify(context.Background(), plan) - if err != nil || verified.Matched { - t.Fatalf("missing batch member passed verification: result=%#v err=%v", verified, err) - } -} - -func TestParseOpaqueRulePreservesRawExpression(t *testing.T) { - scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC") - backend := &fakeBackend{output: "fib saddr . iif oif missing drop # handle 9\n"} - snapshot, err := NewAdapterWithBackend(backend).Observe(context.Background(), scope) - if err != nil || len(snapshot.Rules) != 1 { - t.Fatalf("observe opaque: snapshot=%#v err=%v", snapshot, err) - } - if snapshot.Rules[0].ParseStatus != filter.ParseStatusOpaque || snapshot.Rules[0].Raw == "" { - t.Fatalf("opaque rule was not preserved: %#v", snapshot.Rules[0]) - } -} - -func testScope(family filter.Family, chain string) filter.Scope { - return filter.Scope{Provider: filter.ProviderNftables, Family: family, Table: "filter", Chain: chain, Direction: filter.DirectionInput} -} diff --git a/agent/utils/firewall/filter/providers/ufw/adapter_test.go b/agent/utils/firewall/filter/providers/ufw/adapter_test.go deleted file mode 100644 index b14d38b42a20..000000000000 --- a/agent/utils/firewall/filter/providers/ufw/adapter_test.go +++ /dev/null @@ -1,782 +0,0 @@ -package ufw - -import ( - "context" - "errors" - "fmt" - "reflect" - "strings" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -const numberedFixture = `Status: active - - To Action From - -- ------ ---- -[ 1] 22/tcp ALLOW IN Anywhere # 1panel-rule:rule-v4 -[ 2] Anywhere DENY IN 172.16.10.111 -[ 3] OpenSSH ALLOW IN Anywhere -[ 4] 443/tcp (v6) ALLOW IN Anywhere (v6) # web v6 -[ 5] Anywhere (v6) on eth0 REJECT IN 2001:db8::/64 (v6) -[ 6] Anywhere ALLOW FWD Anywhere on lxdbr0 -[ 7] 53 ALLOW IN Anywhere -[ 8] 6000:6010/udp on eth1 ALLOW IN 10.0.0.0/8 -[ 9] 22,80,443/tcp ALLOW IN Anywhere -[10] 2222/tcp LIMIT IN Anywhere # SSH limit -[11] 25/tcp DENY IN 2001:db8::/32 -[12] a-very-long-application-profile-name ALLOW IN Anywhere -[13] Anywhere on eth0 ALLOW IN Anywhere (log) -[14] 443/tcp ALLOW OUT Anywhere on eth0 (out) -` - -type fakeReader struct { - outputs map[string]string - errors map[string]error - calls [][]string -} - -type scriptedBackend struct { - numbered []string - writes []filter.NativeCommand - failAt int -} - -func (b *scriptedBackend) Read(_ context.Context, args ...string) (string, error) { - if stringsKey(args) == "status verbose" { - return "Status: active\nDefault: deny (incoming), allow (outgoing), deny (routed)\n", nil - } - if stringsKey(args) != "status numbered" || len(b.numbered) == 0 { - return "", fmt.Errorf("unexpected read: %v", args) - } - output := b.numbered[0] - b.numbered = b.numbered[1:] - return output, nil -} - -func (b *scriptedBackend) Run(_ context.Context, command filter.NativeCommand) error { - b.writes = append(b.writes, command) - if b.failAt > 0 && len(b.writes) == b.failAt { - return errors.New("write failed") - } - return nil -} - -func (f *fakeReader) Read(_ context.Context, args ...string) (string, error) { - f.calls = append(f.calls, append([]string(nil), args...)) - key := stringsKey(args) - return f.outputs[key], f.errors[key] -} - -func TestObserveNumberedIncomingIPv4PreservesGlobalPositions(t *testing.T) { - reader := &fakeReader{outputs: map[string]string{ - "status numbered": numberedFixture, - "status verbose": "Status: active\nDefault: deny (incoming), allow (outgoing), deny (routed)\n", - }} - adapter := NewAdapterWithReader(reader) - snapshot, err := adapter.Observe(context.Background(), ufwScope(filter.FamilyIPv4)) - if err != nil { - t.Fatalf("observe: %v", err) - } - if got, want := len(snapshot.Rules), 9; got != want { - t.Fatalf("expected %d IPv4 rules, got %d", want, got) - } - positions := make([]int, 0, len(snapshot.Rules)) - for _, rule := range snapshot.Rules { - positions = append(positions, *rule.Locator.Position) - } - if !reflect.DeepEqual(positions, []int{1, 2, 3, 7, 8, 9, 10, 12, 13}) { - t.Fatalf("unexpected positions: %v", positions) - } - first := snapshot.Rules[0] - if first.ParseStatus != filter.ParseStatusSupported || first.Marker != "1panel-rule:rule-v4" || - first.Rule.DestinationPort != "22" || first.Rule.Protocol != "tcp" || first.Rule.Action != filter.ActionAccept { - t.Fatalf("unexpected first rule: %#v", first) - } - second := snapshot.Rules[1] - if second.ParseStatus != filter.ParseStatusSupported || second.Rule.SourceAddress != "172.16.10.111/32" || second.Rule.Action != filter.ActionDrop { - t.Fatalf("unexpected source deny: %#v", second) - } - application := snapshot.Rules[2] - if application.ParseStatus != filter.ParseStatusOpaque || application.Rule.NativeKind != filter.NativeKindUFWApplication || - application.Rule.Description != "OpenSSH" || application.Rule.Protocol != "" || application.Rule.Action != filter.ActionAccept { - t.Fatalf("UFW application profile was not identified: %#v", application) - } - eighth := snapshot.Rules[4] - if eighth.ParseStatus != filter.ParseStatusSupported || eighth.Rule.DestinationPort != "6000-6010" || eighth.Rule.Interface != "eth1" { - t.Fatalf("unexpected range rule: %#v", eighth) - } - barePort := snapshot.Rules[3] - if barePort.ParseStatus != filter.ParseStatusSupported || barePort.Rule.Protocol != "all" || barePort.Rule.DestinationPort != "53" { - t.Fatalf("unexpected bare-port rule: %#v", barePort) - } - for _, index := range []int{2, 7} { - if snapshot.Rules[index].ParseStatus != filter.ParseStatusOpaque { - t.Fatalf("expected rule at slice index %d to be opaque: %#v", index, snapshot.Rules[index]) - } - } - multiPort := snapshot.Rules[5] - if multiPort.ParseStatus != filter.ParseStatusSupported || multiPort.Rule.Protocol != "tcp" || multiPort.Rule.DestinationPort != "22,80,443" { - t.Fatalf("multi-port display fields were not preserved: %#v", multiPort) - } - limited := snapshot.Rules[6] - if limited.ParseStatus != filter.ParseStatusPartial || limited.Rule.Protocol != "tcp" || limited.Rule.DestinationPort != "2222" || limited.Rule.Action != filter.ActionAccept { - t.Fatalf("limited rule display fields were not preserved: %#v", limited) - } - logged := snapshot.Rules[8] - if logged.ParseStatus != filter.ParseStatusPartial || logged.Rule.Protocol != "all" || logged.Rule.Interface != "eth0" { - t.Fatalf("logged rule display fields were not preserved: %#v", logged) - } - longApplication := snapshot.Rules[7] - if longApplication.Rule.NativeKind != filter.NativeKindUFWApplication || - longApplication.Rule.Description != "a-very-long-application-profile-name" { - t.Fatalf("long UFW application profile was not preserved: %#v", longApplication) - } - if len(snapshot.Notices) != 0 { - t.Fatalf("active UFW default policy should not create a notice: %#v", snapshot.Notices) - } - if !reflect.DeepEqual(reader.calls, [][]string{{"status", "numbered"}}) { - t.Fatalf("unexpected commands: %#v", reader.calls) - } -} - -func TestNativeDetailRunsUFWAppInfoWithCompleteProfileName(t *testing.T) { - info := "Profile: Nginx Full\nTitle: Web Server (Nginx, HTTP + HTTPS)\nDescription: Small, but very powerful and efficient web server\n\nPorts:\n 80,443/tcp" - reader := &fakeReader{outputs: map[string]string{"app info Nginx Full": info}} - - got, err := NewAdapterWithReader(reader).NativeDetail(context.Background(), "Nginx Full", false) - if err != nil { - t.Fatalf("load UFW application detail: %v", err) - } - if got != info || !reflect.DeepEqual(reader.calls, [][]string{{"app", "info", "Nginx Full"}}) { - t.Fatalf("unexpected UFW application detail: output=%q calls=%#v", got, reader.calls) - } -} - -func TestNativeDetailRejectsInvalidUFWProfileName(t *testing.T) { - if _, err := NewAdapterWithReader(&fakeReader{}).NativeDetail(context.Background(), "../OpenSSH", false); !errors.Is(err, filter.ErrInvalidRule) { - t.Fatalf("expected invalid UFW profile rejection, got %v", err) - } -} - -func TestObserveNumberedIncomingIPv6KeepsFamilyGap(t *testing.T) { - reader := &fakeReader{outputs: map[string]string{ - "status numbered": numberedFixture, - "status verbose": "Status: active\nDefault: deny (incoming), allow (outgoing), deny (routed)\n", - }} - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), ufwScope(filter.FamilyIPv6)) - if err != nil { - t.Fatalf("observe: %v", err) - } - if len(snapshot.Rules) != 3 || *snapshot.Rules[0].Locator.Position != 4 || *snapshot.Rules[1].Locator.Position != 5 || - *snapshot.Rules[2].Locator.Position != 11 { - t.Fatalf("unexpected IPv6 rules: %#v", snapshot.Rules) - } - first := snapshot.Rules[0] - if first.ParseStatus != filter.ParseStatusSupported || first.Rule.Description != "web v6" || first.Rule.DestinationPort != "443" { - t.Fatalf("unexpected IPv6 port rule: %#v", first) - } - second := snapshot.Rules[1] - if second.ParseStatus != filter.ParseStatusSupported || second.Rule.Interface != "eth0" || - second.Rule.SourceAddress != "2001:db8::/64" || second.Rule.Action != filter.ActionReject { - t.Fatalf("unexpected IPv6 reject: %#v", second) - } - third := snapshot.Rules[2] - if third.ParseStatus != filter.ParseStatusSupported || third.Rule.SourceAddress != "2001:db8::/32" || third.Rule.DestinationPort != "25" { - t.Fatalf("unexpected explicit IPv6 rule without family suffix: %#v", third) - } -} - -func TestObserveScopesReadsNumberedRulesOnce(t *testing.T) { - reader := &fakeReader{outputs: map[string]string{"status numbered": numberedFixture}} - snapshots, err := NewAdapterWithReader(reader).ObserveScopes(context.Background(), []filter.Scope{ - ufwScope(filter.FamilyIPv4), - ufwScope(filter.FamilyIPv6), - }) - if err != nil { - t.Fatalf("observe UFW families: %v", err) - } - if !reflect.DeepEqual(reader.calls, [][]string{{"status", "numbered"}}) { - t.Fatalf("UFW numbered rules were not read exactly once: %#v", reader.calls) - } - if len(snapshots) != 2 || snapshots[0].Scope.Family != filter.FamilyIPv4 || len(snapshots[0].Rules) != 9 || - snapshots[1].Scope.Family != filter.FamilyIPv6 || len(snapshots[1].Rules) != 3 { - t.Fatalf("unexpected multi-family snapshots: %#v", snapshots) - } -} - -func TestObserveBarePortRulesAreSupportedForBothFamilies(t *testing.T) { - output := `Status: active -[ 1] 22 ALLOW IN Anywhere -[ 2] 22 (v6) ALLOW IN Anywhere (v6) -` - for _, test := range []struct { - family filter.Family - position int - }{ - {family: filter.FamilyIPv4, position: 1}, - {family: filter.FamilyIPv6, position: 2}, - } { - rules := parseNumberedRules(ufwScope(test.family), output) - if len(rules) != 1 { - t.Fatalf("expected one %s bare-port rule, got %#v", test.family, rules) - } - rule := rules[0] - if rule.ParseStatus != filter.ParseStatusSupported || rule.Rule.Protocol != "all" || - rule.Rule.DestinationPort != "22" || rule.Locator.Position == nil || *rule.Locator.Position != test.position { - t.Fatalf("unexpected %s bare-port rule: %#v", test.family, rule) - } - } -} - -func TestObserveInactiveUFWReturnsNoticeAndEmptyInventory(t *testing.T) { - reader := &fakeReader{outputs: map[string]string{ - "status numbered": "Status: inactive\n", - "status verbose": "Status: inactive\n", - }} - snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), ufwScope(filter.FamilyIPv4)) - if err != nil { - t.Fatalf("observe: %v", err) - } - if len(snapshot.Rules) != 0 || !hasNotice(snapshot.Notices, filter.ScopeNoticeManagedScopeInactive) { - t.Fatalf("unexpected inactive snapshot: %#v", snapshot) - } -} - -func TestParseAnnotatedMultiPortRulesForBothFamilies(t *testing.T) { - output := `Status: active -[ 1] 80,443/tcp ALLOW IN Anywhere -[ 2] 137,138/udp (Samba) ALLOW IN Anywhere -[ 3] 80,443/tcp (v6) ALLOW IN Anywhere (v6) -[ 4] 137,138/udp (Samba (v6)) ALLOW IN Anywhere (v6)` - tests := []struct { - family filter.Family - positions []int - }{ - {family: filter.FamilyIPv4, positions: []int{1, 2}}, - {family: filter.FamilyIPv6, positions: []int{3, 4}}, - } - for _, test := range tests { - t.Run(string(test.family), func(t *testing.T) { - rules := parseNumberedRules(ufwScope(test.family), output) - if len(rules) != 2 { - t.Fatalf("expected two %s rules, got %#v", test.family, rules) - } - for index, observed := range rules { - if observed.Locator.Position == nil || *observed.Locator.Position != test.positions[index] { - t.Fatalf("unexpected %s rule %d: %#v", test.family, index, observed) - } - } - if rules[0].ParseStatus != filter.ParseStatusSupported || rules[0].Rule.NativeKind != filter.NativeKindUFWRule || - rules[0].Rule.Protocol != "tcp" || rules[0].Rule.DestinationPort != "80,443" || - rules[0].Rule.Description != "" { - t.Fatalf("plain multi-port fields were not preserved: %#v", rules[0]) - } - if rules[1].ParseStatus != filter.ParseStatusPartial || rules[1].Rule.NativeKind != filter.NativeKindUFWApplication || - rules[1].Rule.Protocol != "udp" || rules[1].Rule.DestinationPort != "137,138" || - rules[1].Rule.Description != "Samba" { - t.Fatalf("annotated application fields were not preserved: %#v", rules[1]) - } - }) - } -} - -func TestObserveRejectsScopeOutsideIncoming(t *testing.T) { - adapter := NewAdapterWithReader(&fakeReader{}) - _, err := adapter.Observe(context.Background(), filter.Scope{ - Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, Chain: "outgoing", Direction: filter.Direction("output"), - }) - if !errors.Is(err, filter.ErrInvalidScope) { - t.Fatalf("expected invalid output scope, got %v", err) - } -} - -func TestCapabilitiesOnlyAdvertiseAtomicIncomingFamilies(t *testing.T) { - capabilities, err := NewAdapterWithReader(&fakeReader{}).Capabilities(context.Background()) - if err != nil { - t.Fatalf("capabilities: %v", err) - } - if !capabilities.Marker || !capabilities.SupportsScope(ufwScope(filter.FamilyIPv4)) || - !capabilities.SupportsScope(ufwScope(filter.FamilyIPv6)) || - !capabilities.ExplicitPosition || - capabilities.SupportsScope(filter.Scope{Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, Chain: "outgoing", Direction: filter.Direction("output")}) { - t.Fatalf("unexpected capabilities: %#v", capabilities) - } -} - -func TestPrepareRuleRejectsOptionsUFWCannotSynchronize(t *testing.T) { - rule := writableRule(filter.FamilyIPv4, "", "8080") - rule.SourcePort = "1024" - if _, err := NewAdapterWithReader(&fakeReader{}).PrepareRule(rule); !errors.Is(err, filter.ErrInvalidRule) { - t.Fatalf("unsupported source port passed UFW preview validation: %v", err) - } -} - -func TestCompileCreateUsesFamilyExplicitFullSyntax(t *testing.T) { - tests := []struct { - name string - family filter.Family - address string - }{ - {name: "IPv4", family: filter.FamilyIPv4, address: "0.0.0.0/0"}, - {name: "IPv6", family: filter.FamilyIPv6, address: "::/0"}, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - snapshot := mustSnapshot(t, ufwScope(test.family), nil) - rule := writableRule(test.family, "new-rule", "8080") - plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeCreate, After: &rule, Append: true, - }}) - if err != nil { - t.Fatalf("compile: %v", err) - } - want := []string{"allow", "in", "proto", "tcp", "from", test.address, "to", test.address, "port", "8080", "comment", "1panel-rule:new-rule"} - if !reflect.DeepEqual(plan.Rules[0].Commands[0].Args, want) { - t.Fatalf("unexpected command:\nwant: %#v\n got: %#v", want, plan.Rules[0].Commands[0].Args) - } - if got := plan.Rules[0].RollbackCommands[0].Args[:3]; !reflect.DeepEqual(got, []string{"--force", "delete", "allow"}) { - t.Fatalf("unexpected rollback prefix: %#v", got) - } - }) - } -} - -func TestCompileCreateBarePortUsesFamilyExplicitSyntaxWithoutProtocol(t *testing.T) { - tests := []struct { - family filter.Family - address string - }{ - {family: filter.FamilyIPv4, address: "0.0.0.0/0"}, - {family: filter.FamilyIPv6, address: "::/0"}, - } - for _, test := range tests { - t.Run(string(test.family), func(t *testing.T) { - snapshot := mustSnapshot(t, ufwScope(test.family), nil) - rule := filter.FirewallRule{ - UUID: "dns", Scope: ufwScope(test.family), NativeKind: filter.NativeKindUFWRule, - Protocol: "all", DestinationPort: "53", Action: filter.ActionAccept, - } - plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeCreate, After: &rule, Append: true, - }}) - if err != nil { - t.Fatalf("compile bare-port create: %v", err) - } - want := []string{ - "allow", "in", "from", test.address, "to", test.address, - "port", "53", "comment", "1panel-rule:dns", - } - if !reflect.DeepEqual(plan.Rules[0].Commands[0].Args, want) { - t.Fatalf("unexpected bare-port command:\nwant: %#v\n got: %#v", want, plan.Rules[0].Commands[0].Args) - } - }) - } -} - -func TestCompileCreateMultiportUsesNativeUFWPortSet(t *testing.T) { - snapshot := mustSnapshot(t, ufwScope(filter.FamilyIPv4), nil) - rule := writableRule(filter.FamilyIPv4, "web", "80,443,8080-8090") - plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeCreate, After: &rule, Append: true, - }}) - if err != nil { - t.Fatalf("compile multiport create: %v", err) - } - wantPort := []string{"port", "80,443,8080:8090"} - args := plan.Rules[0].Commands[0].Args - found := false - for index := 0; index+1 < len(args); index++ { - if reflect.DeepEqual(args[index:index+2], wantPort) { - found = true - break - } - } - if !found { - t.Fatalf("UFW port set was not preserved: %#v", args) - } -} - -func TestCompileCreateAppendsWithoutInvalidNextPosition(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - existing := parseNumberedRules(scope, "Status: active\n[ 8] 80/tcp ALLOW IN Anywhere\n") - snapshot := mustSnapshot(t, scope, existing) - order := int64(9) - rule := writableRule(filter.FamilyIPv4, "appended", "8080") - rule.OrderIndex = &order - - plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeCreate, - After: &rule, - Append: true, - }}) - if err != nil { - t.Fatalf("compile append: %v", err) - } - command := plan.Rules[0].Commands[0] - if len(command.Args) == 0 || command.Args[0] != "allow" { - t.Fatalf("append must not use invalid insert 9: %#v", command.Args) - } - if plan.Rules[0].Expected.Locator.Position == nil || *plan.Rules[0].Expected.Locator.Position != 9 { - t.Fatalf("append verification position was not preserved: %#v", plan.Rules[0].Expected.Locator) - } -} - -func TestCompileLastRuleMutationUsesAppendForWriteAndRestore(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - observed := parseNumberedRules(scope, "Status: active\n[ 8] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n") - snapshot := mustSnapshot(t, scope, observed) - locator := observed[0].Locator - order := int64(8) - updated := writableRule(filter.FamilyIPv4, "managed", "443") - updated.OrderIndex = &order - - updatePlan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeUpdate, - After: &updated, - Locator: &locator, - Append: true, - RestoreAtEnd: true, - }}) - if err != nil { - t.Fatalf("compile last-rule update: %v", err) - } - if got := updatePlan.Rules[0]; got.Commands[1].Args[0] != "allow" || got.RollbackCommands[0].Args[0] != "allow" { - t.Fatalf("last-rule update must append both new and restored rules: %#v", got) - } - - before := observed[0].Rule - before.UUID = "managed" - deletePlan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeDelete, - Before: &before, - Locator: &locator, - RestoreAtEnd: true, - }}) - if err != nil { - t.Fatalf("compile last-rule delete: %v", err) - } - if got := deletePlan.Rules[0].RollbackCommands[0].Args; len(got) == 0 || got[0] != "allow" { - t.Fatalf("last-rule delete rollback must append: %#v", got) - } -} - -func TestCompileAdoptUpdatesCommentWithoutChangingNumber(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - observed := parseNumberedRules(scope, "Status: active\n[ 7] Anywhere DENY IN 172.16.10.111 # imported\n") - if len(observed) != 1 { - t.Fatalf("expected one observed rule: %#v", observed) - } - snapshot := mustSnapshot(t, scope, observed) - rule := observed[0].Rule - rule.UUID = "adopted-rule" - locator := observed[0].Locator - plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeAdopt, After: &rule, Locator: &locator, - }}) - if err != nil { - t.Fatalf("compile adoption: %v", err) - } - rulePlan := plan.Rules[0] - if len(rulePlan.Commands) != 1 || rulePlan.Commands[0].Args[0] != "deny" || - rulePlan.Commands[0].Args[len(rulePlan.Commands[0].Args)-1] != "1panel-rule:adopted-rule" { - t.Fatalf("unexpected adoption command: %#v", rulePlan.Commands) - } - if rulePlan.RollbackCommands[0].Args[len(rulePlan.RollbackCommands[0].Args)-1] != "imported" || - rulePlan.Expected.Locator.Position == nil || *rulePlan.Expected.Locator.Position != 7 { - t.Fatalf("adoption did not preserve comment and position: %#v", rulePlan) - } -} - -func TestCompileUpdateAndDeleteRequireOwnedNumberedRule(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - observed := parseNumberedRules(scope, "Status: active\n[ 3] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n") - snapshot := mustSnapshot(t, scope, observed) - locator := observed[0].Locator - updated := writableRule(filter.FamilyIPv4, "managed", "443") - update, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeUpdate, After: &updated, Locator: &locator, - }}) - if err != nil { - t.Fatalf("compile update: %v", err) - } - if got := update.Rules[0].Commands; len(got) != 2 || !reflect.DeepEqual(got[0].Args, []string{"--force", "delete", "3"}) || - !reflect.DeepEqual(got[1].Args[:2], []string{"insert", "3"}) { - t.Fatalf("unexpected update commands: %#v", got) - } - before := observed[0].Rule - before.UUID = "managed" - deletePlan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeDelete, Before: &before, Locator: &locator, - }}) - if err != nil { - t.Fatalf("compile delete: %v", err) - } - if !reflect.DeepEqual(deletePlan.Rules[0].Commands[0].Args, []string{"--force", "delete", "3"}) { - t.Fatalf("unexpected delete command: %#v", deletePlan.Rules[0].Commands) - } - - external := observed[0] - external.Marker = "" - external.Rule.Description = "external" - externalSnapshot := mustSnapshot(t, scope, []filter.ObservedRule{external}) - _, err = NewAdapterWithReader(&fakeReader{}).Compile(externalSnapshot, []filter.DesiredChange{{ - Operation: filter.ChangeDelete, Before: &before, Locator: &locator, - }}) - if err == nil { - t.Fatal("expected deletion of an external rule to be rejected") - } -} - -func TestCompileUpdateUsesRequestedGlobalPosition(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - observed := parseNumberedRules(scope, "Status: active\n[ 3] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n") - snapshot := mustSnapshot(t, scope, observed) - locator := observed[0].Locator - updated := writableRule(filter.FamilyIPv4, "managed", "443") - target := int64(1) - updated.OrderIndex = &target - plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeUpdate, Before: &observed[0].Rule, After: &updated, Locator: &locator, - }}) - if err != nil { - t.Fatalf("compile positioned update: %v", err) - } - rulePlan := plan.Rules[0] - if len(rulePlan.Commands) != 2 || !reflect.DeepEqual(rulePlan.Commands[0].Args, []string{"--force", "delete", "3"}) || - !reflect.DeepEqual(rulePlan.Commands[1].Args[:2], []string{"insert", "1"}) || - rulePlan.Expected.Locator.Position == nil || *rulePlan.Expected.Locator.Position != 1 { - t.Fatalf("unexpected positioned update: %#v", rulePlan) - } -} - -func TestCompileRejectsInactiveAndProtected(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - rule := writableRule(filter.FamilyIPv4, "rule", "8080") - inactive := mustSnapshot(t, scope, nil) - inactive.Notices = []filter.ScopeNotice{{Code: filter.ScopeNoticeManagedScopeInactive}} - _, err := NewAdapterWithReader(&fakeReader{}).Compile(inactive, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if !errors.Is(err, filter.ErrProviderUnavailable) { - t.Fatalf("expected inactive provider error, got %v", err) - } - - observed := parseNumberedRules(scope, "Status: active\n[ 1] 8080/tcp ALLOW IN Anywhere\n") - observed[0].Protected = true - protected := mustSnapshot(t, scope, observed) - locator := observed[0].Locator - rule.UUID = "protected" - _, err = NewAdapterWithReader(&fakeReader{}).Compile(protected, []filter.DesiredChange{{Operation: filter.ChangeAdopt, After: &rule, Locator: &locator}}) - if !errors.Is(err, filter.ErrProtectedRule) { - t.Fatalf("expected protected rule error, got %v", err) - } - _, err = NewAdapterWithReader(&fakeReader{}).Compile(mustSnapshot(t, scope, nil), []filter.DesiredChange{{Operation: filter.ChangeReorder, After: &rule}}) - if !errors.Is(err, filter.ErrUnsupportedScope) { - t.Fatalf("expected unsupported standalone reorder error, got %v", err) - } - broadDeny := filter.FirewallRule{ - UUID: "deny-all", Scope: scope, NativeKind: filter.NativeKindUFWRule, Protocol: "all", Action: filter.ActionDrop, - } - _, err = NewAdapterWithReader(&fakeReader{}).Compile(mustSnapshot(t, scope, nil), []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &broadDeny}}) - if !errors.Is(err, filter.ErrLockoutRisk) { - t.Fatalf("expected broad deny lockout error, got %v", err) - } -} - -func TestApplyVerifiesMarkerAcrossBothFamilies(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - snapshot := mustSnapshot(t, scope, nil) - rule := writableRule(filter.FamilyIPv4, "created", "8080") - backend := &scriptedBackend{numbered: []string{ - "Status: active\n[ 1] 8080/tcp ALLOW IN Anywhere # 1panel-rule:created\n", - "Status: active\n", - }} - adapter := NewAdapterWithBackend(backend, backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile: %v", err) - } - if _, err = adapter.Apply(context.Background(), plan); err != nil { - t.Fatalf("apply: %v", err) - } - if len(backend.writes) != 1 { - t.Fatalf("unexpected writes: %#v", backend.writes) - } -} - -func TestApplyVerifiesMultiportUpdate(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - beforeOutput := "Status: active\n[ 4] 4422,8088/tcp ALLOW IN Anywhere # 1panel-rule:managed\n" - observed := parseNumberedRules(scope, beforeOutput) - if len(observed) != 1 || observed[0].ParseStatus != filter.ParseStatusSupported { - t.Fatalf("unexpected existing multiport rule: %#v", observed) - } - snapshot := mustSnapshot(t, scope, observed) - updated := writableRule(filter.FamilyIPv4, "managed", "4422,8088,7944") - order := int64(4) - updated.OrderIndex = &order - backend := &scriptedBackend{numbered: []string{ - "Status: active\n[ 4] 4422,8088,7944/tcp ALLOW IN Anywhere # 1panel-rule:managed\n", - "Status: active\n", - }} - adapter := NewAdapterWithBackend(backend, backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeUpdate, After: &updated, Locator: &observed[0].Locator, - }}) - if err != nil { - t.Fatalf("compile multiport update: %v", err) - } - if _, err = adapter.Apply(context.Background(), plan); err != nil { - t.Fatalf("verify multiport update: %v", err) - } - if len(backend.writes) != 2 || backend.writes[1].Args[0] != "insert" || backend.writes[1].Args[1] != "4" { - t.Fatalf("unexpected multiport update writes: %#v", backend.writes) - } -} - -func TestApplyCompensatesFamilyExpansion(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - snapshot := mustSnapshot(t, scope, nil) - rule := writableRule(filter.FamilyIPv4, "expanded", "8080") - backend := &scriptedBackend{numbered: []string{ - "Status: active\n[ 1] 8080/tcp ALLOW IN Anywhere # 1panel-rule:expanded\n", - "Status: active\n[ 2] 8080/tcp (v6) ALLOW IN Anywhere (v6) # 1panel-rule:expanded\n", - }} - adapter := NewAdapterWithBackend(backend, backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile: %v", err) - } - if _, err = adapter.Apply(context.Background(), plan); err == nil { - t.Fatal("expected one-to-many verification failure") - } - if len(backend.writes) != 2 || backend.writes[1].Args[0] != "--force" || backend.writes[1].Args[1] != "delete" { - t.Fatalf("expected compensating delete, got %#v", backend.writes) - } -} - -func TestRollbackReversesFullyAppliedUFWPlan(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - snapshot := mustSnapshot(t, scope, nil) - rule := writableRule(filter.FamilyIPv4, "rollback", "8080") - backend := &scriptedBackend{} - adapter := NewAdapterWithBackend(backend, backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile rollback plan: %v", err) - } - if err := adapter.Rollback(context.Background(), plan); err != nil { - t.Fatalf("rollback applied plan: %v", err) - } - if len(backend.writes) != 1 || !reflect.DeepEqual(backend.writes[0].Args[:2], []string{"--force", "delete"}) { - t.Fatalf("unexpected rollback writes: %#v", backend.writes) - } -} - -func TestApplyCompensatesFailedUpdate(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - beforeOutput := "Status: active\n[ 3] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n" - observed := parseNumberedRules(scope, beforeOutput) - snapshot := mustSnapshot(t, scope, observed) - locator := observed[0].Locator - updated := writableRule(filter.FamilyIPv4, "managed", "443") - backend := &scriptedBackend{numbered: []string{ - "Status: active\n[ 3] 443/tcp ALLOW IN Anywhere # 1panel-rule:managed\n", - }, failAt: 2} - adapter := NewAdapterWithBackend(backend, backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeUpdate, After: &updated, Locator: &locator}}) - if err != nil { - t.Fatalf("compile: %v", err) - } - if _, err = adapter.Apply(context.Background(), plan); err == nil || !strings.Contains(err.Error(), "write failed") { - t.Fatalf("expected write failure, got %v", err) - } - if len(backend.writes) != 4 || - !reflect.DeepEqual(backend.writes[2].Args[:3], []string{"--force", "delete", "allow"}) || - !reflect.DeepEqual(backend.writes[3].Args[:2], []string{"insert", "3"}) { - t.Fatalf("expected possibly-applied new rule removal and old rule restore after failed update: %#v", backend.writes) - } -} - -func TestApplyDoesNotCompensateFailedUpdateCommandWithoutSideEffect(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - beforeOutput := "Status: active\n[ 3] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n" - observed := parseNumberedRules(scope, beforeOutput) - snapshot := mustSnapshot(t, scope, observed) - updated := writableRule(filter.FamilyIPv4, "managed", "443") - backend := &scriptedBackend{numbered: []string{"Status: active\n"}, failAt: 2} - adapter := NewAdapterWithBackend(backend, backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeUpdate, After: &updated, Locator: &observed[0].Locator, - }}) - if err != nil { - t.Fatalf("compile: %v", err) - } - if _, err = adapter.Apply(context.Background(), plan); err == nil || !strings.Contains(err.Error(), "write failed") { - t.Fatalf("expected write failure, got %v", err) - } - if len(backend.writes) != 3 || !reflect.DeepEqual(backend.writes[2].Args[:2], []string{"insert", "3"}) { - t.Fatalf("expected only the successfully deleted old rule to be restored: %#v", backend.writes) - } -} - -func TestApplyUsesCompiledNumberedSnapshotWithoutPreflight(t *testing.T) { - scope := ufwScope(filter.FamilyIPv4) - snapshot := mustSnapshot(t, scope, nil) - rule := writableRule(filter.FamilyIPv4, "stale", "8080") - backend := &scriptedBackend{numbered: []string{ - "Status: active\n[ 1] 8080/tcp ALLOW IN Anywhere # 1panel-rule:stale\n", - "Status: active\n", - }} - adapter := NewAdapterWithBackend(backend, backend) - plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - if err != nil { - t.Fatalf("compile: %v", err) - } - if _, err = adapter.Apply(context.Background(), plan); err != nil { - t.Fatalf("apply compiled plan: %v", err) - } - if len(backend.writes) != 1 { - t.Fatalf("compiled plan was not written: %#v", backend.writes) - } -} - -func ufwScope(family filter.Family) filter.Scope { - return filter.Scope{Provider: filter.ProviderUFW, Family: family, Chain: "incoming", Direction: filter.DirectionInput} -} - -func writableRule(family filter.Family, uuid, port string) filter.FirewallRule { - return filter.FirewallRule{ - UUID: uuid, Scope: ufwScope(family), NativeKind: filter.NativeKindUFWRule, - Protocol: "tcp", DestinationPort: port, Action: filter.ActionAccept, - } -} - -func mustSnapshot(t *testing.T, scope filter.Scope, rules []filter.ObservedRule) filter.Snapshot { - t.Helper() - snapshot, err := filter.NewSnapshot(scope, rules) - if err != nil { - t.Fatalf("snapshot: %v", err) - } - return snapshot -} - -func hasNotice(notices []filter.ScopeNotice, code filter.ScopeNoticeCode) bool { - for _, notice := range notices { - if notice.Code == code { - return true - } - } - return false -} - -func stringsKey(values []string) string { - result := "" - for index, value := range values { - if index != 0 { - result += " " - } - result += value - } - return result -} diff --git a/agent/utils/firewall/filter/runtime/runtime.go b/agent/utils/firewall/filter/runtime/runtime.go new file mode 100644 index 000000000000..1a3eaacec89f --- /dev/null +++ b/agent/utils/firewall/filter/runtime/runtime.go @@ -0,0 +1,307 @@ +package runtime + +import ( + "context" + "errors" + "fmt" + + "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" + filterfirewalld "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/firewalld" + filteriptables "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/iptables" + filternftables "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/nftables" + filterufw "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/ufw" +) + +type SnapshotPolicy func(context.Context, filter.Snapshot) (filter.Snapshot, error) + +type Engine struct { + adapter filter.Adapter + policy SnapshotPolicy +} + +type Registry map[filter.Provider]*Engine + +func NewRegistry(policy SnapshotPolicy) Registry { + return Registry{ + filter.ProviderIptables: New(filteriptables.NewAdapter(), policy), + filter.ProviderNftables: New(filternftables.NewAdapter(), policy), + filter.ProviderFirewalld: New(filterfirewalld.NewAdapter(), policy), + filter.ProviderUFW: New(filterufw.NewAdapter(), policy), + } +} + +func New(adapter filter.Adapter, policy SnapshotPolicy) *Engine { + return &Engine{adapter: adapter, policy: policy} +} + +func (r Registry) Resolve(provider filter.Provider) (*Engine, error) { + engine, exists := r[provider] + if !exists || engine == nil || engine.adapter == nil { + return nil, fmt.Errorf("%w: %s", filter.ErrAdapterUnavailable, provider) + } + return engine, nil +} + +func (r Registry) Providers() []filter.Provider { + providers := make([]filter.Provider, 0, len(r)) + for provider := range r { + providers = append(providers, provider) + } + return providers +} + +func (e *Engine) Provider() filter.Provider { + if e == nil || e.adapter == nil { + return "" + } + return e.adapter.Provider() +} + +func (e *Engine) Observe(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) { + snapshot, err := e.adapter.Observe(ctx, scope) + if err != nil { + return filter.Snapshot{}, err + } + if e.policy == nil { + return snapshot, nil + } + return e.policy(ctx, snapshot) +} + +func (e *Engine) ObserveScopes(ctx context.Context, scopes []filter.Scope) ([]filter.Snapshot, error) { + observer, ok := e.adapter.(filter.MultiScopeObserver) + if !ok { + return nil, fmt.Errorf("%w: %s multi-scope inventory", filter.ErrAdapterUnavailable, e.adapter.Provider()) + } + snapshots, err := observer.ObserveScopes(ctx, scopes) + if err != nil { + return nil, err + } + if e.policy == nil { + return snapshots, nil + } + for index := range snapshots { + snapshots[index], err = e.policy(ctx, snapshots[index]) + if err != nil { + return nil, err + } + } + return snapshots, nil +} + +func (e *Engine) ObserveMutation(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) { + snapshot, err := e.Observe(ctx, scope) + if err != nil { + return filter.Snapshot{}, err + } + for _, notice := range snapshot.Notices { + if notice.Code == filter.ScopeNoticeManagedScopeInactive || notice.Code == filter.ScopeNoticeManagedScopeMissing { + return filter.Snapshot{}, fmt.Errorf("%w: managed firewall scope is unavailable", filter.ErrProviderUnavailable) + } + } + return snapshot, nil +} + +func (e *Engine) Prepare(rule filter.FirewallRule) (filter.FirewallRule, error) { + preparer, ok := e.adapter.(filter.RulePreparer) + if !ok { + return rule, nil + } + return preparer.PrepareRule(rule) +} + +func (e *Engine) CheckRule(ctx context.Context, rule filter.FirewallRule) error { + checker, ok := e.adapter.(filter.RuleChecker) + if !ok { + return nil + } + return checker.CheckRule(ctx, rule) +} + +func (e *Engine) CompileDesired( + ctx context.Context, + policyUUID string, + origin filter.RuleOrigin, + rules []filter.FirewallRule, +) ([]filter.DesiredRule, error) { + result := make([]filter.DesiredRule, 0, len(rules)) + scopeOrdinals := make(map[string]int) + for _, rule := range rules { + prepared, err := e.Prepare(rule) + if err != nil { + return nil, err + } + if err = e.CheckRule(ctx, prepared); err != nil { + return nil, err + } + ruleKey, err := filter.RuleKey(prepared) + if err != nil { + return nil, err + } + scopeKey := prepared.Scope.Key() + ordinal := scopeOrdinals[scopeKey] + scopeOrdinals[scopeKey] = ordinal + 1 + prepared.UUID = compiledRuleUUID(policyUUID, ruleKey, ordinal) + result = append(result, filter.DesiredRule{ + UUID: policyUUID, Rule: prepared, RuleKey: ruleKey, Origin: origin, + Marker: "1panel-rule:" + prepared.UUID, + }) + } + return result, nil +} + +func (e *Engine) ValidatePosition( + ctx context.Context, + snapshot filter.Snapshot, + rule filter.FirewallRule, + target int64, +) error { + if rule.Scope.Provider == filter.ProviderUFW { + minimum, maximum := positionBounds(snapshot) + if target < minimum || target > maximum { + return fmt.Errorf( + "%w: target position %d is outside the %s range %d-%d", + filter.ErrInvalidRule, target, rule.Scope.Family, minimum, maximum, + ) + } + return nil + } + maximum, err := e.MaxPosition(ctx, snapshot, rule) + if err != nil { + return err + } + if target > maximum { + return fmt.Errorf("%w: target position %d is out of range 1-%d", filter.ErrInvalidRule, target, maximum) + } + return nil +} + +func (e *Engine) AppendPosition(ctx context.Context, snapshot filter.Snapshot, rule filter.FirewallRule) (int64, error) { + if rule.Scope.Family == filter.FamilyIPv4 { + return snapshotMaxPosition(snapshot) + 1, nil + } + maximum, err := e.MaxPosition(ctx, snapshot, rule) + if err != nil { + return 0, err + } + return maximum + 1, nil +} + +func (e *Engine) MaxPosition( + ctx context.Context, + snapshot filter.Snapshot, + rule filter.FirewallRule, +) (int64, error) { + maximum := snapshotMaxPosition(snapshot) + if rule.Scope.Provider != filter.ProviderUFW { + return maximum, nil + } + relatedScope := rule.Scope + if relatedScope.Family == filter.FamilyIPv4 { + relatedScope.Family = filter.FamilyIPv6 + } else { + relatedScope.Family = filter.FamilyIPv4 + } + relatedSnapshot, err := e.ObserveMutation(ctx, relatedScope) + if err != nil { + return 0, err + } + if relatedMaximum := snapshotMaxPosition(relatedSnapshot); relatedMaximum > maximum { + maximum = relatedMaximum + } + return maximum, nil +} + +func (e *Engine) NativeDetail(ctx context.Context, name string, permanent bool) (string, error) { + reader, ok := e.adapter.(filter.NativeDetailReader) + if !ok { + return "", fmt.Errorf("%w: native details for %s", filter.ErrAdapterUnavailable, e.Provider()) + } + return reader.NativeDetail(ctx, name, permanent) +} + +func (e *Engine) Capabilities(ctx context.Context) (filter.Capabilities, error) { + return e.adapter.Capabilities(ctx) +} + +func (e *Engine) Execute(ctx context.Context, snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, filter.VerifyResult, error) { + plan, err := e.adapter.Compile(snapshot, changes) + if err != nil { + return filter.BackendPlan{}, filter.VerifyResult{}, err + } + result, err := e.adapter.Apply(ctx, plan) + if err != nil { + return plan, filter.VerifyResult{}, err + } + if result.Verification != nil { + if !result.Verification.Matched { + if rollbackErr := e.Rollback(ctx, plan); rollbackErr != nil { + return plan, *result.Verification, errors.Join(filter.ErrVerificationFailed, rollbackErr) + } + } + return plan, *result.Verification, nil + } + verification, err := e.adapter.Verify(ctx, plan) + if err != nil { + return plan, verification, e.rollback(ctx, plan, err) + } + if !verification.Matched { + if rollbackErr := e.Rollback(ctx, plan); rollbackErr != nil { + return plan, verification, errors.Join(filter.ErrVerificationFailed, rollbackErr) + } + } + return plan, verification, nil +} + +func (e *Engine) Rollback(ctx context.Context, plan filter.BackendPlan) error { + rollbacker, ok := e.adapter.(filter.PlanRollbacker) + if !ok { + return fmt.Errorf("%w: provider %s does not support applied-plan rollback", filter.ErrAdapterUnavailable, e.adapter.Provider()) + } + return rollbacker.Rollback(ctx, plan) +} + +func (e *Engine) rollback(ctx context.Context, plan filter.BackendPlan, cause error) error { + if err := e.Rollback(ctx, plan); err != nil { + return errors.Join(cause, fmt.Errorf("rollback applied firewall plan: %w", err)) + } + return cause +} + +func positionBounds(snapshot filter.Snapshot) (int64, int64) { + minimum, maximum := int64(0), int64(0) + for _, observed := range snapshot.Rules { + if observed.Locator.Position == nil { + continue + } + position := int64(*observed.Locator.Position) + if minimum == 0 || position < minimum { + minimum = position + } + if position > maximum { + maximum = position + } + } + return minimum, maximum +} + +func snapshotMaxPosition(snapshot filter.Snapshot) int64 { + var maximum int64 + for _, observed := range snapshot.Rules { + if observed.Locator.Position != nil && int64(*observed.Locator.Position) > maximum { + maximum = int64(*observed.Locator.Position) + } + } + return maximum +} + +func compiledRuleUUID(policyUUID, ruleKey string, scopeOrdinal int) string { + if scopeOrdinal == 0 { + return policyUUID + } + const suffixLength = 12 + if len(ruleKey) > suffixLength { + ruleKey = ruleKey[:suffixLength] + } + return fmt.Sprintf("%s-%d-%s", policyUUID, scopeOrdinal+1, ruleKey) +} diff --git a/agent/utils/firewall/filter/safety_test.go b/agent/utils/firewall/filter/safety_test.go deleted file mode 100644 index afb907069dd5..000000000000 --- a/agent/utils/firewall/filter/safety_test.go +++ /dev/null @@ -1,139 +0,0 @@ -package filter - -import "testing" - -func TestProtectSnapshotTreatsAllProtocolAsCoveringProtectedTransport(t *testing.T) { - scope := Scope{Provider: ProviderUFW, Family: FamilyIPv4, Chain: UFWInputChain, Direction: DirectionInput} - rules := []ObservedRule{ - protectedPortTestRule(scope, "all", "22", 1), - protectedPortTestRule(scope, "udp", "22", 2), - protectedPortTestRule(scope, "all", "53", 3), - } - snapshot, err := NewSnapshot(scope, rules) - if err != nil { - t.Fatalf("create snapshot: %v", err) - } - - protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Port: "22", Protocol: "tcp"}}) - if err != nil { - t.Fatalf("protect snapshot: %v", err) - } - if !protected.Rules[0].Protected { - t.Fatal("all-protocol port 22 did not cover protected 22/tcp") - } - if protected.Rules[1].Protected { - t.Fatal("22/udp incorrectly matched protected 22/tcp") - } - if protected.Rules[2].Protected { - t.Fatal("unrelated all-protocol port was protected") - } -} - -func TestProtectSnapshotMarksBareUFWPortForBothFamilies(t *testing.T) { - for _, family := range []Family{FamilyIPv4, FamilyIPv6} { - scope := Scope{Provider: ProviderUFW, Family: family, Chain: UFWInputChain, Direction: DirectionInput} - snapshot, err := NewSnapshot(scope, []ObservedRule{protectedPortTestRule(scope, "all", "22", 1)}) - if err != nil { - t.Fatalf("create %s snapshot: %v", family, err) - } - protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Port: "22", Protocol: "tcp"}}) - if err != nil { - t.Fatalf("protect %s snapshot: %v", family, err) - } - if !protected.Rules[0].Protected { - t.Fatalf("bare UFW port 22 was not protected for %s", family) - } - } -} - -func TestProtectSnapshotMatchesPortSetsAndRanges(t *testing.T) { - scope := Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput} - rules := []ObservedRule{ - protectedPortTestRule(scope, "tcp", "22,80,443", 1), - protectedPortTestRule(scope, "tcp", "8000-9000", 2), - protectedPortTestRule(scope, "udp", "22,443", 3), - } - snapshot, err := NewSnapshot(scope, rules) - if err != nil { - t.Fatalf("create snapshot: %v", err) - } - protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Port: "22", Protocol: "tcp"}, {Port: "8080", Protocol: "tcp"}}) - if err != nil { - t.Fatalf("protect snapshot: %v", err) - } - if !protected.Rules[0].Protected || !protected.Rules[1].Protected || protected.Rules[2].Protected { - t.Fatalf("unexpected port-set protection: %#v", protected.Rules) - } -} - -func TestProtectSnapshotRespectsConfiguredFamily(t *testing.T) { - for _, family := range []Family{FamilyIPv4, FamilyIPv6} { - scope := Scope{Provider: ProviderIptables, Family: family, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput} - snapshot, err := NewSnapshot(scope, []ObservedRule{protectedPortTestRule(scope, "tcp", "443", 1)}) - if err != nil { - t.Fatalf("create %s snapshot: %v", family, err) - } - protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Family: "ipv6", Port: "443", Protocol: "tcp"}}) - if err != nil { - t.Fatalf("protect %s snapshot: %v", family, err) - } - if protected.Rules[0].Protected != (family == FamilyIPv6) { - t.Fatalf("unexpected %s protection state: %#v", family, protected.Rules[0]) - } - } -} - -func TestProtectSnapshotDoesNotReclassifyManagedRule(t *testing.T) { - scope := Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput} - managed := protectedPortTestRule(scope, "all", "", 1) - managed.Marker = "1panel-rule:managed-rule" - snapshot, err := NewSnapshot(scope, []ObservedRule{managed}) - if err != nil { - t.Fatalf("create snapshot: %v", err) - } - - protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Port: "9999", Protocol: "tcp"}}) - if err != nil { - t.Fatalf("protect snapshot: %v", err) - } - if protected.Rules[0].Protected { - t.Fatal("managed rule was reclassified as a protected system rule") - } -} - -func TestManagedBroadAllowDoesNotBlockManagedRuleEdit(t *testing.T) { - scope := Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput} - broadAllow := protectedPortTestRule(scope, "all", "", 1) - broadAllow.Marker = "1panel-rule:broad-allow" - target := protectedPortTestRule(scope, "tcp", "55101", 2) - target.Marker = "1panel-rule:edit-target" - snapshot, err := NewSnapshot(scope, []ObservedRule{broadAllow, target}) - if err != nil { - t.Fatalf("create snapshot: %v", err) - } - snapshot, err = ProtectSnapshot(snapshot, []PortWhitelist{{Port: "9999", Protocol: "tcp"}}) - if err != nil { - t.Fatalf("protect snapshot: %v", err) - } - after := target.Rule - after.Protocol = "udp" - after.SourceAddress = "198.51.100.21/32" - after.DestinationPort = "55113" - after.Action = ActionDrop - if err := GuardMutation(snapshot, target, after, "", PortWhitelist{Port: "9999", Protocol: "tcp"}); err != nil { - t.Fatalf("managed rule edit was blocked by another managed allow rule: %v", err) - } -} - -func protectedPortTestRule(scope Scope, protocol, port string, position int) ObservedRule { - return ObservedRule{ - Rule: FirewallRule{ - Scope: scope, NativeKind: NativeKindUFWRule, Protocol: protocol, - DestinationPort: port, Action: ActionAccept, - }, - Locator: Locator{ - Provider: scope.Provider, ScopeKey: scope.Key(), Position: &position, - }, - ParseStatus: ParseStatusSupported, - } -} diff --git a/agent/utils/firewall/forwarding/adapter.go b/agent/utils/firewall/forwarding/adapter.go deleted file mode 100644 index 11de736356c1..000000000000 --- a/agent/utils/firewall/forwarding/adapter.go +++ /dev/null @@ -1,52 +0,0 @@ -package forwarding - -import ( - "strings" - - "github.com/1Panel-dev/1Panel/agent/constant" -) - -const ( - FamilyIPv4 = constant.FirewallFamilyIPv4 - FamilyIPv6 = constant.FirewallFamilyIPv6 - - ChainPreRouting = "1PANEL_PREROUTING" - ChainPostRouting = "1PANEL_POSTROUTING" - ChainForward = "1PANEL_FORWARD" - - ForwardFile = "1panel_forward.rules" - PreRoutingFile = "1panel_forward_pre.rules" - PostRoutingFile = "1panel_forward_post.rules" -) - -type Rule struct { - Num string - Family string - Protocol string - Port string - TargetIP string - TargetPort string - Interface string -} - -func (r Rule) Identity() string { - return strings.Join([]string{r.Family, r.Protocol, r.Port, r.TargetIP, r.TargetPort, r.Interface}, "\x00") -} - -type OperationType string - -const ( - OperationAdd OperationType = "add" - OperationRemove OperationType = "remove" -) - -type Adapter interface { - Name() string - List() ([]Rule, error) - Reconcile(rules []Rule) error - Enable() error - Cleanup() error - InitStatus() (bool, bool, error) - FamilyStatus(family string) (bool, bool, error) - Replay() error -} diff --git a/agent/utils/firewall/forwarding/adapter_test.go b/agent/utils/firewall/forwarding/adapter_test.go deleted file mode 100644 index 5eff7394c63c..000000000000 --- a/agent/utils/firewall/forwarding/adapter_test.go +++ /dev/null @@ -1,20 +0,0 @@ -package forwarding - -import "testing" - -func TestRuleIdentityIncludesEveryIdentityField(t *testing.T) { - base := Rule{Family: FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80", Interface: "eth0"} - variants := []Rule{ - {Family: FamilyIPv6, Protocol: base.Protocol, Port: base.Port, TargetIP: base.TargetIP, TargetPort: base.TargetPort, Interface: base.Interface}, - {Family: base.Family, Protocol: "udp", Port: base.Port, TargetIP: base.TargetIP, TargetPort: base.TargetPort, Interface: base.Interface}, - {Family: base.Family, Protocol: base.Protocol, Port: "8081", TargetIP: base.TargetIP, TargetPort: base.TargetPort, Interface: base.Interface}, - {Family: base.Family, Protocol: base.Protocol, Port: base.Port, TargetIP: "127.0.0.2", TargetPort: base.TargetPort, Interface: base.Interface}, - {Family: base.Family, Protocol: base.Protocol, Port: base.Port, TargetIP: base.TargetIP, TargetPort: "81", Interface: base.Interface}, - {Family: base.Family, Protocol: base.Protocol, Port: base.Port, TargetIP: base.TargetIP, TargetPort: base.TargetPort, Interface: "eth1"}, - } - for _, variant := range variants { - if variant.Identity() == base.Identity() { - t.Fatalf("identity collision for %#v", variant) - } - } -} diff --git a/agent/utils/firewall/forwarding/forwarding.go b/agent/utils/firewall/forwarding/forwarding.go new file mode 100644 index 000000000000..1f2f13b85d05 --- /dev/null +++ b/agent/utils/firewall/forwarding/forwarding.go @@ -0,0 +1,203 @@ +package forwarding + +import ( + "errors" + "fmt" + "net/netip" + "strconv" + "strings" + "sync" + + "github.com/1Panel-dev/1Panel/agent/constant" + "github.com/1Panel-dev/1Panel/agent/utils/re" +) + +var ErrRuleExists = errors.New("forwarding rule already exists") + +const ( + FamilyIPv4 = constant.FirewallFamilyIPv4 + FamilyIPv6 = constant.FirewallFamilyIPv6 + + ChainPreRouting = "1PANEL_PREROUTING" + ChainPostRouting = "1PANEL_POSTROUTING" + ChainForward = "1PANEL_FORWARD" + + ForwardFile = "1panel_forward.rules" + PreRoutingFile = "1panel_forward_pre.rules" + PostRoutingFile = "1panel_forward_post.rules" +) + +type Rule struct { + Num string + Family string + Protocol string + Port string + TargetIP string + TargetPort string + Interface string +} + +func (r Rule) Identity() string { + return strings.Join([]string{r.Family, r.Protocol, r.Port, r.TargetIP, r.TargetPort, r.Interface}, "\x00") +} + +type OperationType string + +const ( + OperationAdd OperationType = "add" + OperationRemove OperationType = "remove" +) + +type Adapter interface { + Name() string + List() ([]Rule, error) + Reconcile(rules []Rule) error + Enable() error + Cleanup() error + InitStatus() (bool, bool, error) + FamilyStatus(family string) (bool, bool, error) + Replay() error +} + +type Status struct { + Name string + Version string + IsInit bool + IsBind bool +} + +type RuntimeClient interface { + Version() (string, error) +} + +type Manager struct { + adapter Adapter + runtime RuntimeClient +} + +func NewManager(adapter Adapter, runtime RuntimeClient) *Manager { + return &Manager{adapter: adapter, runtime: runtime} +} + +func (m *Manager) Status() (Status, error) { + status := Status{Name: m.adapter.Name(), Version: "-"} + var versionErr error + var initErr error + var wg sync.WaitGroup + wg.Add(1) + if m.runtime != nil { + wg.Add(1) + go func() { + defer wg.Done() + status.Version, versionErr = m.runtime.Version() + }() + } + go func() { + defer wg.Done() + status.IsInit, status.IsBind, initErr = m.adapter.InitStatus() + }() + wg.Wait() + return status, errors.Join(versionErr, initErr) +} + +func (m *Manager) List(info, strategy string) ([]Rule, error) { + rules, err := m.adapter.List() + if err != nil { + return nil, err + } + if strategy != "" { + return []Rule{}, nil + } + filtered := make([]Rule, 0, len(rules)) + for _, rule := range rules { + if info != "" && !strings.Contains(rule.Port, info) && + !strings.Contains(rule.TargetPort, info) && !strings.Contains(rule.TargetIP, info) { + continue + } + filtered = append(filtered, rule) + } + return filtered, nil +} + +func (m *Manager) Enable() error { return m.adapter.Enable() } + +func (m *Manager) Reconcile(rules []Rule) error { return m.adapter.Reconcile(rules) } + +func (m *Manager) Cleanup() error { return m.adapter.Cleanup() } + +func (m *Manager) FamilyStatus(family string) (bool, bool, error) { + return m.adapter.FamilyStatus(family) +} + +func (m *Manager) Replay() error { return m.adapter.Replay() } + +func (m *Manager) Name() string { return m.adapter.Name() } + +func NormalizeRule(rule Rule) (Rule, error) { + rule.Family = strings.ToLower(strings.TrimSpace(rule.Family)) + if rule.Family == "" { + rule.Family = FamilyIPv4 + } + if rule.Family != FamilyIPv4 && rule.Family != FamilyIPv6 { + return Rule{}, fmt.Errorf("unsupported forwarding family %q", rule.Family) + } + rule.Protocol = strings.ToLower(strings.TrimSpace(rule.Protocol)) + if rule.Protocol != "tcp" && rule.Protocol != "udp" { + return Rule{}, fmt.Errorf("unsupported forwarding protocol %q", rule.Protocol) + } + var err error + if rule.Port, err = normalizeForwardPort(rule.Port); err != nil { + return Rule{}, fmt.Errorf("invalid forwarding port: %w", err) + } + if rule.TargetPort, err = normalizeForwardPort(rule.TargetPort); err != nil { + return Rule{}, fmt.Errorf("invalid forwarding target port: %w", err) + } + rule.TargetIP = strings.TrimSpace(rule.TargetIP) + if rule.TargetIP == "" || strings.EqualFold(rule.TargetIP, "localhost") { + if rule.Family == FamilyIPv6 { + rule.TargetIP = "::1" + } else { + rule.TargetIP = "127.0.0.1" + } + } + address, err := netip.ParseAddr(rule.TargetIP) + if err == nil { + address = address.Unmap() + } + if err != nil || (rule.Family == FamilyIPv4) != address.Is4() { + return Rule{}, fmt.Errorf("invalid %s forwarding target %q", rule.Family, rule.TargetIP) + } + rule.TargetIP = address.String() + rule.Interface = strings.TrimSpace(rule.Interface) + if rule.Interface == "all" || rule.Interface == "*" { + rule.Interface = "" + } + if rule.Interface != "" && !re.ForwardInterfaceRegex.MatchString(rule.Interface) { + return Rule{}, fmt.Errorf("invalid forwarding interface %q", rule.Interface) + } + return rule, nil +} + +func normalizeForwardPort(value string) (string, error) { + parts := strings.Split(strings.TrimSpace(value), "-") + if len(parts) < 1 || len(parts) > 2 { + return "", fmt.Errorf("invalid port range %q", value) + } + ports := make([]int, len(parts)) + for index, part := range parts { + port, err := strconv.Atoi(strings.TrimSpace(part)) + if err != nil || port < 1 || port > 65535 { + return "", fmt.Errorf("invalid port %q", part) + } + ports[index] = port + } + if len(ports) == 2 { + if ports[0] > ports[1] { + return "", fmt.Errorf("descending port range %q", value) + } + if ports[0] != ports[1] { + return strconv.Itoa(ports[0]) + "-" + strconv.Itoa(ports[1]), nil + } + } + return strconv.Itoa(ports[0]), nil +} diff --git a/agent/utils/firewall/forwarding/manager.go b/agent/utils/firewall/forwarding/manager.go deleted file mode 100644 index 6e80fea408af..000000000000 --- a/agent/utils/firewall/forwarding/manager.go +++ /dev/null @@ -1,83 +0,0 @@ -package forwarding - -import ( - "errors" - "strings" - "sync" -) - -var ErrRuleExists = errors.New("forwarding rule already exists") - -type Status struct { - Name string - Version string - IsInit bool - IsBind bool -} - -type RuntimeClient interface { - Version() (string, error) -} - -type Manager struct { - adapter Adapter - runtime RuntimeClient -} - -func NewManager(adapter Adapter, runtime RuntimeClient) *Manager { - return &Manager{adapter: adapter, runtime: runtime} -} - -func (m *Manager) Status() (Status, error) { - status := Status{Name: m.adapter.Name(), Version: "-"} - var versionErr error - var initErr error - var wg sync.WaitGroup - wg.Add(1) - if m.runtime != nil { - wg.Add(1) - go func() { - defer wg.Done() - status.Version, versionErr = m.runtime.Version() - }() - } - go func() { - defer wg.Done() - status.IsInit, status.IsBind, initErr = m.adapter.InitStatus() - }() - wg.Wait() - return status, errors.Join(versionErr, initErr) -} - -func (m *Manager) List(info, strategy string) ([]Rule, error) { - rules, err := m.adapter.List() - if err != nil { - return nil, err - } - if strategy != "" { - return []Rule{}, nil - } - filtered := make([]Rule, 0, len(rules)) - for _, rule := range rules { - if info != "" && !strings.Contains(rule.Port, info) && - !strings.Contains(rule.TargetPort, info) && !strings.Contains(rule.TargetIP, info) { - continue - } - filtered = append(filtered, rule) - } - return filtered, nil -} - -func (m *Manager) Enable() error { return m.adapter.Enable() } - -func (m *Manager) Reconcile(rules []Rule) error { return m.adapter.Reconcile(rules) } - -func (m *Manager) Cleanup() error { return m.adapter.Cleanup() } - -func (m *Manager) FamilyStatus(family string) (bool, bool, error) { - return m.adapter.FamilyStatus(family) -} - -func (m *Manager) Replay() error { return m.adapter.Replay() } - -func (m *Manager) Name() string { return m.adapter.Name() } diff --git a/agent/utils/firewall/forwarding/providers/adapter_contract_test.go b/agent/utils/firewall/forwarding/providers/adapter_contract_test.go deleted file mode 100644 index 438dd5164843..000000000000 --- a/agent/utils/firewall/forwarding/providers/adapter_contract_test.go +++ /dev/null @@ -1,329 +0,0 @@ -package providers - -import ( - "os" - "reflect" - "strings" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/iptables_helper" -) - -type commandCall struct { - name string - args []string -} - -func commandKey(name string, args ...string) string { - return strings.Join(append([]string{name}, args...), " ") -} - -func TestForwardingAdapterFactoryContract(t *testing.T) { - for _, name := range []string{"iptables", "nftables"} { - adapter, err := New(name) - if err != nil { - t.Fatalf("%s: %v", name, err) - } - if adapter.Name() != name { - t.Fatalf("got adapter %q want %q", adapter.Name(), name) - } - if name == "nftables" { - if _, ok := adapter.(*nftablesAdapter); !ok { - t.Fatalf("nftables must use its native forwarding adapter, got %T", adapter) - } - } else if _, ok := adapter.(*iptablesNATAdapter); !ok { - t.Fatalf("%s must use the iptables forwarding adapter, got %T", name, adapter) - } - } - for _, name := range []string{"firewalld", "ufw", "unknown"} { - if _, err := New(name); err == nil { - t.Fatalf("unsupported forwarding provider %q must be rejected", name) - } - } -} - -type backendCall struct { - method string - table string - args []string -} - -type fakeIptablesBackend struct { - calls []backendCall - stdout map[string]string - err error - ipv6 bool -} - -func (f *fakeIptablesBackend) IPv6Available() bool { return f.ipv6 } - -func (f *fakeIptablesBackend) Run(table string, args ...string) error { - f.calls = append(f.calls, backendCall{method: "run", table: table, args: append([]string(nil), args...)}) - return f.err -} - -func (f *fakeIptablesBackend) RunWithStd(table string, args ...string) (string, error) { - f.calls = append(f.calls, backendCall{method: "stdout", table: table, args: append([]string(nil), args...)}) - return f.stdout[commandKey(table, args...)], f.err -} - -func (f *fakeIptablesBackend) RunIPv6(table string, args ...string) error { - f.calls = append(f.calls, backendCall{method: "run6", table: table, args: append([]string(nil), args...)}) - return f.err -} - -func (f *fakeIptablesBackend) RunIPv6WithStd(table string, args ...string) (string, error) { - f.calls = append(f.calls, backendCall{method: "stdout6", table: table, args: append([]string(nil), args...)}) - return f.stdout["ipv6 "+commandKey(table, args...)], f.err -} - -func (f *fakeIptablesBackend) AddChainWithAppend(table, parentChain, chain string) error { - f.calls = append(f.calls, backendCall{method: "add-chain", table: table, args: []string{parentChain, chain}}) - return f.err -} - -func (f *fakeIptablesBackend) AddIPv6ChainWithAppend(table, parentChain, chain string) error { - f.calls = append(f.calls, backendCall{method: "add-chain6", table: table, args: []string{parentChain, chain}}) - return f.err -} - -func (f *fakeIptablesBackend) Restore(family, input string) error { - f.calls = append(f.calls, backendCall{method: "restore", table: family, args: []string{input}}) - return f.err -} - -func (f *fakeIptablesBackend) LoadRulesFromFile(table, chain, fileName string) error { - f.calls = append(f.calls, backendCall{method: "load", table: table, args: []string{chain, fileName}}) - return f.err -} - -func (f *fakeIptablesBackend) LoadIPv6RulesFromFile(table, chain, fileName string) error { - f.calls = append(f.calls, backendCall{method: "load6", table: table, args: []string{chain, fileName}}) - return f.err -} - -func TestIptablesReconcileRebuildsOwnedChains(t *testing.T) { - backend := &fakeIptablesBackend{stdout: map[string]string{}} - adapter := &iptablesNATAdapter{provider: "iptables", backend: backend, system: &fakeForwardingSystem{}} - if err := adapter.Reconcile([]forwarding.Rule{{ - Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80", - }}); err != nil { - t.Fatal(err) - } - wantScript := "*nat\n" + - "-F " + forwarding.ChainPreRouting + "\n" + - "-F " + forwarding.ChainPostRouting + "\n" + - "-A " + forwarding.ChainPreRouting + " -p tcp --dport 8080 -j REDIRECT --to-port 80\n" + - "COMMIT\n" + - "*filter\n" + - "-F " + forwarding.ChainForward + "\n" + - "COMMIT\n" - want := []backendCall{{method: "restore", table: forwarding.FamilyIPv4, args: []string{wantScript}}} - if !reflect.DeepEqual(backend.calls, want) { - t.Fatalf("reconcile transcript changed\ngot %#v\nwant %#v", backend.calls, want) - } -} - -func TestIptablesReconcileBatchesIPv4AndIPv6Separately(t *testing.T) { - backend := &fakeIptablesBackend{stdout: map[string]string{}, ipv6: true} - adapter := &iptablesNATAdapter{provider: "iptables", backend: backend, system: &fakeForwardingSystem{}} - if err := adapter.Reconcile([]forwarding.Rule{{ - Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "8443", TargetIP: "2001:db8::20", TargetPort: "443", - }}); err != nil { - t.Fatal(err) - } - if len(backend.calls) != 2 || backend.calls[0].method != "restore" || backend.calls[0].table != forwarding.FamilyIPv4 || - backend.calls[1].method != "restore" || backend.calls[1].table != forwarding.FamilyIPv6 { - t.Fatalf("expected one restore call per family, got %#v", backend.calls) - } - if !strings.Contains(backend.calls[1].args[0], "--to-destination [2001:db8::20]:443") { - t.Fatalf("unexpected IPv6 restore script:\n%s", backend.calls[1].args[0]) - } -} - -func TestIptablesForwardLifecycleUsesSingleRestoreScript(t *testing.T) { - script := buildIptablesForwardLifecycleScript(map[string]string{ - iptables_helper.NatTab: "-N " + forwarding.ChainPreRouting + "\n-A PREROUTING -j " + forwarding.ChainPreRouting, - iptables_helper.FilterTab: "", - }, true) - if strings.Count(script, "*nat\n") != 1 || strings.Count(script, "*filter\n") != 1 || strings.Count(script, "COMMIT\n") != 2 { - t.Fatalf("unexpected lifecycle restore transaction:\n%s", script) - } - if strings.Contains(script, "-N "+forwarding.ChainPreRouting+"\n") || - !strings.Contains(script, "-N "+forwarding.ChainPostRouting+"\n") || - !strings.Contains(script, "-A FORWARD -j "+forwarding.ChainForward+"\n") { - t.Fatalf("lifecycle restore did not preserve/create the expected chains:\n%s", script) - } -} - -func TestIptablesForwardCleanupUsesSingleRestoreScript(t *testing.T) { - script := buildIptablesForwardLifecycleScript(map[string]string{ - iptables_helper.NatTab: strings.Join([]string{ - "-N " + forwarding.ChainPreRouting, - "-A PREROUTING -j " + forwarding.ChainPreRouting, - "-N " + forwarding.ChainPostRouting, - "-A POSTROUTING -j " + forwarding.ChainPostRouting, - }, "\n"), - iptables_helper.FilterTab: "-N " + forwarding.ChainForward + "\n-A FORWARD -j " + forwarding.ChainForward, - }, false) - for _, line := range []string{ - "-D PREROUTING -j " + forwarding.ChainPreRouting, - "-F " + forwarding.ChainPostRouting, - "-X " + forwarding.ChainForward, - } { - if !strings.Contains(script, line+"\n") { - t.Fatalf("cleanup restore is missing %q:\n%s", line, script) - } - } -} - -type fileWrite struct { - name string - data string -} - -type fakeForwardingSystem struct { - reads map[string][]byte - writes []fileWrite - runs []commandCall -} - -func (f *fakeForwardingSystem) ReadFile(name string) ([]byte, error) { - data, ok := f.reads[name] - if !ok { - return nil, os.ErrNotExist - } - return data, nil -} - -func (f *fakeForwardingSystem) WriteFile(name string, data []byte, _ os.FileMode) error { - f.writes = append(f.writes, fileWrite{name: name, data: string(data)}) - return nil -} - -func (f *fakeForwardingSystem) RunWithOptionalSudo(name string, args ...string) error { - f.runs = append(f.runs, commandCall{name: name, args: append([]string(nil), args...)}) - return nil -} - -func TestIptablesNATEnableReplayAndStatusContract(t *testing.T) { - natStatus := strings.Join([]string{ - "-N THIRD_PARTY_DNAT", - "-A PREROUTING -j THIRD_PARTY_DNAT", - "-N " + forwarding.ChainPreRouting, - "-N " + forwarding.ChainPostRouting, - "-A PREROUTING -j " + forwarding.ChainPreRouting, - "-A POSTROUTING -j " + forwarding.ChainPostRouting, - }, "\n") - filterStatus := "-N " + forwarding.ChainForward + "\n-A FORWARD -j " + forwarding.ChainForward + "\n-A FORWARD -j DOCKER-USER\n" - backend := &fakeIptablesBackend{stdout: map[string]string{ - "nat -S": natStatus, - "filter -S": filterStatus, - }} - system := &fakeForwardingSystem{reads: map[string][]byte{ - "/proc/sys/net/ipv4/ip_forward": []byte("1\n"), - "/etc/sysctl.conf": []byte("net.ipv4.tcp_syncookies = 1\n"), - }} - adapter := &iptablesNATAdapter{provider: "iptables", backend: backend, system: system} - if err := adapter.Enable(); err != nil { - t.Fatal(err) - } - if len(system.writes) != 2 || system.writes[0].name != "/proc/sys/net/ipv4/ip_forward" || - !strings.Contains(system.writes[1].data, "net.ipv4.ip_forward = 1") { - t.Fatalf("sysctl writes changed: %#v", system.writes) - } - if !reflect.DeepEqual(system.runs, []commandCall{{name: "sysctl", args: []string{"-p"}}}) { - t.Fatalf("sysctl transcript changed: %#v", system.runs) - } - init, bind, err := adapter.InitStatus() - if err != nil { - t.Fatal(err) - } - if !init || !bind { - t.Fatalf("expected initialized and bound, got %v %v", init, bind) - } - system.reads["/proc/sys/net/ipv4/ip_forward"] = []byte("0\n") - init, bind, err = adapter.InitStatus() - if err != nil { - t.Fatal(err) - } - if !init || bind { - t.Fatalf("disabled IP forwarding must preserve initialization without reporting a binding, got %v %v", init, bind) - } - backend.calls = nil - if err := adapter.Replay(); err != nil { - t.Fatal(err) - } - wantLoads := []backendCall{ - {method: "load", table: iptables_helper.FilterTab, args: []string{forwarding.ChainForward, forwarding.ForwardFile}}, - {method: "load", table: iptables_helper.NatTab, args: []string{forwarding.ChainPreRouting, forwarding.PreRoutingFile}}, - {method: "load", table: iptables_helper.NatTab, args: []string{forwarding.ChainPostRouting, forwarding.PostRoutingFile}}, - } - if !reflect.DeepEqual(backend.calls, wantLoads) { - t.Fatalf("replay transcript changed: %#v", backend.calls) - } -} - -func TestIptablesListParsingContract(t *testing.T) { - stdout := strings.Join([]string{ - "1 0 0 DNAT tcp -- eth0 * 0.0.0.0/0 0.0.0.0/0 tcp dpt:8080 to:10.0.0.2:80", - "2 0 0 REDIRECT udp -- * * 0.0.0.0/0 0.0.0.0/0 udp dpts:9000:9001 redir ports 53", - }, "\n") - rules := parseIptablesRules(stdout, forwarding.FamilyIPv4) - want := []forwarding.Rule{ - {Num: "1", Family: forwarding.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", Interface: "eth0"}, - {Num: "2", Family: forwarding.FamilyIPv4, Protocol: "udp", Port: "9000-9001", TargetIP: "127.0.0.1", TargetPort: "53", Interface: "*"}, - } - if !reflect.DeepEqual(rules, want) { - t.Fatalf("got %#v want %#v", rules, want) - } -} - -func TestIptablesListParsingAllowsExtraIPv6MatchColumns(t *testing.T) { - stdout := strings.Join([]string{ - "3 0 0 REDIRECT tcp -- * * ::/0 ::/0 tcp dpt:55204 ctstate NEW redir ports 80", - "4 0 0 REDIRECT tcp -- * * ::/0 ::/0 tcp dpt:55205 redir ports", - }, "\n") - rules := parseIptablesRules(stdout, forwarding.FamilyIPv6) - want := []forwarding.Rule{{ - Num: "3", Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "55204", - TargetIP: "::1", TargetPort: "80", Interface: "*", - }} - if !reflect.DeepEqual(rules, want) { - t.Fatalf("got %#v want %#v", rules, want) - } -} - -func TestIptablesIPv6NATContract(t *testing.T) { - stdout := "1 0 0 DNAT tcp -- eth0 * ::/0 ::/0 tcp dpt:8443 to:[2001:db8::20]:443" - rules := parseIptablesRules(stdout, forwarding.FamilyIPv6) - if len(rules) != 1 || rules[0].Family != forwarding.FamilyIPv6 || rules[0].TargetIP != "2001:db8::20" || rules[0].TargetPort != "443" { - t.Fatalf("unexpected parsed IPv6 rules: %#v", rules) - } -} - -func TestEnableIPv4ForwardingReplacesDisabledSetting(t *testing.T) { - content := strings.Join([]string{ - "# net.ipv4.ip_forward = 0", - "net.ipv4.ip_forward=0", - "net.ipv4.tcp_syncookies = 1", - }, "\n") - want := strings.Join([]string{ - "# net.ipv4.ip_forward = 0", - "net.ipv4.ip_forward = 1", - "net.ipv4.tcp_syncookies = 1", - "", - }, "\n") - if got := enableIPv4Forwarding(content); got != want { - t.Fatalf("got %q want %q", got, want) - } -} - -func TestEnableForwardingSysctlsAddsIPv6(t *testing.T) { - content := "net.ipv4.ip_forward = 0\nnet.ipv6.conf.all.forwarding=0\n" - want := "net.ipv4.ip_forward = 1\nnet.ipv6.conf.all.forwarding = 1\n" - if got := enableForwardingSysctls(content, true); got != want { - t.Fatalf("got %q want %q", got, want) - } -} diff --git a/agent/utils/firewall/forwarding/providers/iptables.go b/agent/utils/firewall/forwarding/providers/iptables.go index 586627da7e9c..be8570357614 100644 --- a/agent/utils/firewall/forwarding/providers/iptables.go +++ b/agent/utils/firewall/forwarding/providers/iptables.go @@ -134,7 +134,7 @@ func (l *iptablesNATAdapter) Reconcile(rules []forwarding.Rule) error { forwarding.FamilyIPv6: nil, } for _, rule := range rules { - normalized, err := NormalizeRule(rule) + normalized, err := forwarding.NormalizeRule(rule) if err != nil { return err } diff --git a/agent/utils/firewall/forwarding/providers/nftables.go b/agent/utils/firewall/forwarding/providers/nftables.go index 8588205d3de7..bd982ce2dd52 100644 --- a/agent/utils/firewall/forwarding/providers/nftables.go +++ b/agent/utils/firewall/forwarding/providers/nftables.go @@ -185,7 +185,7 @@ func rebuildNftForwardCommands(rules []forwarding.Rule) ([][]string, error) { } } for _, rule := range rules { - normalized, err := NormalizeRule(rule) + normalized, err := forwarding.NormalizeRule(rule) if err != nil { return nil, err } diff --git a/agent/utils/firewall/forwarding/providers/nftables_test.go b/agent/utils/firewall/forwarding/providers/nftables_test.go deleted file mode 100644 index d6f61fcee9c8..000000000000 --- a/agent/utils/firewall/forwarding/providers/nftables_test.go +++ /dev/null @@ -1,137 +0,0 @@ -package providers - -import ( - "strings" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" -) - -func TestNftForwardNamingContract(t *testing.T) { - if nftForwardTable != "nft_1panel_forward" { - t.Fatalf("unexpected nftables forwarding table %q", nftForwardTable) - } - if nftForwardFile != "1panel_forward.nft" { - t.Fatalf("unexpected nftables forwarding rules file %q", nftForwardFile) - } - wantChains := map[string]string{ - forwarding.ChainPreRouting: "NFT_1PANEL_PREROUTING", - forwarding.ChainPostRouting: "NFT_1PANEL_POSTROUTING", - forwarding.ChainForward: "NFT_1PANEL_FORWARD", - } - for logical, want := range wantChains { - if got := nftForwardChain(logical); got != want { - t.Fatalf("nftForwardChain(%q) = %q, want %q", logical, got, want) - } - } -} - -func TestNormalizeNftForwardRuleRejectsScriptTokens(t *testing.T) { - base := forwarding.Rule{Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80"} - tests := []struct { - name string - mutate func(*forwarding.Rule) - }{ - {name: "source port", mutate: func(rule *forwarding.Rule) { rule.Port = "8080\nflush ruleset" }}, - {name: "target address", mutate: func(rule *forwarding.Rule) { rule.TargetIP = "10.0.0.2; flush ruleset" }}, - {name: "target port", mutate: func(rule *forwarding.Rule) { rule.TargetPort = "80; flush ruleset" }}, - {name: "interface", mutate: func(rule *forwarding.Rule) { rule.Interface = "eth0;flush" }}, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - rule := base - test.mutate(&rule) - if _, err := NormalizeRule(rule); err == nil { - t.Fatal("expected invalid nftables forwarding rule") - } - }) - } -} - -func TestNormalizeForwardRuleTreatsWildcardInterfaceAsAll(t *testing.T) { - rule, err := NormalizeRule(forwarding.Rule{ - Protocol: "udp", Port: "9000-9001", TargetPort: "53", Interface: "*", - }) - if err != nil { - t.Fatalf("normalize wildcard forwarding interface: %v", err) - } - if rule.Interface != "" { - t.Fatalf("wildcard forwarding interface normalized to %q, want empty", rule.Interface) - } -} - -func TestRebuildNftForwardCommandsUseArguments(t *testing.T) { - commands, err := rebuildNftForwardCommands([]forwarding.Rule{{ - Protocol: "tcp", Port: "08080", TargetIP: "10.0.0.2", TargetPort: "080", Interface: "eth0", - }}) - if err != nil { - t.Fatalf("build commands: %v", err) - } - if len(commands) != 10 { - t.Fatalf("unexpected command count: %d", len(commands)) - } - joined := strings.Join(commands[6], " ") - if !strings.Contains(joined, "add rule ip nft_1panel_forward NFT_1PANEL_PREROUTING") || - !strings.Contains(joined, `iifname "eth0"`) || !strings.Contains(joined, "dport 8080") || - !strings.Contains(joined, "dnat to 10.0.0.2:80") { - t.Fatalf("unexpected prerouting command: %q", joined) - } -} - -func TestNftForwardCommandsBuildSingleBatchScript(t *testing.T) { - commands, err := rebuildNftForwardCommands([]forwarding.Rule{{ - Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80", - }}) - if err != nil { - t.Fatal(err) - } - script, err := nftCommandsScript(commands) - if err != nil { - t.Fatal(err) - } - if lines := strings.Count(script, "\n"); lines != len(commands) { - t.Fatalf("batch script has %d lines, want %d:\n%s", lines, len(commands), script) - } - if !strings.Contains(script, "flush chain ip nft_1panel_forward NFT_1PANEL_PREROUTING\n") || - !strings.Contains(script, "add rule ip nft_1panel_forward NFT_1PANEL_PREROUTING meta l4proto tcp tcp dport 8080 redirect to :80") { - t.Fatalf("unexpected nftables batch script:\n%s", script) - } - if _, err := nftCommandsScript([][]string{{"add", "rule\nflush ruleset"}}); err == nil { - t.Fatal("batch script accepted a newline token") - } -} - -func TestRebuildNftIPv6ForwardCommands(t *testing.T) { - commands, err := rebuildNftForwardCommands([]forwarding.Rule{{ - Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "8443", TargetIP: "2001:db8::20", TargetPort: "443", Interface: "eth0", - }}) - if err != nil { - t.Fatalf("build IPv6 commands: %v", err) - } - if len(commands) != 10 { - t.Fatalf("unexpected command count: %d", len(commands)) - } - preRouting := strings.Join(commands[6], " ") - if !strings.Contains(preRouting, "add rule ip6 nft_1panel_forward NFT_1PANEL_PREROUTING") || - !strings.Contains(preRouting, "dnat to [2001:db8::20]:443") { - t.Fatalf("unexpected IPv6 prerouting command: %q", preRouting) - } - forward := strings.Join(commands[8], " ") - if !strings.Contains(forward, "ip6 daddr 2001:db8::20") { - t.Fatalf("unexpected IPv6 forward command: %q", forward) - } -} - -func TestNormalizeForwardRuleEnforcesAddressFamily(t *testing.T) { - ipv6, err := NormalizeRule(forwarding.Rule{Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "8080", TargetIP: "2001:db8::2", TargetPort: "80"}) - if err != nil || ipv6.TargetIP != "2001:db8::2" { - t.Fatalf("normalize IPv6 rule = %#v, %v", ipv6, err) - } - if _, err := NormalizeRule(forwarding.Rule{Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80"}); err == nil { - t.Fatal("expected an IPv4 target to be rejected for an IPv6 rule") - } - loopback, err := NormalizeRule(forwarding.Rule{Family: forwarding.FamilyIPv6, Protocol: "udp", Port: "5353", TargetPort: "53"}) - if err != nil || loopback.TargetIP != "::1" { - t.Fatalf("IPv6 loopback normalization = %#v, %v", loopback, err) - } -} diff --git a/agent/utils/firewall/forwarding/providers/normalize.go b/agent/utils/firewall/forwarding/providers/normalize.go deleted file mode 100644 index e1d3f308af8f..000000000000 --- a/agent/utils/firewall/forwarding/providers/normalize.go +++ /dev/null @@ -1,82 +0,0 @@ -package providers - -import ( - "fmt" - "net/netip" - "regexp" - "strconv" - "strings" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" -) - -var forwardInterfacePattern = regexp.MustCompile(`^[A-Za-z0-9_.:@-]{1,15}$`) - -func NormalizeRule(rule forwarding.Rule) (forwarding.Rule, error) { - rule.Family = strings.ToLower(strings.TrimSpace(rule.Family)) - if rule.Family == "" { - rule.Family = forwarding.FamilyIPv4 - } - if rule.Family != forwarding.FamilyIPv4 && rule.Family != forwarding.FamilyIPv6 { - return forwarding.Rule{}, fmt.Errorf("unsupported forwarding family %q", rule.Family) - } - rule.Protocol = strings.ToLower(strings.TrimSpace(rule.Protocol)) - if rule.Protocol != "tcp" && rule.Protocol != "udp" { - return forwarding.Rule{}, fmt.Errorf("unsupported forwarding protocol %q", rule.Protocol) - } - var err error - if rule.Port, err = normalizeForwardPort(rule.Port); err != nil { - return forwarding.Rule{}, fmt.Errorf("invalid forwarding port: %w", err) - } - if rule.TargetPort, err = normalizeForwardPort(rule.TargetPort); err != nil { - return forwarding.Rule{}, fmt.Errorf("invalid forwarding target port: %w", err) - } - rule.TargetIP = strings.TrimSpace(rule.TargetIP) - if rule.TargetIP == "" || strings.EqualFold(rule.TargetIP, "localhost") { - if rule.Family == forwarding.FamilyIPv6 { - rule.TargetIP = "::1" - } else { - rule.TargetIP = "127.0.0.1" - } - } - address, err := netip.ParseAddr(rule.TargetIP) - if err == nil { - address = address.Unmap() - } - if err != nil || (rule.Family == forwarding.FamilyIPv4) != address.Is4() { - return forwarding.Rule{}, fmt.Errorf("invalid %s forwarding target %q", rule.Family, rule.TargetIP) - } - rule.TargetIP = address.String() - rule.Interface = strings.TrimSpace(rule.Interface) - if rule.Interface == "all" || rule.Interface == "*" { - rule.Interface = "" - } - if rule.Interface != "" && !forwardInterfacePattern.MatchString(rule.Interface) { - return forwarding.Rule{}, fmt.Errorf("invalid forwarding interface %q", rule.Interface) - } - return rule, nil -} - -func normalizeForwardPort(value string) (string, error) { - parts := strings.Split(strings.TrimSpace(value), "-") - if len(parts) < 1 || len(parts) > 2 { - return "", fmt.Errorf("invalid port range %q", value) - } - ports := make([]int, len(parts)) - for index, part := range parts { - port, err := strconv.Atoi(strings.TrimSpace(part)) - if err != nil || port < 1 || port > 65535 { - return "", fmt.Errorf("invalid port %q", part) - } - ports[index] = port - } - if len(ports) == 2 { - if ports[0] > ports[1] { - return "", fmt.Errorf("descending port range %q", value) - } - if ports[0] != ports[1] { - return strconv.Itoa(ports[0]) + "-" + strconv.Itoa(ports[1]), nil - } - } - return strconv.Itoa(ports[0]), nil -} diff --git a/agent/utils/firewall/iptables_helper/inspect_test.go b/agent/utils/firewall/iptables_helper/inspect_test.go deleted file mode 100644 index cd8cf43744ad..000000000000 --- a/agent/utils/firewall/iptables_helper/inspect_test.go +++ /dev/null @@ -1,12 +0,0 @@ -package iptables_helper - -import "testing" - -func TestHasBaseChainBinding(t *testing.T) { - if hasBaseChainBinding("-P INPUT ACCEPT") { - t.Fatal("input policy was treated as a 1Panel binding") - } - if !hasBaseChainBinding("-A INPUT -j " + BasicAfterChain) { - t.Fatal("partial iptables binding was not detected") - } -} diff --git a/agent/utils/firewall/iptables_helper/manager_restore_test.go b/agent/utils/firewall/iptables_helper/manager_restore_test.go deleted file mode 100644 index 09933e5c8e12..000000000000 --- a/agent/utils/firewall/iptables_helper/manager_restore_test.go +++ /dev/null @@ -1,214 +0,0 @@ -package iptables_helper - -import ( - "errors" - "os" - "path/filepath" - "strings" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall" -) - -func TestBuildBaseChainsRestoreScriptBatchesPersistedRules(t *testing.T) { - dir := t.TempDir() - files := map[string]string{ - BasicBeforeFileName: "-A " + BasicBeforeChain + " -i lo -j ACCEPT\n", - BasicFileName: "-A " + BasicChain + " -p tcp --dport 8080 -j ACCEPT\n", - BasicAfterFileName: "-A " + BasicAfterChain + " -p tcp -j DROP\n", - } - for name, content := range files { - if err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o600); err != nil { - t.Fatal(err) - } - } - script, err := buildBaseChainsRestoreScript(dir, "9443", false) - if err != nil { - t.Fatal(err) - } - for _, expected := range []string{ - "*filter\n", - "-F " + BasicBeforeChain + "\n", - "-F " + BasicChain + "\n", - "-F " + BasicAfterChain + "\n", - files[BasicBeforeFileName], - files[BasicFileName], - files[BasicAfterFileName], - "-A " + BasicBeforeChain + " -p tcp -m tcp --dport 9443 -j ACCEPT\n", - "COMMIT\n", - } { - if !strings.Contains(script, expected) { - t.Fatalf("restore script does not contain %q:\n%s", expected, script) - } - } - if strings.Count(script, "COMMIT\n") != 1 { - t.Fatalf("restore script is not a single batch:\n%s", script) - } -} - -func TestBuildBaseChainsRestoreScriptUsesIPv6Files(t *testing.T) { - dir := t.TempDir() - want := "-A " + BasicChain + " -p ipv6-icmp -j ACCEPT\n" - if err := os.WriteFile(filepath.Join(dir, IPv6FileName(BasicFileName)), []byte(want), 0o600); err != nil { - t.Fatal(err) - } - script, err := buildBaseChainsRestoreScript(dir, "9443", true) - if err != nil { - t.Fatal(err) - } - if !strings.Contains(script, want) { - t.Fatalf("IPv6 persisted rule was not restored:\n%s", script) - } -} - -func TestBuildRequiredPortsRestoreScriptBatchesAddsAndDeletes(t *testing.T) { - desired := []firewall.PortWhitelist{{Protocol: "tcp", Port: "22"}, {Protocol: "udp", Port: "53"}} - before := []FilterRules{ - {Chain: BasicBeforeChain, Protocol: "tcp", DstPort: "22", Strategy: "accept"}, - {Chain: BasicBeforeChain, Protocol: "tcp", DstPort: "22", Strategy: "accept"}, - {Chain: BasicBeforeChain, Protocol: "tcp", DstPort: "80", Strategy: "accept"}, - } - after := []FilterRules{{Chain: BasicAfterChain, Protocol: "udp", DstPort: "5353", Strategy: "accept"}} - script := buildRequiredPortsRestoreScript(desired, before, after, "", "", true) - - for _, line := range []string{ - "-D 1PANEL_BASIC_BEFORE -p tcp -m tcp --dport 22 -j ACCEPT", - "-D 1PANEL_BASIC_BEFORE -p tcp -m tcp --dport 80 -j ACCEPT", - "-D 1PANEL_BASIC_AFTER -p udp -m udp --dport 5353 -j ACCEPT", - "-A 1PANEL_BASIC_BEFORE -p udp -m udp --dport 53 -j ACCEPT", - "-A 1PANEL_BASIC_AFTER -p tcp -j DROP", - "-A 1PANEL_BASIC_AFTER -p udp -j DROP", - } { - if !strings.Contains(script, line+"\n") { - t.Fatalf("batch script is missing %q:\n%s", line, script) - } - } - if got := strings.Count(script, "--dport 22 -j ACCEPT"); got != 1 { - t.Fatalf("duplicate desired port was not removed exactly once: count=%d\n%s", got, script) - } - if !strings.HasPrefix(script, "*filter\n") || !strings.HasSuffix(script, "COMMIT\n") { - t.Fatalf("invalid restore transaction:\n%s", script) - } -} - -func TestBuildRequiredPortsRestoreScriptSkipsUnchangedState(t *testing.T) { - desired := []firewall.PortWhitelist{{Protocol: "tcp", Port: "22"}} - before := []FilterRules{{Chain: BasicBeforeChain, Protocol: "tcp", DstPort: "22", Strategy: "accept"}} - if script := buildRequiredPortsRestoreScript(desired, before, nil, "", "", false); script != "" { - t.Fatalf("unchanged required ports generated a restore transaction:\n%s", script) - } -} - -func TestBuildBaseChainBindingsRestoreScriptRebindsInOneTransaction(t *testing.T) { - output := strings.Join([]string{ - "-A INPUT -j " + BasicBeforeChain, - "-A INPUT -j external", - "-A INPUT -j " + BasicAfterChain, - }, "\n") - script := buildBaseChainBindingsRestoreScript(output, true) - for _, line := range []string{ - "-D INPUT -j " + BasicBeforeChain, - "-D INPUT -j " + BasicAfterChain, - "-I INPUT 1 -j " + BasicBeforeChain, - "-I INPUT 2 -j " + BasicChain, - "-I INPUT 3 -j " + BasicAfterChain, - } { - if !strings.Contains(script, line+"\n") { - t.Fatalf("binding batch is missing %q:\n%s", line, script) - } - } - if strings.Contains(script, "external") || strings.Count(script, "COMMIT\n") != 1 { - t.Fatalf("binding batch modified an external rule or is not atomic:\n%s", script) - } -} - -func TestLoadInitStatusUsesIPv6BaselineWithoutIPv4TerminalRules(t *testing.T) { - output := strings.Join([]string{ - "-N " + BasicBeforeChain, - "-N " + BasicChain, - "-N " + BasicAfterChain, - "-A " + BasicBeforeChain + " -i lo -m comment --comment \"Loopback Whitelist\" -j ACCEPT", - "-A " + BasicBeforeChain + " -m conntrack --ctstate RELATED,ESTABLISHED -m comment --comment \"ESTABLISHED Whitelist\" -j ACCEPT", - "-A " + InputChain + " -j " + BasicBeforeChain, - "-A " + InputChain + " -j " + BasicChain, - "-A " + InputChain + " -j " + BasicAfterChain, - }, "\n") - runner := func(string, ...string) (string, error) { return output, nil } - initialized, bound, err := loadInitStatus("base", runner, false) - if err != nil || !initialized || !bound { - t.Fatalf("IPv6 baseline status = initialized:%v bound:%v err:%v", initialized, bound, err) - } - initialized, bound, err = loadInitStatus("base", runner, true) - if err != nil || initialized || bound { - t.Fatalf("IPv4 status ignored missing terminal rules: initialized:%v bound:%v err:%v", initialized, bound, err) - } -} - -func TestRepairIPv6BaseChainsOnlyBindsInitializedChains(t *testing.T) { - bindCalls, ensureCalls := 0, 0 - err := repairIPv6BaseChains(true, false, func() error { - bindCalls++ - return nil - }, func() error { - ensureCalls++ - return nil - }) - if err != nil { - t.Fatalf("repair initialized IPv6 base chains: %v", err) - } - if bindCalls != 1 || ensureCalls != 0 { - t.Fatalf("repair calls = bind:%d ensure:%d, want bind:1 ensure:0", bindCalls, ensureCalls) - } -} - -func TestRepairIPv6BaseChainsRebuildsMissingChains(t *testing.T) { - bindCalls, ensureCalls := 0, 0 - err := repairIPv6BaseChains(false, false, func() error { - bindCalls++ - return nil - }, func() error { - ensureCalls++ - return nil - }) - if err != nil { - t.Fatalf("repair missing IPv6 base chains: %v", err) - } - if bindCalls != 0 || ensureCalls != 1 { - t.Fatalf("repair calls = bind:%d ensure:%d, want bind:0 ensure:1", bindCalls, ensureCalls) - } -} - -func TestRepairIPv6BaseChainsLeavesHealthyChainsUnchanged(t *testing.T) { - bindCalls, ensureCalls := 0, 0 - err := repairIPv6BaseChains(true, true, func() error { - bindCalls++ - return nil - }, func() error { - ensureCalls++ - return nil - }) - if err != nil { - t.Fatalf("repair healthy IPv6 base chains: %v", err) - } - if bindCalls != 0 || ensureCalls != 0 { - t.Fatalf("repair calls = bind:%d ensure:%d, want no operation", bindCalls, ensureCalls) - } -} - -func TestRepairIPv6BaseChainsPropagatesOperationErrors(t *testing.T) { - wantBindErr := errors.New("bind failed") - if err := repairIPv6BaseChains(true, false, func() error { return wantBindErr }, func() error { return nil }); !errors.Is(err, wantBindErr) { - t.Fatalf("bind error = %v, want %v", err, wantBindErr) - } - wantEnsureErr := errors.New("ensure failed") - if err := repairIPv6BaseChains(false, false, func() error { return nil }, func() error { return wantEnsureErr }); !errors.Is(err, wantEnsureErr) { - t.Fatalf("ensure error = %v, want %v", err, wantEnsureErr) - } -} - -func TestLoadFamilyInitStatusRejectsUnknownFamily(t *testing.T) { - initialized, bound, err := LoadFamilyInitStatus("inet", "base") - if err == nil || initialized || bound { - t.Fatalf("unknown family status = initialized:%v bound:%v err:%v", initialized, bound, err) - } -} diff --git a/agent/utils/firewall/lifecycle/client.go b/agent/utils/firewall/lifecycle/client.go deleted file mode 100644 index 45f340c1f6c7..000000000000 --- a/agent/utils/firewall/lifecycle/client.go +++ /dev/null @@ -1,105 +0,0 @@ -package lifecycle - -import ( - "fmt" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle/providers" -) - -type Client interface { - Name() string - Start() error - Stop() error - Restart() error - Status() (bool, error) - Version() (string, error) -} - -// Resetter restores a service-backed firewall to its installation defaults. -// Implementations must leave the firewall disabled after a successful reset. -type Resetter interface { - Reset() error -} - -func NewClient() (Client, error) { - runtime, err := DetectRuntime() - if err != nil { - return nil, err - } - return NewClientFor(runtime.Provider) -} - -func NewClientFor(provider string) (Client, error) { - switch provider { - case "firewalld": - if !which("firewalld") { - return nil, fmt.Errorf("firewalld is not installed") - } - return providers.NewFirewalld() - case "ufw": - if !which("ufw") { - return nil, fmt.Errorf("ufw is not installed") - } - return providers.NewUFW() - case "iptables": - commands, err := ResolveIptablesCommands() - if err != nil { - return nil, err - } - return providers.NewIptables(commands.IPv4) - case "nftables": - if !which("nft") { - return nil, fmt.Errorf("nftables is not installed") - } - return providers.NewNftables() - default: - return nil, fmt.Errorf("unsupported firewall provider: %s", provider) - } -} - -func InstalledProviders() []string { - providers := make([]string, 0, 4) - if which("firewalld") { - providers = append(providers, ProviderFirewalld) - } - if which("ufw") { - providers = append(providers, ProviderUFW) - } - if _, err := ResolveIptablesCommands(); err == nil { - providers = append(providers, ProviderIptables) - } - if which("nft") { - providers = append(providers, ProviderNftables) - } - return providers -} - -func NewNetfilterClients() ([]Client, error) { - clients := make([]Client, 0, 2) - if which("nft") { - client, err := providers.NewNftables() - if err != nil { - return nil, err - } - clients = append(clients, client) - } - if commands, err := ResolveIptablesCommands(); err == nil { - client, err := providers.NewIptables(commands.IPv4) - if err != nil { - return nil, err - } - clients = append(clients, client) - } - if len(clients) == 0 { - return nil, fmt.Errorf("no supported forwarding backend detected (iptables/iptables-nft/nftables)") - } - return clients, nil -} - -func DetectProvider() (string, error) { - runtime, err := DetectRuntime() - if err != nil { - return "", err - } - return runtime.Provider, nil -} diff --git a/agent/utils/firewall/lifecycle/client_test.go b/agent/utils/firewall/lifecycle/client_test.go deleted file mode 100644 index 0d7306d10fe1..000000000000 --- a/agent/utils/firewall/lifecycle/client_test.go +++ /dev/null @@ -1,61 +0,0 @@ -package lifecycle - -import "testing" - -func TestNewNetfilterClientsIgnoreHostFirewallService(t *testing.T) { - original := which - t.Cleanup(func() { which = original }) - - tests := []struct { - name string - commands map[string]bool - want []string - wantErr bool - }{ - { - name: "firewalld host with iptables", - commands: map[string]bool{ - "firewalld": true, "iptables": true, "iptables-restore": true, - }, - want: []string{ProviderIptables}, - }, - { - name: "ufw host with iptables nft", - commands: map[string]bool{ - "ufw": true, "iptables-nft": true, "iptables-nft-restore": true, "nft": true, - }, - want: []string{ProviderNftables, ProviderIptables}, - }, - { - name: "firewalld host with native nft", - commands: map[string]bool{"firewalld": true, "nft": true}, - want: []string{ProviderNftables}, - }, - { - name: "service without netfilter command backend", - commands: map[string]bool{"firewalld": true}, - wantErr: true, - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - which = func(name string) bool { return test.commands[name] } - clients, err := NewNetfilterClients() - if (err != nil) != test.wantErr { - t.Fatalf("NewNetfilterClients() error = %v, wantErr %v", err, test.wantErr) - } - if test.wantErr { - return - } - if len(clients) != len(test.want) { - t.Fatalf("NewNetfilterClients() returned %d clients, want %d", len(clients), len(test.want)) - } - for index, client := range clients { - if client.Name() != test.want[index] { - t.Fatalf("NewNetfilterClients()[%d] = %q, want %q", index, client.Name(), test.want[index]) - } - } - }) - } -} diff --git a/agent/utils/firewall/lifecycle/lifecycle.go b/agent/utils/firewall/lifecycle/lifecycle.go new file mode 100644 index 000000000000..4637b03363d1 --- /dev/null +++ b/agent/utils/firewall/lifecycle/lifecycle.go @@ -0,0 +1,230 @@ +package lifecycle + +import ( + "errors" + "fmt" + "sync" + + "github.com/1Panel-dev/1Panel/agent/constant" + "github.com/1Panel-dev/1Panel/agent/utils/cmd" + "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle/providers" +) + +const ( + ProviderFirewalld = constant.FirewallProviderFirewalld + ProviderUFW = constant.FirewallProviderUFW + ProviderIptables = constant.FirewallProviderIptables + ProviderNftables = constant.FirewallProviderNftables +) + +var ErrNotInstalled = errors.New("is not installed") + +type IptablesCommands struct { + IPv4 string + IPv6 string + Restore4 string + Restore6 string +} + +func (c IptablesCommands) IPv6Available() bool { + return c.IPv6 != "" && c.Restore6 != "" +} + +type Runtime struct { + Provider string + Iptables IptablesCommands +} + +var which = cmd.Which + +func DetectRuntime() (Runtime, error) { + hasFirewalld := which("firewalld") + hasUFW := which("ufw") + if hasFirewalld && hasUFW { + return Runtime{}, errors.New("it is detected that the system has both firewalld and ufw services. To avoid conflicts, please uninstall and try again") + } + if hasFirewalld { + return Runtime{Provider: ProviderFirewalld}, nil + } + if hasUFW { + return Runtime{Provider: ProviderUFW}, nil + } + if commands, ok := detectIptablesCommands(""); ok { + return Runtime{Provider: ProviderIptables, Iptables: commands}, nil + } + if commands, ok := detectIptablesCommands("-nft"); ok { + return Runtime{Provider: ProviderIptables, Iptables: commands}, nil + } + if which("nft") { + return Runtime{Provider: ProviderNftables}, nil + } + return Runtime{}, errors.New("no system firewall service detected (firewalld/ufw/iptables/iptables-nft/nft), please check and try again") +} + +func detectIptablesCommands(suffix string) (IptablesCommands, bool) { + ipv4 := "iptables" + suffix + restore4 := "iptables" + suffix + "-restore" + if !which(ipv4) || !which(restore4) { + return IptablesCommands{}, false + } + commands := IptablesCommands{IPv4: ipv4, Restore4: restore4} + ipv6 := "ip6tables" + suffix + restore6 := "ip6tables" + suffix + "-restore" + if which(ipv6) && which(restore6) { + commands.IPv6 = ipv6 + commands.Restore6 = restore6 + } + return commands, true +} + +func ResolveIptablesCommands() (IptablesCommands, error) { + if commands, ok := detectIptablesCommands(""); ok { + return commands, nil + } + if commands, ok := detectIptablesCommands("-nft"); ok { + return commands, nil + } + return IptablesCommands{}, fmt.Errorf("no complete iptables command family is available") +} + +type Client interface { + Name() string + Start() error + Stop() error + Restart() error + Status() (bool, error) + Version() (string, error) +} + +// Resetter restores a service-backed firewall to its installation defaults. +// Implementations must leave the firewall disabled after a successful reset. +type Resetter interface { + Reset() error +} + +func NewClient() (Client, error) { + runtime, err := DetectRuntime() + if err != nil { + return nil, err + } + return NewClientFor(runtime.Provider) +} + +func NewClientFor(provider string) (Client, error) { + switch provider { + case "firewalld": + if !which("firewalld") { + return nil, fmt.Errorf("firewalld %w", ErrNotInstalled) + } + return providers.NewFirewalld() + case "ufw": + if !which("ufw") { + return nil, fmt.Errorf("ufw %w", ErrNotInstalled) + } + return providers.NewUFW() + case "iptables": + commands, err := ResolveIptablesCommands() + if err != nil { + return nil, err + } + return providers.NewIptables(commands.IPv4) + case "nftables": + if !which("nft") { + return nil, fmt.Errorf("nftables %w", ErrNotInstalled) + } + return providers.NewNftables() + default: + return nil, fmt.Errorf("unsupported firewall provider: %s", provider) + } +} + +func InstalledProviders() []string { + providers := make([]string, 0, 4) + if which("firewalld") { + providers = append(providers, ProviderFirewalld) + } + if which("ufw") { + providers = append(providers, ProviderUFW) + } + if _, err := ResolveIptablesCommands(); err == nil { + providers = append(providers, ProviderIptables) + } + if which("nft") { + providers = append(providers, ProviderNftables) + } + return providers +} + +func NewNetfilterClients() ([]Client, error) { + clients := make([]Client, 0, 2) + if which("nft") { + client, err := providers.NewNftables() + if err != nil { + return nil, err + } + clients = append(clients, client) + } + if commands, err := ResolveIptablesCommands(); err == nil { + client, err := providers.NewIptables(commands.IPv4) + if err != nil { + return nil, err + } + clients = append(clients, client) + } + if len(clients) == 0 { + return nil, fmt.Errorf("no supported forwarding backend detected (iptables/iptables-nft/nftables)") + } + return clients, nil +} + +func DetectProvider() (string, error) { + runtime, err := DetectRuntime() + if err != nil { + return "", err + } + return runtime.Provider, nil +} + +type State struct { + Name string + IsActive bool +} + +type Status struct { + State + Version string +} + +func LoadState(client Client) (State, error) { + state := State{Name: client.Name()} + var err error + state.IsActive, err = client.Status() + return state, err +} + +func LoadStatus(client Client) (Status, error) { + status := Status{Version: "-"} + var state State + var version string + var stateErr, versionErr error + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + state, stateErr = LoadState(client) + }() + go func() { + defer wg.Done() + version, versionErr = client.Version() + }() + wg.Wait() + status.State = state + if stateErr != nil { + return status, errors.Join(stateErr, versionErr) + } + if !status.IsActive { + return status, nil + } + status.Version = version + return status, versionErr +} diff --git a/agent/utils/firewall/lifecycle/operator_test.go b/agent/utils/firewall/lifecycle/operator_test.go deleted file mode 100644 index a9517c3fbc06..000000000000 --- a/agent/utils/firewall/lifecycle/operator_test.go +++ /dev/null @@ -1,48 +0,0 @@ -package lifecycle - -import ( - "errors" - "testing" -) - -type operatorTestClient struct { - name string - started bool - stopped bool -} - -func (f *operatorTestClient) Name() string { return f.name } -func (f *operatorTestClient) Start() error { f.started = true; return nil } -func (f *operatorTestClient) Stop() error { f.stopped = true; return nil } -func (f *operatorTestClient) Restart() error { return nil } -func (f *operatorTestClient) Status() (bool, error) { return true, nil } -func (f *operatorTestClient) Version() (string, error) { return "test", nil } -func TestOperatorDelegatesLifecycleAndPreparesStart(t *testing.T) { - client := &operatorTestClient{name: "iptables"} - operator := NewOperator(client) - prepared := false - if err := operator.Operate("start", false, func(got Client) error { - prepared = got == client - return nil - }); err != nil { - t.Fatal(err) - } - if !client.started || !prepared { - t.Fatalf("start was not fully coordinated: started=%v prepared=%v", client.started, prepared) - } -} - -func TestOperatorKeepsFirewallRunningWhenPostStartPreparationFails(t *testing.T) { - client := &operatorTestClient{name: "firewalld"} - wantErr := errors.New("sync accepted ports") - - err := NewOperator(client).Operate(OperationStart, false, func(Client) error { - return wantErr - }) - if !errors.Is(err, wantErr) { - t.Fatalf("start returned error %v, want %v", err, wantErr) - } - if !client.started || client.stopped { - t.Fatalf("post-start failure changed firewall state: started=%v stopped=%v", client.started, client.stopped) - } -} diff --git a/agent/utils/firewall/lifecycle/provider_test.go b/agent/utils/firewall/lifecycle/provider_test.go deleted file mode 100644 index 09a34566324c..000000000000 --- a/agent/utils/firewall/lifecycle/provider_test.go +++ /dev/null @@ -1,43 +0,0 @@ -package lifecycle - -import ( - "os" - "path/filepath" - "testing" -) - -func TestDetectProvider(t *testing.T) { - tests := []struct { - name string - executables []string - want string - wantErr bool - }{ - {name: "none", wantErr: true}, - {name: "iptables", executables: []string{"iptables", "iptables-restore"}, want: "iptables"}, - {name: "iptables-nft", executables: []string{"iptables-nft", "iptables-nft-restore"}, want: "iptables"}, - {name: "nftables", executables: []string{"nft"}, want: "nftables"}, - {name: "ufw", executables: []string{"iptables", "iptables-restore", "ufw"}, want: "ufw"}, - {name: "firewalld", executables: []string{"iptables", "iptables-restore", "firewalld"}, want: "firewalld"}, - {name: "conflict", executables: []string{"firewalld", "ufw"}, wantErr: true}, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - directory := t.TempDir() - for _, executable := range test.executables { - name := filepath.Join(directory, executable) - if err := os.WriteFile(name, []byte("#!/bin/sh\n"), 0755); err != nil { - t.Fatal(err) - } - } - t.Setenv("PATH", directory) - got, err := DetectProvider() - if (err != nil) != test.wantErr { - t.Fatalf("DetectProvider() error = %v, wantErr %v", err, test.wantErr) - } - if got != test.want { - t.Fatalf("DetectProvider() = %q, want %q", got, test.want) - } - }) - } -} diff --git a/agent/utils/firewall/lifecycle/providers/firewalld_test.go b/agent/utils/firewall/lifecycle/providers/firewalld_test.go deleted file mode 100644 index 994e29f0bd39..000000000000 --- a/agent/utils/firewall/lifecycle/providers/firewalld_test.go +++ /dev/null @@ -1,107 +0,0 @@ -package providers - -import ( - "errors" - "os" - "path/filepath" - "testing" -) - -func TestFirewalldStoppedRecognizesNormalInactiveResult(t *testing.T) { - for _, test := range []struct { - stdout string - err error - }{ - {stdout: "not running\n"}, - {err: errors.New("stderr: FirewallD is not running, exit status 252")}, - } { - if !firewalldStopped(test.stdout, test.err) { - t.Fatalf("expected stopped result for stdout=%q err=%v", test.stdout, test.err) - } - } -} - -func TestFirewalldStoppedKeepsUnexpectedFailures(t *testing.T) { - if firewalldStopped("", errors.New("permission denied")) { - t.Fatal("unexpected command failures must not be treated as an inactive firewall") - } -} - -func TestReplaceFirewalldConfigCreatesCleanConfigurationAndBackup(t *testing.T) { - root := t.TempDir() - configDir := filepath.Join(root, "firewalld") - backupDir := filepath.Join(root, "firewalld.backup") - originalZone := filepath.Join(configDir, "zones", "custom.xml") - if err := os.MkdirAll(filepath.Dir(originalZone), 0750); err != nil { - t.Fatal(err) - } - if err := os.WriteFile(originalZone, []byte("custom"), 0600); err != nil { - t.Fatal(err) - } - - prepared := false - rollback, err := replaceFirewalldConfig( - configDir, - backupDir, - func(path string) error { - prepared = path == configDir - return nil - }, - func() error { - for _, directory := range firewalldConfigSubdirectories { - if info, err := os.Stat(filepath.Join(configDir, directory)); err != nil || !info.IsDir() { - t.Fatalf("expected clean %s directory, info=%v err=%v", directory, info, err) - } - } - if _, err := os.Stat(originalZone); !os.IsNotExist(err) { - t.Fatalf("custom zone must not remain in clean configuration: %v", err) - } - return nil - }, - ) - if err != nil { - t.Fatalf("replace firewalld configuration: %v", err) - } - if !prepared { - t.Fatal("expected clean configuration to be prepared") - } - if content, err := os.ReadFile(filepath.Join(backupDir, "zones", "custom.xml")); err != nil || string(content) != "custom" { - t.Fatalf("expected original configuration in backup, content=%q err=%v", content, err) - } - - if err := rollback(); err != nil { - t.Fatalf("rollback firewalld configuration: %v", err) - } - if content, err := os.ReadFile(originalZone); err != nil || string(content) != "custom" { - t.Fatalf("expected original configuration after rollback, content=%q err=%v", content, err) - } -} - -func TestReplaceFirewalldConfigRollsBackValidationFailure(t *testing.T) { - root := t.TempDir() - configDir := filepath.Join(root, "firewalld") - backupDir := filepath.Join(root, "firewalld.backup") - originalConfig := filepath.Join(configDir, "firewalld.conf") - if err := os.Mkdir(configDir, 0750); err != nil { - t.Fatal(err) - } - if err := os.WriteFile(originalConfig, []byte("DefaultZone=custom\n"), 0600); err != nil { - t.Fatal(err) - } - - rollback, err := replaceFirewalldConfig( - configDir, - backupDir, - nil, - func() error { return errors.New("invalid defaults") }, - ) - if err == nil || rollback != nil { - t.Fatalf("expected validation failure with automatic rollback, hasRollback=%t err=%v", rollback != nil, err) - } - if content, readErr := os.ReadFile(originalConfig); readErr != nil || string(content) != "DefaultZone=custom\n" { - t.Fatalf("expected original configuration after failed validation, content=%q err=%v", content, readErr) - } - if _, statErr := os.Stat(backupDir); !os.IsNotExist(statErr) { - t.Fatalf("backup must be restored after failed validation: %v", statErr) - } -} diff --git a/agent/utils/firewall/lifecycle/runtime.go b/agent/utils/firewall/lifecycle/runtime.go deleted file mode 100644 index caec9219695d..000000000000 --- a/agent/utils/firewall/lifecycle/runtime.go +++ /dev/null @@ -1,84 +0,0 @@ -package lifecycle - -import ( - "errors" - "fmt" - - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/utils/cmd" -) - -const ( - ProviderFirewalld = constant.FirewallProviderFirewalld - ProviderUFW = constant.FirewallProviderUFW - ProviderIptables = constant.FirewallProviderIptables - ProviderNftables = constant.FirewallProviderNftables -) - -type IptablesCommands struct { - IPv4 string - IPv6 string - Restore4 string - Restore6 string -} - -func (c IptablesCommands) IPv6Available() bool { - return c.IPv6 != "" && c.Restore6 != "" -} - -type Runtime struct { - Provider string - Iptables IptablesCommands -} - -var which = cmd.Which - -func DetectRuntime() (Runtime, error) { - hasFirewalld := which("firewalld") - hasUFW := which("ufw") - if hasFirewalld && hasUFW { - return Runtime{}, errors.New("it is detected that the system has both firewalld and ufw services. To avoid conflicts, please uninstall and try again") - } - if hasFirewalld { - return Runtime{Provider: ProviderFirewalld}, nil - } - if hasUFW { - return Runtime{Provider: ProviderUFW}, nil - } - if commands, ok := detectIptablesCommands(""); ok { - return Runtime{Provider: ProviderIptables, Iptables: commands}, nil - } - if commands, ok := detectIptablesCommands("-nft"); ok { - return Runtime{Provider: ProviderIptables, Iptables: commands}, nil - } - if which("nft") { - return Runtime{Provider: ProviderNftables}, nil - } - return Runtime{}, errors.New("no system firewall service detected (firewalld/ufw/iptables/iptables-nft/nft), please check and try again") -} - -func detectIptablesCommands(suffix string) (IptablesCommands, bool) { - ipv4 := "iptables" + suffix - restore4 := "iptables" + suffix + "-restore" - if !which(ipv4) || !which(restore4) { - return IptablesCommands{}, false - } - commands := IptablesCommands{IPv4: ipv4, Restore4: restore4} - ipv6 := "ip6tables" + suffix - restore6 := "ip6tables" + suffix + "-restore" - if which(ipv6) && which(restore6) { - commands.IPv6 = ipv6 - commands.Restore6 = restore6 - } - return commands, true -} - -func ResolveIptablesCommands() (IptablesCommands, error) { - if commands, ok := detectIptablesCommands(""); ok { - return commands, nil - } - if commands, ok := detectIptablesCommands("-nft"); ok { - return commands, nil - } - return IptablesCommands{}, fmt.Errorf("no complete iptables command family is available") -} diff --git a/agent/utils/firewall/lifecycle/runtime_test.go b/agent/utils/firewall/lifecycle/runtime_test.go deleted file mode 100644 index 49270d1f5022..000000000000 --- a/agent/utils/firewall/lifecycle/runtime_test.go +++ /dev/null @@ -1,41 +0,0 @@ -package lifecycle - -import "testing" - -func TestDetectRuntimePriority(t *testing.T) { - original := which - t.Cleanup(func() { which = original }) - - tests := []struct { - name string - commands map[string]bool - provider string - executable string - }{ - {name: "default iptables", commands: map[string]bool{"iptables": true, "iptables-restore": true, "iptables-nft": true, "iptables-nft-restore": true, "nft": true}, provider: ProviderIptables, executable: "iptables"}, - {name: "explicit iptables nft", commands: map[string]bool{"iptables-nft": true, "iptables-nft-restore": true, "nft": true}, provider: ProviderIptables, executable: "iptables-nft"}, - {name: "native nft", commands: map[string]bool{"nft": true}, provider: ProviderNftables}, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - which = func(name string) bool { return test.commands[name] } - runtime, err := DetectRuntime() - if err != nil { - t.Fatalf("detect: %v", err) - } - if runtime.Provider != test.provider || runtime.Iptables.IPv4 != test.executable { - t.Fatalf("unexpected runtime: %#v", runtime) - } - }) - } -} - -func TestDetectRuntimeRequiresRestoreCommand(t *testing.T) { - original := which - t.Cleanup(func() { which = original }) - which = func(name string) bool { return name == "iptables" || name == "nft" } - runtime, err := DetectRuntime() - if err != nil || runtime.Provider != ProviderNftables { - t.Fatalf("incomplete iptables family should fall back to nft: runtime=%#v err=%v", runtime, err) - } -} diff --git a/agent/utils/firewall/lifecycle/status.go b/agent/utils/firewall/lifecycle/status.go deleted file mode 100644 index b2c3eaeeb9d1..000000000000 --- a/agent/utils/firewall/lifecycle/status.go +++ /dev/null @@ -1,50 +0,0 @@ -package lifecycle - -import ( - "errors" - "sync" -) - -type State struct { - Name string - IsActive bool -} - -type Status struct { - State - Version string -} - -func LoadState(client Client) (State, error) { - state := State{Name: client.Name()} - var err error - state.IsActive, err = client.Status() - return state, err -} - -func LoadStatus(client Client) (Status, error) { - status := Status{Version: "-"} - var state State - var version string - var stateErr, versionErr error - var wg sync.WaitGroup - wg.Add(2) - go func() { - defer wg.Done() - state, stateErr = LoadState(client) - }() - go func() { - defer wg.Done() - version, versionErr = client.Version() - }() - wg.Wait() - status.State = state - if stateErr != nil { - return status, errors.Join(stateErr, versionErr) - } - if !status.IsActive { - return status, nil - } - status.Version = version - return status, versionErr -} diff --git a/agent/utils/firewall/lifecycle/status_test.go b/agent/utils/firewall/lifecycle/status_test.go deleted file mode 100644 index 534ae7508b99..000000000000 --- a/agent/utils/firewall/lifecycle/status_test.go +++ /dev/null @@ -1,55 +0,0 @@ -package lifecycle - -import ( - "errors" - "testing" -) - -type statusTestClient struct { - active bool - statusErr error - versionErr error - versioned *bool -} - -func (statusTestClient) Name() string { return "ufw" } -func (statusTestClient) Start() error { return nil } -func (statusTestClient) Stop() error { return nil } -func (statusTestClient) Restart() error { return nil } -func (c statusTestClient) Status() (bool, error) { return c.active, c.statusErr } -func (c statusTestClient) Version() (string, error) { - if c.versioned != nil { - *c.versioned = true - } - return "1.0", c.versionErr -} - -func TestLoadStatusAggregatesClientState(t *testing.T) { - status, err := LoadStatus(statusTestClient{active: true}) - if err != nil { - t.Fatal(err) - } - if status.Name != "ufw" || status.Version != "1.0" || !status.IsActive { - t.Fatalf("unexpected status: %#v", status) - } -} - -func TestLoadStatusReturnsClientErrors(t *testing.T) { - statusErr := errors.New("status failed") - versionErr := errors.New("version failed") - _, err := LoadStatus(statusTestClient{active: true, statusErr: statusErr, versionErr: versionErr}) - if !errors.Is(err, statusErr) || !errors.Is(err, versionErr) { - t.Fatalf("got %v, want joined status and version errors", err) - } -} - -func TestLoadStatusIgnoresVersionFailureForInactiveFirewall(t *testing.T) { - versioned := false - status, err := LoadStatus(statusTestClient{versionErr: errors.New("firewall is stopped"), versioned: &versioned}) - if err != nil { - t.Fatal(err) - } - if status.Name != "ufw" || status.Version != "-" || status.IsActive || !versioned { - t.Fatalf("unexpected inactive status: %#v versioned=%v", status, versioned) - } -} diff --git a/agent/utils/firewall/nftables_helper/command.go b/agent/utils/firewall/nftables_helper/command.go deleted file mode 100644 index c4075dec305c..000000000000 --- a/agent/utils/firewall/nftables_helper/command.go +++ /dev/null @@ -1,66 +0,0 @@ -package nftables_helper - -import ( - "fmt" - "strings" - "time" - - "github.com/1Panel-dev/1Panel/agent/utils/cmd" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -const ( - TableName = "nft_1panel_filter" - InputChain = "NFT_1PANEL_INPUT" - BasicBeforeChain = "NFT_1PANEL_BASIC_BEFORE" - BasicChain = "NFT_1PANEL_BASIC" - BasicAfterChain = "NFT_1PANEL_BASIC_AFTER" -) - -func TableFamily(family filter.Family) string { - if family == filter.FamilyIPv6 { - return "ip6" - } - return "ip" -} - -func BasicChains() []string { - return []string{BasicBeforeChain, BasicChain, BasicAfterChain} -} - -func run(args ...string) (string, error) { - return cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudoAndStdout("nft", args...) -} - -func runCommand(args ...string) error { - return cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudo("nft", args...) -} - -func runBatch(commands ...[]string) error { - script, err := buildBatchScript(commands...) - if err != nil || script == "" { - return err - } - manager := cmd.NewCommandMgr( - cmd.WithTimeout(60*time.Second), - cmd.WithStdin(strings.NewReader(script)), - ) - return manager.RunWithOptionalSudo("nft", "-f", "-") -} - -func buildBatchScript(commands ...[]string) (string, error) { - var script strings.Builder - for _, command := range commands { - if len(command) == 0 { - continue - } - for _, token := range command { - if strings.ContainsAny(token, "\r\n") { - return "", fmt.Errorf("invalid newline in nftables batch command") - } - } - script.WriteString(strings.Join(command, " ")) - script.WriteByte('\n') - } - return script.String(), nil -} diff --git a/agent/utils/firewall/nftables_helper/command_test.go b/agent/utils/firewall/nftables_helper/command_test.go deleted file mode 100644 index e31ddfab19a3..000000000000 --- a/agent/utils/firewall/nftables_helper/command_test.go +++ /dev/null @@ -1,51 +0,0 @@ -package nftables_helper - -import ( - "strings" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -func TestNativeNames(t *testing.T) { - tests := []struct { - family filter.Family - tableFamily string - }{ - {family: filter.FamilyIPv4, tableFamily: "ip"}, - {family: filter.FamilyIPv6, tableFamily: "ip6"}, - } - for _, test := range tests { - if got := TableFamily(test.family); got != test.tableFamily { - t.Fatalf("TableFamily(%s) = %q, want %q", test.family, got, test.tableFamily) - } - } - - wantChains := []string{"NFT_1PANEL_BASIC_BEFORE", "NFT_1PANEL_BASIC", "NFT_1PANEL_BASIC_AFTER"} - for index, got := range BasicChains() { - if got != wantChains[index] { - t.Fatalf("BasicChains()[%d] = %q, want %q", index, got, wantChains[index]) - } - } -} - -func TestBuildBatchScriptCombinesCommands(t *testing.T) { - script, err := buildBatchScript( - []string{"flush", "chain", "ip", TableName, BasicBeforeChain}, - []string{"add", "rule", "ip", TableName, BasicBeforeChain, "tcp", "dport", "443", "accept"}, - []string{"delete", "rule", "ip6", TableName, BasicBeforeChain, "handle", "12"}, - ) - if err != nil { - t.Fatal(err) - } - if strings.Count(script, "\n") != 3 || !strings.Contains(script, "tcp dport 443 accept\n") || - !strings.Contains(script, "delete rule ip6 "+TableName+" "+BasicBeforeChain+" handle 12\n") { - t.Fatalf("unexpected nftables batch script:\n%s", script) - } -} - -func TestBuildBatchScriptRejectsNewline(t *testing.T) { - if _, err := buildBatchScript([]string{"add", "rule", "ip", TableName, BasicChain, "unsafe\nrule"}); err == nil { - t.Fatal("expected newline validation error") - } -} diff --git a/agent/utils/firewall/nftables_helper/manager_test.go b/agent/utils/firewall/nftables_helper/manager_test.go deleted file mode 100644 index f1809a7c6950..000000000000 --- a/agent/utils/firewall/nftables_helper/manager_test.go +++ /dev/null @@ -1,45 +0,0 @@ -package nftables_helper - -import ( - "reflect" - "testing" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall" -) - -func TestHasBaseChainBinding(t *testing.T) { - if hasBaseChainBinding(`chain NFT_1PANEL_INPUT { policy accept; }`) { - t.Fatal("empty input chain was treated as bound") - } - if !hasBaseChainBinding(`jump NFT_1PANEL_BASIC`) { - t.Fatal("partial nftables binding was not detected") - } -} - -func TestRequiredPortChangesPreserveExistingRules(t *testing.T) { - existing := []requiredPortRule{ - {Key: "22/tcp", Handle: "12"}, - {Key: "22/tcp", Handle: "13"}, - {Key: "80/tcp", Handle: "14"}, - } - desired := []firewall.PortWhitelist{{Port: "22", Protocol: "tcp"}, {Port: "53", Protocol: "udp"}} - missing, stale := requiredPortChanges(existing, desired) - if want := []firewall.PortWhitelist{{Port: "53", Protocol: "udp"}}; !reflect.DeepEqual(missing, want) { - t.Fatalf("missing=%#v, want %#v", missing, want) - } - if want := []string{"13", "14"}; !reflect.DeepEqual(stale, want) { - t.Fatalf("stale=%#v, want %#v", stale, want) - } -} - -func TestRequiredPortRules(t *testing.T) { - output := ` - tcp dport 22 accept comment "1Panel Port Whitelist" # handle 12 - tcp dport 80 accept comment "external" # handle 13 - udp dport 53 accept comment "1Panel Port Whitelist" # handle 14 - ` - want := []requiredPortRule{{Key: "22/tcp", Handle: "12"}, {Key: "53/udp", Handle: "14"}} - if got := requiredPortRules(output); !reflect.DeepEqual(got, want) { - t.Fatalf("rules=%#v, want %#v", got, want) - } -} diff --git a/agent/utils/firewall/nftables_helper/inspect.go b/agent/utils/firewall/nftables_helper/runtime.go similarity index 53% rename from agent/utils/firewall/nftables_helper/inspect.go rename to agent/utils/firewall/nftables_helper/runtime.go index 3ce2c05fb7a0..42de25fd281e 100644 --- a/agent/utils/firewall/nftables_helper/inspect.go +++ b/agent/utils/firewall/nftables_helper/runtime.go @@ -1,11 +1,70 @@ package nftables_helper import ( + "fmt" "strings" + "time" + "github.com/1Panel-dev/1Panel/agent/utils/cmd" "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" ) +const ( + TableName = "nft_1panel_filter" + InputChain = "NFT_1PANEL_INPUT" + BasicBeforeChain = "NFT_1PANEL_BASIC_BEFORE" + BasicChain = "NFT_1PANEL_BASIC" + BasicAfterChain = "NFT_1PANEL_BASIC_AFTER" +) + +func TableFamily(family filter.Family) string { + if family == filter.FamilyIPv6 { + return "ip6" + } + return "ip" +} + +func BasicChains() []string { + return []string{BasicBeforeChain, BasicChain, BasicAfterChain} +} + +func run(args ...string) (string, error) { + return cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudoAndStdout("nft", args...) +} + +func runCommand(args ...string) error { + return cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudo("nft", args...) +} + +func runBatch(commands ...[]string) error { + script, err := buildBatchScript(commands...) + if err != nil || script == "" { + return err + } + manager := cmd.NewCommandMgr( + cmd.WithTimeout(60*time.Second), + cmd.WithStdin(strings.NewReader(script)), + ) + return manager.RunWithOptionalSudo("nft", "-f", "-") +} + +func buildBatchScript(commands ...[]string) (string, error) { + var script strings.Builder + for _, command := range commands { + if len(command) == 0 { + continue + } + for _, token := range command { + if strings.ContainsAny(token, "\r\n") { + return "", fmt.Errorf("invalid newline in nftables batch command") + } + } + script.WriteString(strings.Join(command, " ")) + script.WriteByte('\n') + } + return script.String(), nil +} + func LoadInitStatus(tab string) (bool, bool, error) { if tab != "base" { return false, false, nil diff --git a/agent/utils/firewall/port_whitelist.go b/agent/utils/firewall/port_whitelist.go index 1a60fd40eb66..60e70dc9b8e1 100644 --- a/agent/utils/firewall/port_whitelist.go +++ b/agent/utils/firewall/port_whitelist.go @@ -3,6 +3,7 @@ package firewall import ( "encoding/json" "fmt" + "sort" "strconv" "strings" @@ -193,3 +194,93 @@ func PortWhitelistKey(item PortWhitelist) string { } return key } + +type SystemPort struct { + Family string + Port string + Protocol string +} + +func RuleForSystemPort(provider filter.Provider, port SystemPort) filter.FirewallRule { + scope := filter.Scope{Provider: provider, Direction: filter.DirectionInput} + family := filter.Family(strings.ToLower(strings.TrimSpace(port.Family))) + switch provider { + case filter.ProviderIptables, filter.ProviderNftables: + if family != filter.FamilyIPv6 { + family = filter.FamilyIPv4 + } + scope.Family, scope.Table = family, "filter" + case filter.ProviderFirewalld: + if family != filter.FamilyIPv4 && family != filter.FamilyIPv6 { + family = filter.FamilyInet + } + scope.Family, scope.Zone = family, filter.FirewalldInputZone + case filter.ProviderUFW: + if family != filter.FamilyIPv6 { + family = filter.FamilyIPv4 + } + scope.Family = family + } + return filter.FirewallRule{ + Scope: scope, Protocol: port.Protocol, DestinationPort: port.Port, + Action: filter.ActionAccept, Description: "1Panel managed accepted port", + } +} + +func NormalizeSystemPorts(ports []SystemPort) (map[string]SystemPort, error) { + result := make(map[string]SystemPort, len(ports)) + for _, port := range ports { + normalized, err := filter.NormalizeRule(RuleForSystemPort(filter.ProviderIptables, port)) + if err != nil { + return nil, err + } + family := strings.ToLower(strings.TrimSpace(port.Family)) + if family != "" { + family = string(normalized.Scope.Family) + } + item := SystemPort{Family: family, Port: normalized.DestinationPort, Protocol: normalized.Protocol} + result[SystemPortKey(item)] = item + } + return result, nil +} + +func SystemPortKey(port SystemPort) string { + key := LegacySystemPortKey(port) + if family := strings.ToLower(strings.TrimSpace(port.Family)); family != "" { + return family + "/" + key + } + return key +} + +func LegacySystemPortKey(port SystemPort) string { + return strings.ToLower(strings.TrimSpace(port.Protocol)) + "/" + strings.TrimSpace(port.Port) +} + +func SortedSystemPortKeys(ports map[string]SystemPort) []string { + keys := make([]string, 0, len(ports)) + for key := range ports { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + +func ContainsPort(ports []PortWhitelist, target PortWhitelist) bool { + for _, port := range ports { + familyMatches := port.Family == "" || target.Family == "" || port.Family == target.Family + if familyMatches && port.Port == target.Port && port.Protocol == target.Protocol { + return true + } + } + return false +} + +func ExcludePorts(ports, excluded []PortWhitelist) []PortWhitelist { + result := make([]PortWhitelist, 0, len(ports)) + for _, port := range ports { + if !ContainsPort(excluded, port) { + result = append(result, port) + } + } + return result +} diff --git a/agent/utils/firewall/port_whitelist_test.go b/agent/utils/firewall/port_whitelist_test.go deleted file mode 100644 index d3d20651f01c..000000000000 --- a/agent/utils/firewall/port_whitelist_test.go +++ /dev/null @@ -1,72 +0,0 @@ -package firewall - -import "testing" - -func TestParsePortWhitelistLegacyAndStructured(t *testing.T) { - legacy, err := ParsePortWhitelist("80/tcp,53/udp") - if err != nil { - t.Fatalf("parse legacy whitelist: %v", err) - } - if len(legacy) != 2 || legacy[0] != (PortWhitelist{Family: "ipv4", Port: "80", Protocol: "tcp"}) { - t.Fatalf("unexpected legacy whitelist: %#v", legacy) - } - - structured, err := ParsePortWhitelist(`[ - {"family":"ipv6","protocol":"TCP","port":"8080:8090"}, - {"family":"ipv4","protocol":"udp","port":"53"} - ]`) - if err != nil { - t.Fatalf("parse structured whitelist: %v", err) - } - want := []PortWhitelist{ - {Family: "ipv6", Port: "8080-8090", Protocol: "tcp"}, - {Family: "ipv4", Port: "53", Protocol: "udp"}, - } - if len(structured) != len(want) { - t.Fatalf("unexpected structured whitelist: %#v", structured) - } - for index := range want { - if structured[index] != want[index] { - t.Fatalf("rule %d = %#v, want %#v", index, structured[index], want[index]) - } - } -} - -func TestParsePortWhitelistRejectsInvalidRule(t *testing.T) { - for _, value := range []string{ - `[{"family":"inet","protocol":"tcp","port":"80"}]`, - `[{"family":"ipv4","protocol":"icmp","port":"80"}]`, - `[{"family":"ipv4","protocol":"tcp","port":"9000-8000"}]`, - `[{"family":"ipv6","protocol":"udp","port":"65536"}]`, - `[{"family":"ipv4","protocol":"tcp","port":"8000-8100"},{"family":"ipv4","protocol":"tcp","port":"8080"}]`, - } { - if _, err := ParsePortWhitelist(value); err == nil { - t.Fatalf("expected %q to fail", value) - } - } -} - -func TestParsePortWhitelistKeepsFamiliesDistinct(t *testing.T) { - rules, err := ParsePortWhitelist(`[ - {"family":"ipv4","protocol":"tcp","port":"443"}, - {"family":"ipv6","protocol":"tcp","port":"443"}, - {"family":"ipv4","protocol":"tcp","port":"443"} - ]`) - if err != nil { - t.Fatalf("parse whitelist: %v", err) - } - if len(rules) != 2 { - t.Fatalf("expected family-specific deduplication, got %#v", rules) - } -} - -func TestNormalizePortWhitelistPrefersFamilyNeutralRequiredRule(t *testing.T) { - rules := NormalizePortWhitelist([]PortWhitelist{ - {Family: "ipv4", Port: "22", Protocol: "tcp"}, - {Family: "ipv6", Port: "22", Protocol: "tcp"}, - {Port: "22", Protocol: "tcp"}, - }) - if len(rules) != 1 || rules[0].Family != "" { - t.Fatalf("family-neutral rule should replace family-specific duplicates: %#v", rules) - } -} diff --git a/agent/utils/firewall/sync/diff.go b/agent/utils/firewall/sync/diff.go new file mode 100644 index 000000000000..e7f2a768a672 --- /dev/null +++ b/agent/utils/firewall/sync/diff.go @@ -0,0 +1,120 @@ +package sync + +type Status string + +const ( + StatusReady Status = "ready" + StatusExisting Status = "existing" + StatusRemove Status = "remove" + StatusBlocked Status = "blocked" +) + +type Outcome string + +const ( + OutcomeApplied Outcome = "applied" + OutcomeSkipped Outcome = "skipped" + OutcomeRemoved Outcome = "removed" + OutcomeFailed Outcome = "failed" +) + +type ReasonCode string + +const ( + ReasonInvalidPolicy ReasonCode = "invalid_policy" + ReasonAlreadyExists ReasonCode = "already_exists_in_target" + ReasonOnlyExistsInTarget ReasonCode = "only_exists_in_target" + ReasonManagedOnlyInTarget ReasonCode = "managed_only_exists_in_target" + ReasonUnsafeRemoval ReasonCode = "unsafe_managed_rule_removal" +) + +func ReasonMessage(code ReasonCode) string { + switch code { + case ReasonAlreadyExists: + return "rule already exists in target backend" + case ReasonOnlyExistsInTarget: + return "rule exists only in target backend" + case ReasonManagedOnlyInTarget: + return "managed rule exists only in target backend" + case ReasonUnsafeRemoval: + return "managed runtime rule cannot be safely removed" + default: + return "" + } +} + +type Desired[T any, P any] struct { + Value T + Payload P + Err error +} + +type Item[P any] struct { + Payload P + Status Status + ReasonCode ReasonCode + Reason string +} + +func Diff[T any, P any](desired []Desired[T, P], actual []T, key func(T) string, actualPayload func(T) P) []Item[P] { + items := make([]Item[P], 0, len(desired)+len(actual)) + actualByKey := make(map[string][]int, len(actual)) + for index, value := range actual { + actualByKey[key(value)] = append(actualByKey[key(value)], index) + } + matched := make([]bool, len(actual)) + for _, candidate := range desired { + item := Item[P]{Payload: candidate.Payload} + switch { + case candidate.Err != nil: + item.Status, item.ReasonCode, item.Reason = StatusBlocked, ReasonInvalidPolicy, candidate.Err.Error() + default: + match := unmatchedIndex(actualByKey[key(candidate.Value)], matched) + if match >= 0 { + matched[match] = true + item.Status, item.ReasonCode = StatusExisting, ReasonAlreadyExists + item.Reason = ReasonMessage(item.ReasonCode) + } else { + item.Status = StatusReady + } + } + items = append(items, item) + } + for index, value := range actual { + if matched[index] { + continue + } + items = append(items, Item[P]{ + Payload: actualPayload(value), Status: StatusRemove, ReasonCode: ReasonOnlyExistsInTarget, + Reason: ReasonMessage(ReasonOnlyExistsInTarget), + }) + } + return items +} + +func StatesEqual[T any](left, right []T, key func(T) string) bool { + if len(left) != len(right) { + return false + } + counts := make(map[string]int, len(left)) + for _, value := range left { + counts[key(value)]++ + } + for _, value := range right { + valueKey := key(value) + if counts[valueKey] == 0 { + return false + } + counts[valueKey]-- + } + return true +} + +func unmatchedIndex(indices []int, matched []bool) int { + for _, index := range indices { + if !matched[index] { + return index + } + } + return -1 +} diff --git a/agent/utils/firewall/sync/order.go b/agent/utils/firewall/sync/order.go new file mode 100644 index 000000000000..96e324755971 --- /dev/null +++ b/agent/utils/firewall/sync/order.go @@ -0,0 +1,141 @@ +package sync + +import ( + "slices" + "strings" + + "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" +) + +func SupportsManagedOrder(provider filter.Provider) bool { + return provider == filter.ProviderIptables || provider == filter.ProviderNftables || provider == filter.ProviderUFW +} + +func ManagedOrderDrift(snapshot filter.Snapshot, desiredMarkers []string) (map[string]struct{}, bool) { + if !SupportsManagedOrder(snapshot.Scope.Provider) || len(desiredMarkers) < 2 { + return nil, true + } + expected := make(map[string]struct{}, len(desiredMarkers)) + for _, marker := range desiredMarkers { + expected[marker] = struct{}{} + } + actual := make([]string, 0, len(desiredMarkers)) + segments := make(map[string]int, len(desiredMarkers)) + segment := 0 + for _, observed := range snapshot.Rules { + _, wanted := expected[observed.Marker] + if wanted { + actual = append(actual, observed.Marker) + if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque { + segment++ + segments[observed.Marker] = segment + segment++ + } else { + segments[observed.Marker] = segment + } + continue + } + if strings.HasPrefix(observed.Marker, "1panel-rule:") && + !observed.Protected && observed.ParseStatus != filter.ParseStatusOpaque { + continue + } + segment++ + } + + desired := make([]string, 0, len(actual)) + for _, marker := range desiredMarkers { + if _, exists := segments[marker]; exists { + desired = append(desired, marker) + } + } + if slices.Equal(actual, desired) { + return nil, true + } + drifted := make(map[string]struct{}, len(desired)) + for index := range desired { + if actual[index] != desired[index] { + drifted[actual[index]] = struct{}{} + drifted[desired[index]] = struct{}{} + } + } + feasible, previousSegment := true, -1 + for _, marker := range desired { + if segments[marker] < previousSegment { + feasible = false + break + } + previousSegment = segments[marker] + } + return drifted, feasible +} + +func InsertionPosition(snapshot filter.Snapshot, desiredMarkers []string, targetMarker string) (int64, bool) { + if !SupportsManagedOrder(snapshot.Scope.Provider) { + return 0, false + } + targetIndex := slices.Index(desiredMarkers, targetMarker) + if targetIndex < 0 { + return 0, false + } + for index := targetIndex - 1; index >= 0; index-- { + if _, position, exists := ObservedByMarker(snapshot, desiredMarkers[index]); exists { + return int64(position + 1), true + } + } + for index := targetIndex + 1; index < len(desiredMarkers); index++ { + if _, position, exists := ObservedByMarker(snapshot, desiredMarkers[index]); exists { + return int64(position), true + } + } + return 0, false +} + +func NextManagedOrderChange(snapshot filter.Snapshot, desiredMarkers []string) (string, int, bool, error) { + expected := make(map[string]struct{}, len(desiredMarkers)) + for _, marker := range desiredMarkers { + expected[marker] = struct{}{} + } + actual := make([]string, 0, len(desiredMarkers)) + positions := make([]int, 0, len(desiredMarkers)) + for index, observed := range snapshot.Rules { + if _, exists := expected[observed.Marker]; !exists { + continue + } + actual = append(actual, observed.Marker) + position := index + 1 + if observed.Locator.Position != nil { + position = *observed.Locator.Position + } + positions = append(positions, position) + } + if len(actual) != len(desiredMarkers) { + return "", 0, false, filter.ErrRuleStale + } + for index := range desiredMarkers { + if actual[index] != desiredMarkers[index] { + return desiredMarkers[index], positions[index], false, nil + } + } + return "", 0, true, nil +} + +func ObservedByMarker(snapshot filter.Snapshot, marker string) (filter.ObservedRule, int, bool) { + for index, observed := range snapshot.Rules { + if observed.Marker == marker { + position := index + 1 + if observed.Locator.Position != nil { + position = *observed.Locator.Position + } + return observed, position, true + } + } + return filter.ObservedRule{}, 0, false +} + +func ObservedRule(observed filter.ObservedRule) filter.FirewallRule { + rule := observed.Rule + if rule.UUID == "" && strings.HasPrefix(observed.Marker, "1panel-rule:") { + rule.UUID = strings.TrimSpace(strings.TrimPrefix(observed.Marker, "1panel-rule:")) + } + return rule +} diff --git a/agent/utils/re/firewall.go b/agent/utils/re/firewall.go index 8e6f234cc36a..03a6fd4f2d60 100644 --- a/agent/utils/re/firewall.go +++ b/agent/utils/re/firewall.go @@ -5,4 +5,5 @@ import "regexp" var ( UFWNumberedRulePrefixRegex = regexp.MustCompile(`^\s*\[\s*([0-9]+)\]\s+(.+?)\s*$`) UFWNumberedRuleRegex = regexp.MustCompile(`^\s*\[\s*([0-9]+)\]\s+(.+?)\s+(ALLOW|DENY|REJECT|LIMIT)(?:\s+(IN|OUT|FWD))?\s+(.+?)\s*$`) + ForwardInterfaceRegex = regexp.MustCompile(`^[A-Za-z0-9_.:@-]{1,15}$`) ) diff --git a/core/init/migration/migrate_test.go b/core/init/migration/migrate_test.go deleted file mode 100644 index dbac5bac7172..000000000000 --- a/core/init/migration/migrate_test.go +++ /dev/null @@ -1,24 +0,0 @@ -package migration - -import "testing" - -func TestCoreMigrationsRegisterFirewallMenuUpgrade(t *testing.T) { - migrations := coreMigrations() - seen := make(map[string]int, len(migrations)) - foundFirewallMenu := false - for index, migration := range migrations { - if migration == nil || migration.ID == "" { - t.Fatalf("invalid migration at index %d: %#v", index, migration) - } - if previous, exists := seen[migration.ID]; exists { - t.Fatalf("duplicate migration ID %q at indexes %d and %d", migration.ID, previous, index) - } - seen[migration.ID] = index - if migration.ID == "20260819-update-firewall-menu-path" { - foundFirewallMenu = true - } - } - if !foundFirewallMenu { - t.Fatal("firewall menu path migration is not registered") - } -} diff --git a/core/init/migration/migrations/firewall_test.go b/core/init/migration/migrations/firewall_test.go deleted file mode 100644 index 0aed2c605eb7..000000000000 --- a/core/init/migration/migrations/firewall_test.go +++ /dev/null @@ -1,114 +0,0 @@ -package migrations - -import ( - "encoding/json" - "path/filepath" - "testing" - - "github.com/1Panel-dev/1Panel/core/app/dto" - "github.com/1Panel-dev/1Panel/core/app/model" - "github.com/1Panel-dev/1Panel/core/init/migration/helper" - "github.com/glebarez/sqlite" - "gorm.io/gorm" - "gorm.io/gorm/logger" -) - -func TestUpdateFirewallMenuPathMigratesCustomizedMenu(t *testing.T) { - db := newCoreFirewallMigrationTestDB(t) - menus := []dto.ShowMenu{{ - ID: "custom-parent", Label: "Custom", Children: []dto.ShowMenu{ - {ID: "74", Label: "renamed-firewall", Title: "custom.title", Path: "/hosts/firewall/port", Sort: 987, IsShow: false}, - {ID: "other", Label: "Other", Path: "/hosts/firewall/port", Sort: 123, IsShow: true}, - }, - }} - seedCoreHideMenu(t, db, menus) - - if err := UpdateFirewallMenuPath.Migrate(db); err != nil { - t.Fatal(err) - } - after := loadCoreHideMenu(t, db) - firewall := after[0].Children[0] - if firewall.Path != "/hosts/firewall/rules" || firewall.Title != "custom.title" || firewall.Sort != 987 || firewall.IsShow { - t.Fatalf("migration did not preserve customized firewall menu: %#v", firewall) - } - if other := after[0].Children[1]; other.Path != "/hosts/firewall/port" { - t.Fatalf("unrelated menu was changed: %#v", other) - } -} - -func TestUpdateFirewallMenuPathSupportsLegacyLabelAndIsIdempotent(t *testing.T) { - db := newCoreFirewallMigrationTestDB(t) - menus := []dto.ShowMenu{{Children: []dto.ShowMenu{ - {ID: "legacy-id", Label: "FirewallPort", Path: "/hosts/firewall/port"}, - {ID: "74", Label: "FirewallPort", Path: "/hosts/firewall/rules"}, - }}} - seedCoreHideMenu(t, db, menus) - - for i := 0; i < 2; i++ { - if err := UpdateFirewallMenuPath.Migrate(db); err != nil { - t.Fatalf("migrate firewall menu on pass %d: %v", i+1, err) - } - } - after := loadCoreHideMenu(t, db) - for _, menu := range after[0].Children { - if menu.Path != "/hosts/firewall/rules" { - t.Fatalf("legacy firewall menu path remained after migration: %#v", menu) - } - } -} - -func TestDefaultMenuUsesFirewallV2Route(t *testing.T) { - var menus []dto.ShowMenu - if err := json.Unmarshal([]byte(helper.LoadMenus()), &menus); err != nil { - t.Fatal(err) - } - for _, parent := range menus { - for _, child := range parent.Children { - if child.ID == "74" { - if child.Path != "/hosts/firewall/rules" { - t.Fatalf("default firewall path = %q", child.Path) - } - return - } - } - } - t.Fatal("default firewall menu was not found") -} - -func newCoreFirewallMigrationTestDB(t *testing.T) *gorm.DB { - t.Helper() - db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "migration.db")), &gorm.Config{ - Logger: logger.Default.LogMode(logger.Silent), - }) - if err != nil { - t.Fatal(err) - } - if err := db.AutoMigrate(&model.Setting{}); err != nil { - t.Fatal(err) - } - return db -} - -func seedCoreHideMenu(t *testing.T, db *gorm.DB, menus []dto.ShowMenu) { - t.Helper() - value, err := json.Marshal(menus) - if err != nil { - t.Fatal(err) - } - if err := db.Create(&model.Setting{Key: "HideMenu", Value: string(value)}).Error; err != nil { - t.Fatal(err) - } -} - -func loadCoreHideMenu(t *testing.T, db *gorm.DB) []dto.ShowMenu { - t.Helper() - var setting model.Setting - if err := db.Where("key = ?", "HideMenu").First(&setting).Error; err != nil { - t.Fatal(err) - } - var menus []dto.ShowMenu - if err := json.Unmarshal([]byte(setting.Value), &menus); err != nil { - t.Fatal(err) - } - return menus -} diff --git a/frontend/src/api/interface/firewall.ts b/frontend/src/api/interface/firewall.ts index cdefad42dcce..6935ee72237b 100644 --- a/frontend/src/api/interface/firewall.ts +++ b/frontend/src/api/interface/firewall.ts @@ -277,6 +277,7 @@ export namespace Firewall { forwardRule?: RuleForward; dockerRule?: DockerGuardEndpoint; status: RuleSyncStatus; + reasonCode?: string; reason?: string; } diff --git a/frontend/src/views/host/firewall/sync/index.vue b/frontend/src/views/host/firewall/sync/index.vue index 99d3371ec10d..f7ec81d50c87 100644 --- a/frontend/src/views/host/firewall/sync/index.vue +++ b/frontend/src/views/host/firewall/sync/index.vue @@ -147,7 +147,7 @@ @@ -296,7 +296,16 @@ const syncReasonKeys: Record = { 'protected firewall rule cannot be modified': 'protectedRule', }; -const reasonText = (reason?: string) => { +const syncReasonCodeKeys: Record = { + already_exists_in_target: 'alreadyExistsInTarget', + only_exists_in_target: 'onlyInTarget', + managed_only_exists_in_target: 'managedOnlyInTarget', + unsafe_managed_rule_removal: 'managedRuntimeCannotRemove', +}; + +const reasonText = (reasonCode?: string, reason?: string) => { + const codedReasonKey = reasonCode ? syncReasonCodeKeys[reasonCode] : undefined; + if (codedReasonKey) return i18n.global.t(`firewall.ruleSyncReasonDetail.${codedReasonKey}`); if (!reason) return '-'; const reasonKey = syncReasonKeys[reason]; if (reasonKey) return i18n.global.t(`firewall.ruleSyncReasonDetail.${reasonKey}`);