From 153049931fa32a4c4c5797f5b6306bef7c2c8f71 Mon Sep 17 00:00:00 2001 From: ssongliu Date: Mon, 21 Sep 2026 15:05:51 +0800 Subject: [PATCH] refactor(firewall): simplify state reads and batch verification --- agent/app/api/v2/firewall.go | 54 +- agent/app/dto/firewall.go | 15 +- agent/app/model/firewall.go | 227 +- agent/app/repo/docker_port_guard.go | 27 - agent/app/repo/firewall_rule.go | 61 +- agent/app/repo/forwarding_rule.go | 26 +- agent/app/service/docker.go | 8 +- agent/app/service/entry.go | 5 +- agent/app/service/firewall.go | 3313 ++++---------- agent/app/service/firewall_docker.go | 957 +--- agent/app/service/firewall_rule_task.go | 76 - agent/app/service/firewall_selection.go | 65 - agent/app/service/firewall_setting.go | 463 +- agent/app/service/firewall_sync.go | 1455 ++---- agent/app/service/firewall_utils.go | 3948 +++++++++++++++++ agent/app/service/forward.go | 597 ++- agent/app/service/website_domain.go | 5 +- agent/i18n/lang/en.yaml | 4 + agent/i18n/lang/es-ES.yaml | 4 + agent/i18n/lang/fa.yaml | 4 + agent/i18n/lang/ja.yaml | 4 + agent/i18n/lang/ko.yaml | 4 + agent/i18n/lang/lo.yaml | 4 + agent/i18n/lang/ms.yaml | 4 + agent/i18n/lang/pt-BR.yaml | 4 + agent/i18n/lang/ru.yaml | 4 + agent/i18n/lang/tr.yaml | 4 + agent/i18n/lang/zh-Hant.yaml | 4 + agent/i18n/lang/zh.yaml | 4 + agent/init/firewall/firewall.go | 8 +- .../migrations/firewall_whitelist.go | 3 + .../utils/host_firewall_transfer.go | 52 +- agent/utils/cmd/cmdx.go | 9 + agent/utils/firewall/docker_guard/inspect.go | 247 ++ agent/utils/firewall/docker_guard/manager.go | 113 +- agent/utils/firewall/docker_guard/nftables.go | 91 +- agent/utils/firewall/docker_guard/policy.go | 166 +- agent/utils/firewall/docker_guard/runtime.go | 182 +- agent/utils/firewall/filter/adapter.go | 59 +- agent/utils/firewall/filter/external.go | 20 + agent/utils/firewall/filter/identity.go | 130 +- agent/utils/firewall/filter/inventory.go | 273 -- agent/utils/firewall/filter/model.go | 23 +- agent/utils/firewall/filter/normalize.go | 26 +- .../filter/providers/firewalld/adapter.go | 381 +- .../filter/providers/iptables/adapter.go | 575 +-- .../filter/providers/nftables/adapter.go | 413 +- .../firewall/filter/providers/ufw/adapter.go | 305 +- .../firewall/filter/runtime/inventory.go | 123 - .../utils/firewall/filter/runtime/runtime.go | 404 -- agent/utils/firewall/filter/safety.go | 141 +- agent/utils/firewall/forwarding/forwarding.go | 106 +- agent/utils/firewall/forwarding/iptables.go | 107 +- agent/utils/firewall/forwarding/nftables.go | 157 +- agent/utils/firewall/iptables_helper/ipv6.go | 19 +- .../utils/firewall/iptables_helper/manager.go | 101 +- .../firewall/iptables_helper/persistence.go | 9 +- .../utils/firewall/iptables_helper/repair.go | 7 +- agent/utils/firewall/lifecycle/lifecycle.go | 23 +- agent/utils/firewall/lifecycle/operator.go | 178 - .../firewall/lifecycle/providers/firewalld.go | 34 +- .../utils/firewall/nftables_helper/manager.go | 86 +- .../utils/firewall/nftables_helper/runtime.go | 42 +- agent/utils/firewall/sync/diff.go | 93 - agent/utils/firewall/sync/order.go | 136 - frontend/src/api/interface/firewall.ts | 5 +- frontend/src/lang/modules/en.ts | 5 +- frontend/src/lang/modules/es-es.ts | 6 +- frontend/src/lang/modules/fa.ts | 5 +- frontend/src/lang/modules/ja.ts | 6 +- frontend/src/lang/modules/ko.ts | 5 +- frontend/src/lang/modules/lo.ts | 5 +- frontend/src/lang/modules/ms.ts | 7 +- frontend/src/lang/modules/pt-br.ts | 6 +- frontend/src/lang/modules/ru.ts | 7 +- frontend/src/lang/modules/tr.ts | 7 +- frontend/src/lang/modules/zh-Hant.ts | 5 +- frontend/src/lang/modules/zh.ts | 5 +- .../host/firewall/docker/detail/index.vue | 43 +- .../host/firewall/docker/import/index.vue | 70 +- .../src/views/host/firewall/docker/index.vue | 42 +- .../host/firewall/forward/import/index.vue | 70 +- .../src/views/host/firewall/forward/index.vue | 6 +- .../views/host/firewall/rule/import/index.vue | 69 +- .../src/views/host/firewall/rule/index.vue | 37 +- .../host/firewall/rule/operate/index.vue | 23 +- .../src/views/host/firewall/status/index.vue | 14 + .../src/views/host/firewall/sync/index.vue | 21 +- .../views/host/firewall/utils/validation.ts | 3 + 89 files changed, 8128 insertions(+), 8536 deletions(-) delete mode 100644 agent/app/service/firewall_rule_task.go delete mode 100644 agent/app/service/firewall_selection.go create mode 100644 agent/app/service/firewall_utils.go create mode 100644 agent/utils/firewall/docker_guard/inspect.go create mode 100644 agent/utils/firewall/filter/external.go delete mode 100644 agent/utils/firewall/filter/runtime/inventory.go delete mode 100644 agent/utils/firewall/filter/runtime/runtime.go delete mode 100644 agent/utils/firewall/lifecycle/operator.go delete mode 100644 agent/utils/firewall/sync/order.go diff --git a/agent/app/api/v2/firewall.go b/agent/app/api/v2/firewall.go index 5c60f5459a78..b42eaba15f82 100644 --- a/agent/app/api/v2/firewall.go +++ b/agent/app/api/v2/firewall.go @@ -2,14 +2,16 @@ package v2 import ( "errors" + "github.com/1Panel-dev/1Panel/agent/buserr" "net/http" "strings" "github.com/1Panel-dev/1Panel/agent/app/api/v2/helper" "github.com/1Panel-dev/1Panel/agent/app/dto" "github.com/1Panel-dev/1Panel/agent/app/repo" - "github.com/1Panel-dev/1Panel/agent/app/service" + "github.com/1Panel-dev/1Panel/agent/global" + "github.com/1Panel-dev/1Panel/agent/utils/docker" "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" "github.com/gin-gonic/gin" ) @@ -455,6 +457,8 @@ func normalizeFirewallRuleUUID(c *gin.Context, value *string) bool { } func handleFirewallRuleError(c *gin.Context, err error) { + var businessErr buserr.BusinessError + isBusinessError := errors.As(err, &businessErr) switch { case errors.Is(err, filter.ErrProtectedRule): helper.ErrorWithBusinessCode(c, http.StatusBadRequest, "FW_LOCKOUT_RISK", "ErrInvalidParams", err) @@ -462,7 +466,7 @@ func handleFirewallRuleError(c *gin.Context, err error) { helper.ErrorWithBusinessCode(c, http.StatusConflict, "FW_RULE_STALE", "ErrInvalidParams", err) case errors.Is(err, repo.ErrFirewallRuleRevisionConflict): helper.ErrorWithBusinessCode(c, http.StatusConflict, "FW_RULE_REVISION_CONFLICT", "ErrInvalidParams", err) - case errors.Is(err, filter.ErrManagedScopeChange): + case isBusinessError && businessErr.Msg == "ErrFirewallRuleScopeChange": helper.ErrorWithBusinessCode(c, http.StatusBadRequest, "FW_SCOPE_UNSUPPORTED", "ErrFirewallRuleScopeChange", err) case errors.Is(err, filter.ErrUnsupportedScope), errors.Is(err, filter.ErrInvalidScope), errors.Is(err, filter.ErrProviderUnavailable), errors.Is(err, filter.ErrAdapterUnavailable): @@ -470,6 +474,9 @@ func handleFirewallRuleError(c *gin.Context, err error) { case errors.Is(err, filter.ErrInvalidRule), errors.Is(err, filter.ErrRuleOperation), errors.Is(err, filter.ErrRuleConflict), errors.Is(err, repo.ErrFirewallPersistenceInvalid): helper.ErrorWithBusinessCode(c, http.StatusBadRequest, "FW_RULE_UNSUPPORTED", "ErrInvalidParams", err) + case isBusinessError && businessErr.Msg == "ErrInvalidParams": + c.JSON(http.StatusOK, dto.Response{Code: http.StatusBadRequest, ErrorCode: "FW_RULE_UNSUPPORTED", Message: err.Error()}) + c.Abort() case errors.Is(err, filter.ErrVerificationFailed): helper.ErrorWithBusinessCode(c, http.StatusInternalServerError, "FW_VERIFY_FAILED", "ErrInternalServer", err) default: @@ -573,14 +580,10 @@ func (b *BaseApi) OperateFirewallBackend(c *gin.Context) { return } if err := firewallSettingService.Operate(c.Request.Context(), request); err != nil { - if errors.Is(err, service.ErrFirewallBackendCleanupRequired) { - helper.ErrorWithBusinessCode( - c, - http.StatusConflict, - "FW_BACKEND_CLEANUP_REQUIRED", - "ErrInvalidParams", - err, - ) + var businessErr buserr.BusinessError + if errors.As(err, &businessErr) && businessErr.Msg == "ErrFirewallBackendCleanupRequired" { + c.JSON(http.StatusOK, dto.Response{Code: http.StatusConflict, ErrorCode: "FW_BACKEND_CLEANUP_REQUIRED", Message: err.Error()}) + c.Abort() return } helper.InternalServer(c, err) @@ -709,19 +712,26 @@ func (b *BaseApi) UpsertDockerPortGuardPolicies(c *gin.Context) { } func handleDockerPortGuardError(c *gin.Context, err error) { - if errors.Is(err, service.ErrDockerIptablesChainUnavailable) { - helper.ErrorWithBusinessCode(c, http.StatusServiceUnavailable, "FW_DOCKER_IPTABLES_CHAIN_UNAVAILABLE", "ErrDockerIptablesChainUnavailable", err) - return - } - if errors.Is(err, service.ErrDockerNftablesChainUnavailable) { - helper.ErrorWithBusinessCode(c, http.StatusServiceUnavailable, "FW_DOCKER_NFTABLES_CHAIN_UNAVAILABLE", "ErrDockerNftablesChainUnavailable", err) - return - } - if errors.Is(err, service.ErrDockerGuardInvalid) { - helper.ErrorWithBusinessCode(c, http.StatusBadRequest, "FW_DOCKER_GUARD_INVALID", "ErrInvalidParams", err) - return + var businessErr buserr.BusinessError + if errors.As(err, &businessErr) { + code, errorCode := http.StatusInternalServerError, "" + switch businessErr.Msg { + case "ErrDockerIptablesChainUnavailable": + code, errorCode = http.StatusServiceUnavailable, "FW_DOCKER_IPTABLES_CHAIN_UNAVAILABLE" + case "ErrDockerNftablesChainUnavailable": + code, errorCode = http.StatusServiceUnavailable, "FW_DOCKER_NFTABLES_CHAIN_UNAVAILABLE" + case "ErrInvalidParams": + code, errorCode = http.StatusBadRequest, "FW_DOCKER_GUARD_INVALID" + case "ErrDockerFailed": + code, errorCode = http.StatusServiceUnavailable, "FW_DOCKER_UNAVAILABLE" + } + if errorCode != "" { + c.JSON(http.StatusOK, dto.Response{Code: code, ErrorCode: errorCode, Message: err.Error()}) + c.Abort() + return + } } - if errors.Is(err, service.ErrDockerUnavailable) { + if errors.Is(err, docker.ErrUnavailable) { helper.ErrorWithBusinessCode(c, http.StatusServiceUnavailable, "FW_DOCKER_UNAVAILABLE", "ErrDockerFailed", err) return } diff --git a/agent/app/dto/firewall.go b/agent/app/dto/firewall.go index 02981b277634..8ee530aae083 100644 --- a/agent/app/dto/firewall.go +++ b/agent/app/dto/firewall.go @@ -49,10 +49,11 @@ type FirewallBackendOption struct { } type FirewallBackendFamilyStatus struct { - Available bool `json:"available"` - Initialized bool `json:"initialized"` - Bound bool `json:"bound"` - Reason string `json:"reason,omitempty"` + Available bool `json:"available"` + Initialized bool `json:"initialized"` + Bound bool `json:"bound"` + Reason string `json:"reason,omitempty"` + ForwardPolicy string `json:"forwardPolicy,omitempty"` } type FirewallBackendGroup struct { @@ -240,8 +241,10 @@ type DockerPortGuardOperation struct { } type FirewallRuleAdopt struct { - Scope filter.Scope `json:"scope" validate:"required"` - InstanceKey string `json:"instanceKey" validate:"required,max=128"` + Scope filter.Scope `json:"scope" validate:"required"` + InstanceKey string `json:"instanceKey,omitempty" validate:"omitempty,max=128"` + Rule *filter.FirewallRule `json:"rule,omitempty"` + Marker string `json:"marker,omitempty" validate:"max=256"` } type FirewallRuleCreateItem struct { diff --git a/agent/app/model/firewall.go b/agent/app/model/firewall.go index b41b5396f965..eca6db098a6f 100644 --- a/agent/app/model/firewall.go +++ b/agent/app/model/firewall.go @@ -1,222 +1,51 @@ package model -import ( - "crypto/sha256" - "encoding/hex" - "encoding/json" - "fmt" - "sort" - "strings" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -const FirewallRuleSequenceStep int64 = 1 << 32 - type DockerPortGuardPolicy struct { BaseModel - UUID string `gorm:"size:64;not null;uniqueIndex" json:"uuid"` - ReadOnly bool `gorm:"not null;default:false;uniqueIndex:idx_docker_port_guard_endpoint" json:"-"` - Family string `gorm:"size:16;not null;uniqueIndex:idx_docker_port_guard_endpoint" json:"family"` - HostIP string `gorm:"size:64;not null;uniqueIndex:idx_docker_port_guard_endpoint" json:"hostIP"` - HostPort uint16 `gorm:"not null;uniqueIndex:idx_docker_port_guard_endpoint" json:"hostPort"` - Protocol string `gorm:"size:8;not null;uniqueIndex:idx_docker_port_guard_endpoint" json:"protocol"` - Mode string `gorm:"size:32;not null" json:"mode"` + UUID string `gorm:"uniqueIndex" json:"uuid"` + ReadOnly bool `gorm:"default:false;uniqueIndex:idx_docker_port_guard_endpoint" json:"-"` + Family string `gorm:"uniqueIndex:idx_docker_port_guard_endpoint" json:"family"` + HostIP string `gorm:"uniqueIndex:idx_docker_port_guard_endpoint" json:"hostIP"` + HostPort uint16 `gorm:"uniqueIndex:idx_docker_port_guard_endpoint" json:"hostPort"` + Protocol string `gorm:"uniqueIndex:idx_docker_port_guard_endpoint" json:"protocol"` + Mode string `json:"mode"` Sources string `gorm:"type:text" json:"-"` Description string `gorm:"type:text" json:"description"` - NativeAction string `gorm:"size:32;not null;default:''" json:"-"` + NativeAction string `gorm:"default:''" json:"-"` NativeRules string `gorm:"type:text" json:"-"` - Sequence int64 `gorm:"not null;default:0" json:"-"` + Sequence int64 `gorm:"default:0" json:"-"` } type ForwardingRule struct { BaseModel - Family string `gorm:"size:16;not null;uniqueIndex:idx_forwarding_rule_identity" json:"family"` - Protocol string `gorm:"size:8;not null;uniqueIndex:idx_forwarding_rule_identity" json:"protocol"` - Port string `gorm:"size:32;not null;uniqueIndex:idx_forwarding_rule_identity" json:"port"` - TargetIP string `gorm:"size:64;not null;uniqueIndex:idx_forwarding_rule_identity" json:"targetIP"` - TargetPort string `gorm:"size:32;not null;uniqueIndex:idx_forwarding_rule_identity" json:"targetPort"` - Interface string `gorm:"size:32;not null;default:'';uniqueIndex:idx_forwarding_rule_identity" json:"interface"` + Family string `gorm:"uniqueIndex:idx_forwarding_rule_identity" json:"family"` + Protocol string `gorm:"uniqueIndex:idx_forwarding_rule_identity" json:"protocol"` + Port string `gorm:"uniqueIndex:idx_forwarding_rule_identity" json:"port"` + TargetIP string `gorm:"uniqueIndex:idx_forwarding_rule_identity" json:"targetIP"` + TargetPort string `gorm:"uniqueIndex:idx_forwarding_rule_identity" json:"targetPort"` + Interface string `gorm:"default:'';uniqueIndex:idx_forwarding_rule_identity" json:"interface"` } type FirewallRule struct { - UUID string `gorm:"size:64;primaryKey" json:"uuid"` - Family string `gorm:"size:16;not null" json:"family"` - - Protocol string `gorm:"size:32;not null" json:"protocol"` - SourceAddress string `gorm:"size:255" json:"sourceAddress"` - SourcePort string `gorm:"size:64" json:"sourcePort"` - DestinationAddress string `gorm:"size:255" json:"destinationAddress"` - DestinationPort string `gorm:"size:64" json:"destinationPort"` - Interface string `gorm:"size:128" json:"interface"` + UUID string `gorm:"primaryKey" json:"uuid"` + Family string `json:"family"` + + Protocol string `json:"protocol"` + SourceAddress string `json:"sourceAddress"` + SourcePort string `json:"sourcePort"` + DestinationAddress string `json:"destinationAddress"` + DestinationPort string `json:"destinationPort"` + Interface string `json:"interface"` ConnectionStates string `gorm:"type:text" json:"connectionStates"` - Action string `gorm:"size:32;not null" json:"action"` + Action string `json:"action"` Description string `gorm:"type:text" json:"description"` CompatibilityError string `gorm:"type:text" json:"compatibilityError,omitempty"` Priority *int `json:"priority,omitempty"` Sequence *int64 `gorm:"index" json:"sequence,omitempty"` - Origin string `gorm:"size:32;not null" json:"origin"` - Owner string `gorm:"size:320;not null" json:"owner"` - Revision uint `gorm:"not null;default:1" json:"revision"` -} - -func FirewallRuleOwner(sourceKind, sourceID string) string { - sourceKind = strings.TrimSpace(sourceKind) - sourceID = strings.TrimSpace(sourceID) - if sourceID == "" { - return sourceKind - } - return sourceKind + ":" + sourceID -} - -func FirewallRuleFromDomain(rule filter.FirewallRule) (FirewallRule, error) { - normalized, err := filter.NormalizeRule(rule) - if err != nil { - return FirewallRule{}, err - } - switch normalized.NativeKind { - case "", filter.NativeKindRule, filter.NativeKindZonePort, filter.NativeKindRichRule, filter.NativeKindUFWRule: - default: - return FirewallRule{}, fmt.Errorf("%w: native rule %q cannot be stored as a provider-neutral policy", filter.ErrUnsupportedScope, normalized.NativeKind) - } - record := FirewallRule{ - Family: string(normalized.Scope.Family), - Protocol: normalized.Protocol, - SourceAddress: normalized.SourceAddress, - SourcePort: normalized.SourcePort, - DestinationAddress: normalized.DestinationAddress, - DestinationPort: normalized.DestinationPort, - Interface: normalized.Interface, - ConnectionStates: strings.Join(normalized.ConnectionStates, ","), - Action: string(normalized.Action), - Description: normalized.Description, - } - if normalized.Scope.Provider == filter.ProviderFirewalld { - record.Priority = normalized.Priority - } - return record, nil -} - -func (rule FirewallRule) PolicyKey() string { - payload, _ := json.Marshal(struct { - Family string `json:"family"` - Protocol string `json:"protocol"` - SourceAddress string `json:"sourceAddress,omitempty"` - SourcePort string `json:"sourcePort,omitempty"` - DestinationAddress string `json:"destinationAddress,omitempty"` - DestinationPort string `json:"destinationPort,omitempty"` - Interface string `json:"interface,omitempty"` - ConnectionStates string `json:"connectionStates,omitempty"` - Action string `json:"action"` - }{ - Family: rule.Family, Protocol: rule.Protocol, - SourceAddress: rule.SourceAddress, SourcePort: rule.SourcePort, - DestinationAddress: rule.DestinationAddress, DestinationPort: rule.DestinationPort, - Interface: rule.Interface, ConnectionStates: rule.ConnectionStates, Action: rule.Action, - }) - 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.ProviderUFW && strings.EqualFold(strings.TrimSpace(base.Protocol), "all") && - strings.TrimSpace(base.SourcePort) == "" && strings.TrimSpace(base.DestinationPort) != "" { - base.Protocol = "tcp/udp" - } - 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 + Origin string `json:"origin"` + Owner string `json:"owner"` + Revision uint `gorm:"default:1" json:"revision"` } diff --git a/agent/app/repo/docker_port_guard.go b/agent/app/repo/docker_port_guard.go index dfe2b5e4b4ac..9b792b6a330e 100644 --- a/agent/app/repo/docker_port_guard.go +++ b/agent/app/repo/docker_port_guard.go @@ -11,10 +11,8 @@ import ( type IDockerPortGuardRepo interface { ListManaged(context.Context) ([]model.DockerPortGuardPolicy, error) - ListRuntimeReadOnly(context.Context) ([]model.DockerPortGuardPolicy, error) DeleteBatch(context.Context, []string) error UpsertBatch(context.Context, []model.DockerPortGuardPolicy) error - ReplaceRuntimeReadOnly(context.Context, []model.DockerPortGuardPolicy) error } type DockerPortGuardRepo struct{} @@ -30,15 +28,6 @@ func (r *DockerPortGuardRepo) ListManaged(ctx context.Context) ([]model.DockerPo return policies, err } -func (r *DockerPortGuardRepo) ListRuntimeReadOnly(ctx context.Context) ([]model.DockerPortGuardPolicy, error) { - var policies []model.DockerPortGuardPolicy - err := global.DB.WithContext(ctx). - Where("read_only = ?", true). - Order("family, sequence, host_ip, host_port, protocol"). - Find(&policies).Error - return policies, err -} - func (r *DockerPortGuardRepo) DeleteBatch(ctx context.Context, uuids []string) error { return global.DB.WithContext(ctx). Where("read_only = ? AND uuid IN ?", false, uuids). @@ -59,19 +48,3 @@ func (r *DockerPortGuardRepo) UpsertBatch(ctx context.Context, policies []model. return nil }) } - -func (r *DockerPortGuardRepo) ReplaceRuntimeReadOnly(ctx context.Context, policies []model.DockerPortGuardPolicy) error { - return global.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Where("read_only = ?", true). - Delete(&model.DockerPortGuardPolicy{}).Error; err != nil { - return err - } - if len(policies) == 0 { - return nil - } - for i := range policies { - policies[i].ReadOnly = true - } - return tx.Create(&policies).Error - }) -} diff --git a/agent/app/repo/firewall_rule.go b/agent/app/repo/firewall_rule.go index 50413e77502f..422fc5c3f936 100644 --- a/agent/app/repo/firewall_rule.go +++ b/agent/app/repo/firewall_rule.go @@ -10,6 +10,7 @@ import ( "github.com/1Panel-dev/1Panel/agent/global" "github.com/google/uuid" "gorm.io/gorm" + "gorm.io/gorm/clause" ) var ( @@ -23,6 +24,8 @@ type IFirewallRuleRepo interface { List(context.Context, ...DBOption) ([]model.FirewallRule, error) UpdateWithRevision(context.Context, string, uint, map[string]interface{}) error DeleteWithRevision(context.Context, string, uint) error + DeleteBatchWithRevision(context.Context, []model.FirewallRule) map[string]error + SaveResetOrder(context.Context, []model.FirewallRule) error } type FirewallRuleRepo struct { @@ -87,6 +90,51 @@ func (r *FirewallRuleRepo) DeleteWithRevision(ctx context.Context, ruleUUID stri return nil } +func (r *FirewallRuleRepo) DeleteBatchWithRevision(ctx context.Context, rules []model.FirewallRule) map[string]error { + failures := make(map[string]error) + for start := 0; start < len(rules); start += 500 { + batch := rules[start:min(start+500, len(rules))] + ids := make([][]interface{}, 0, len(batch)) + for _, rule := range batch { + ids = append(ids, []interface{}{rule.UUID, rule.Revision}) + failures[rule.UUID] = ErrFirewallRuleRevisionConflict + } + var deleted []model.FirewallRule + err := r.dbFor(ctx).Clauses(clause.Returning{Columns: []clause.Column{{Name: "uuid"}}}). + Where("(uuid, revision) IN ?", ids).Delete(&deleted).Error + if err != nil { + for _, rule := range batch { + failures[rule.UUID] = err + } + continue + } + for _, rule := range deleted { + delete(failures, rule.UUID) + } + } + return failures +} + +func (r *FirewallRuleRepo) SaveResetOrder(ctx context.Context, rules []model.FirewallRule) error { + if len(rules) == 0 { + return nil + } + return r.dbFor(ctx).Transaction(func(tx *gorm.DB) error { + for _, rule := range rules { + result := tx.Model(&model.FirewallRule{}). + Where("uuid = ? AND revision = ?", rule.UUID, rule.Revision). + Updates(map[string]interface{}{"sequence": rule.Sequence, "priority": rule.Priority, "revision": gorm.Expr("revision + 1")}) + if result.Error != nil { + return result.Error + } + if result.RowsAffected == 0 { + return ErrFirewallRuleRevisionConflict + } + } + return nil + }) +} + func (r *FirewallRuleRepo) dbFor(ctx context.Context) *gorm.DB { return firewallDB(ctx, r.db) } @@ -127,18 +175,13 @@ func prepareFirewallRule(rule *model.FirewallRule) error { } func sanitizeRuleUpdates(updates map[string]interface{}) map[string]interface{} { - result := cloneUpdates(updates) - delete(result, "id") - delete(result, "uuid") - delete(result, "revision") - delete(result, "created_at") - return result -} - -func cloneUpdates(updates map[string]interface{}) map[string]interface{} { result := make(map[string]interface{}, len(updates)+1) for key, value := range updates { result[key] = value } + delete(result, "id") + delete(result, "uuid") + delete(result, "revision") + delete(result, "created_at") return result } diff --git a/agent/app/repo/forwarding_rule.go b/agent/app/repo/forwarding_rule.go index 8b07430b8234..92ac0d90fd49 100644 --- a/agent/app/repo/forwarding_rule.go +++ b/agent/app/repo/forwarding_rule.go @@ -5,12 +5,12 @@ import ( "github.com/1Panel-dev/1Panel/agent/app/model" "github.com/1Panel-dev/1Panel/agent/global" - "gorm.io/gorm" ) type IForwardingRuleRepo interface { List(context.Context) ([]model.ForwardingRule, error) - ReplaceAll(context.Context, []model.ForwardingRule) error + CreateBatch(context.Context, []model.ForwardingRule) error + DeleteBatch(context.Context, []uint) error } type ForwardingRuleRepo struct{} @@ -23,14 +23,16 @@ func (r *ForwardingRuleRepo) List(ctx context.Context) ([]model.ForwardingRule, return rules, err } -func (r *ForwardingRuleRepo) ReplaceAll(ctx context.Context, rules []model.ForwardingRule) error { - return global.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&model.ForwardingRule{}).Error; err != nil { - return err - } - if len(rules) == 0 { - return nil - } - return tx.Create(&rules).Error - }) +func (r *ForwardingRuleRepo) CreateBatch(ctx context.Context, rules []model.ForwardingRule) error { + if len(rules) == 0 { + return nil + } + return global.DB.WithContext(ctx).CreateInBatches(&rules, 500).Error +} + +func (r *ForwardingRuleRepo) DeleteBatch(ctx context.Context, ids []uint) error { + if len(ids) == 0 { + return nil + } + return global.DB.WithContext(ctx).Where("id IN ?", ids).Delete(&model.ForwardingRule{}).Error } diff --git a/agent/app/service/docker.go b/agent/app/service/docker.go index ee6d369a15de..3eb284bcc333 100644 --- a/agent/app/service/docker.go +++ b/agent/app/service/docker.go @@ -18,7 +18,8 @@ import ( "github.com/1Panel-dev/1Panel/agent/utils/common" "github.com/1Panel-dev/1Panel/agent/utils/controller" "github.com/1Panel-dev/1Panel/agent/utils/docker" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" + + dockerfirewall "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" ) const dockerNftablesMinVersion = "29.0.0" @@ -84,7 +85,7 @@ func (u *DockerService) UpdateFirewallBackend(backend string) error { return fmt.Errorf("Docker Engine %s or later is required for the nftables firewall backend", dockerNftablesMinVersion) } if backend == constant.FirewallProviderNftables { - if err := docker_guard.CheckIPv4Forwarding(); err != nil { + if err := dockerfirewall.CheckIPv4Forwarding(); err != nil { return err } } @@ -282,7 +283,8 @@ func (u *DockerService) UpdateConf(req dto.SettingUpdate, withRestart bool) erro delete(daemonMap, "ipv6") delete(daemonMap, "fixed-cidr-v6") delete(daemonMap, "ip6tables") - if configuredDockerFirewallBackend() != constant.FirewallProviderNftables { + backend, _ := settingRepo.GetValueByKey(constant.FirewallDockerBackendKey) + if !strings.EqualFold(strings.TrimSpace(backend), constant.FirewallProviderNftables) { delete(daemonMap, "experimental") } } diff --git a/agent/app/service/entry.go b/agent/app/service/entry.go index c3ba1186bbe6..7a6b3d554776 100644 --- a/agent/app/service/entry.go +++ b/agent/app/service/entry.go @@ -37,8 +37,9 @@ var ( clamRepo = repo.NewIClamRepo() monitorRepo = repo.NewIMonitorRepo() - settingRepo = repo.NewISettingRepo() - backupRepo = repo.NewIBackupRepo() + settingRepo = repo.NewISettingRepo() + forwardingRuleRepo = repo.NewIForwardingRuleRepo() + backupRepo = repo.NewIBackupRepo() websiteRepo = repo.NewIWebsiteRepo() websiteDomainRepo = repo.NewIWebsiteDomainRepo() diff --git a/agent/app/service/firewall.go b/agent/app/service/firewall.go index 52b646c4df52..647829167092 100644 --- a/agent/app/service/firewall.go +++ b/agent/app/service/firewall.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "slices" "sort" "strconv" "strings" @@ -18,28 +19,23 @@ import ( "github.com/1Panel-dev/1Panel/agent/constant" "github.com/1Panel-dev/1Panel/agent/global" "github.com/1Panel-dev/1Panel/agent/i18n" - "github.com/1Panel-dev/1Panel/agent/utils/cmd" - "github.com/1Panel-dev/1Panel/agent/utils/controller" "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" 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" - firewallsync "github.com/1Panel-dev/1Panel/agent/utils/firewall/sync" "github.com/google/uuid" "gorm.io/gorm" ) type FirewallService struct { rules repo.IFirewallRuleRepo - adapters firewallRuleRuntimeResolver + adapters map[filter.Provider]filter.Adapter forwardingSync firewallDatabaseSyncAdapter dockerSync firewallDatabaseSyncAdapter selectedProvider func(context.Context) (filter.Provider, error) requiredPorts func() ([]firewall.PortWhitelist, error) - iptablesHelper *iptables_helper.Manager cleanupBackend func(string) error cleanupInactiveBackend func(string) error resetBackend func(string, bool) error @@ -49,11 +45,6 @@ type FirewallService struct { baseClient func() (lifecycle.Client, error) } -type firewallRuleRuntimeResolver interface { - Resolve(filter.Provider) (*filterruntime.Engine, error) - Providers() []filter.Provider -} - var firewallRuleMutationMu sync.Mutex var ( @@ -82,47 +73,44 @@ type IFirewallService interface { CurrentRuleSyncTask() (dto.FirewallRuleSyncTask, error) } -func NewIFirewallService() IFirewallService { - return newFirewallService() +type firewallRuleDeleteItem struct { + index int + stored model.FirewallRule + rules []filter.DesiredRule } -func newFirewallService() *FirewallService { - return &FirewallService{ - rules: repo.NewIFirewallRuleRepo(), - adapters: filterruntime.NewRegistry(firewallRuleSnapshotPolicy), - forwardingSync: newForwardingService(), - dockerSync: newDockerPortGuardService(), - selectedProvider: firewallRuleSelectedProvider, - requiredPorts: LoadRequiredFirewallPortWhiteList, - iptablesHelper: newIptablesHelperManager(), - cleanupBackend: cleanupSystemBackend, - cleanupInactiveBackend: cleanupInactiveSystemBackend, - resetBackend: resetServiceFirewallBackend, - dockerActive: firewallDockerActive, - restoreForwarding: func(ctx context.Context) error { - return newForwardingService().Restore(ctx) - }, - restoreDockerGuard: ReconcileDockerPortGuard, - baseClient: selectedSystemFirewallClient, - } +type preparedManagedUpdate struct { + Stored model.FirewallRule + Before filter.DesiredRule + After filter.FirewallRule + RuleSet filter.RuleSet + Observed filter.ObservedRule + Runtime filter.Adapter } -func firewallDockerActive() (bool, error) { - if !cmd.Which("docker") { - return false, nil +func (s *FirewallService) UpdatePanelPort(ctx context.Context, oldPort, port uint) error { + if oldPort == 0 || oldPort > 65535 || port == 0 || port > 65535 { + return fmt.Errorf("invalid panel port transition %d -> %d", oldPort, port) + } + if LoadPanelPort() != strconv.Itoa(int(oldPort)) { + return fmt.Errorf("panel port changed before firewall update") + } + if oldPort == port { + return nil } - return controller.CheckActive("docker") + return updateSystemAccessPortWhitelist(ctx, firewall.PortWhitelistTypePanel, []string{strconv.Itoa(int(port))}) } func (s *FirewallService) LoadBaseInfo(chainGroup string) (dto.FirewallSubsystemStatus, error) { status := dto.FirewallSubsystemStatus{Version: "-", Name: "-", Backend: "-"} status.LifecycleTaskID = currentFirewallLifecycleTaskID() - if selected := configuredSystemFirewallBackend(); selected != "" { + selected, _ := settingRepo.GetValueByKey(constant.FirewallSystemBackendKey) + if selected = strings.TrimSpace(selected); selected != "" { status.Name, status.Backend = selected, selected } loadClient := s.baseClient if loadClient == nil { - loadClient = selectedSystemFirewallClient + loadClient = NewSelectedSystemFirewallClient } client, err := loadClient() if err != nil { @@ -145,61 +133,17 @@ func (s *FirewallService) LoadBaseInfo(chainGroup string) (dto.FirewallSubsystem status.Name, status.Backend = runtimeStatus.Name, runtimeStatus.Name status.Version, status.PingStatus = runtimeStatus.Version, firewall.LoadPingStatus() status.IsActive = runtimeStatus.IsActive - if supportsManagedFilterChains(runtimeStatus.Name) { - initialized, bound, err := loadFirewallInitStatus(runtimeStatus.Name, chainGroup) + if runtimeStatus.Name == constant.FirewallProviderIptables || runtimeStatus.Name == constant.FirewallProviderNftables { + overview, err := loadSystemFirewallOverview(runtimeStatus.Name, chainGroup) if err != nil { return status, err } - status.IsInit, status.IsBind = initialized, bound - status.IPv4 = loadSystemFirewallFamilyInfo(status.Name, constant.FirewallFamilyIPv4) - status.IPv6 = loadSystemFirewallFamilyInfo(status.Name, constant.FirewallFamilyIPv6) + status.IsInit, status.IsBind = overview.IsInit, overview.IsBind + status.IPv4, status.IPv6 = overview.IPv4, overview.IPv6 } return status, nil } -type firewallLifecycleClient struct{ lifecycle.Client } - -func (c firewallLifecycleClient) Start() error { - filterruntime.InvalidateInventory() - defer filterruntime.InvalidateInventory() - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - return c.Client.Start() -} - -func (c firewallLifecycleClient) Stop() error { - filterruntime.InvalidateInventory() - defer filterruntime.InvalidateInventory() - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - return c.Client.Stop() -} - -func (c firewallLifecycleClient) Restart() error { - filterruntime.InvalidateInventory() - defer filterruntime.InvalidateInventory() - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - return c.Client.Restart() -} - -func currentFirewallLifecycleTaskID() string { - firewallLifecycleTaskMu.Lock() - defer firewallLifecycleTaskMu.Unlock() - return firewallLifecycleTaskID -} - -func lockFirewallLifecycleIdle() error { - if !firewallLifecycleTaskMu.TryLock() { - return buserr.New("TaskIsExecuting") - } - if firewallLifecycleTaskID != "" { - firewallLifecycleTaskMu.Unlock() - return buserr.New("TaskIsExecuting") - } - return nil -} - func (s *FirewallService) QueueFirewallOperation(request dto.FirewallLifecycleOperation) (dto.FirewallLifecycleOperationResponse, error) { response := dto.FirewallLifecycleOperationResponse{} if request.Operation == "disableBanPing" || request.Operation == "enableBanPing" { @@ -224,7 +168,7 @@ func (s *FirewallService) QueueFirewallOperation(request dto.FirewallLifecycleOp } loadClient := s.baseClient if loadClient == nil { - loadClient = selectedSystemFirewallClient + loadClient = NewSelectedSystemFirewallClient } client, err := loadClient() if err != nil { @@ -236,10 +180,10 @@ func (s *FirewallService) QueueFirewallOperation(request dto.FirewallLifecycleOp return response, s.OperateFirewall(request) } operation, label := task.TaskExec, "Start" - switch lifecycle.Operation(request.Operation) { - case lifecycle.OperationStop: + switch request.Operation { + case string(lifecycle.OperationStop): label = "Stop" - case lifecycle.OperationRestart: + case string(lifecycle.OperationRestart): operation, label = task.TaskRestart, task.TaskRestart } name := task.GetTaskName(client.Name(), label, task.TaskScopeFirewall) @@ -275,159 +219,12 @@ func (s *FirewallService) QueueFirewallOperation(request dto.FirewallLifecycleOp return dto.FirewallLifecycleOperationResponse{TaskID: taskItem.TaskID, Queued: true}, nil } -func runFirewallLifecycleAction(t *task.Task, name string, action func() error) error { - t.Log(i18n.GetWithName("TaskStart", name)) - started := time.Now() - err := t.TaskCtx.Err() - if err == nil { - err = action() - } - t.LogWithStatus(fmt.Sprintf("%s (%.2fs)", name, time.Since(started).Seconds()), err) - return err -} - -func (s *FirewallService) runFirewallLifecycleTask(t *task.Task, client lifecycle.Client, request dto.FirewallLifecycleOperation) error { - ctx := t.TaskCtx - provider := filter.Provider(client.Name()) - operator := lifecycle.NewOperator(firewallLifecycleClient{client}) - operator.RunAction = func(operation, name string, action func() error) error { - return runFirewallLifecycleAction(t, task.GetTaskName(name, operation, ""), action) - } - operationErr := operator.Operate(lifecycle.Operation(request.Operation), request.WithDockerRestart, func(lifecycle.Client) error { - rulesErr := runFirewallLifecycleAction(t, i18n.GetWithName("FirewallRestoreRulesStep", client.Name()), func() error { - return s.restoreStoredFirewallRules(ctx, provider, t) - }) - whitelistErr := runFirewallLifecycleAction(t, i18n.GetMsgByKey("FirewallSyncWhitelistStep"), func() error { - return s.SyncPortWhitelist(ctx) - }) - return errors.Join(rulesErr, whitelistErr) - }) - if request.Operation == string(lifecycle.OperationStop) { - return operationErr - } - var recoveryErr *lifecycle.CompletedOperationError - if operationErr != nil && !errors.As(operationErr, &recoveryErr) { - return operationErr - } - var forwardingErr error - if provider == filter.ProviderFirewalld { - forwardingErr = runFirewallLifecycleAction(t, i18n.GetMsgByKey("FirewallRestoreForwardingRulesStep"), func() error { - if s.restoreForwarding != nil { - return s.restoreForwarding(ctx) - } - return newForwardingService().Restore(ctx) - }) - } - dockerErr := runFirewallLifecycleAction(t, i18n.GetMsgByKey("FirewallInspectDockerGuardStep"), func() error { - if provider == filter.ProviderFirewalld { - active := s.dockerActive - if active == nil { - active = firewallDockerActive - } - running, err := active() - if err != nil || !running { - return err - } - } - if s.restoreDockerGuard != nil { - return s.restoreDockerGuard(ctx) - } - return ReconcileDockerPortGuard(ctx) - }) - return errors.Join(operationErr, forwardingErr, dockerErr) -} - -func (s *FirewallService) OperateFirewall(request dto.FirewallLifecycleOperation) error { - switch request.Operation { - case "disableBanPing": - if err := firewall.UpdatePingStatus("0"); err != nil { - return err - } - return settingRepo.Update(constant.FirewallPingStatusKey, constant.StatusDisable) - case "enableBanPing": - if err := firewall.UpdatePingStatus("1"); err != nil { - return err - } - return settingRepo.Update(constant.FirewallPingStatusKey, constant.StatusEnable) - } - baseClient := s.baseClient - if baseClient == nil { - baseClient = selectedSystemFirewallClient - } - client, err := baseClient() - if err != nil { - return err - } - operation := lifecycle.Operation(request.Operation) - operationErr := lifecycle.NewOperator(firewallLifecycleClient{client}).Operate(operation, request.WithDockerRestart, s.restoreFirewallAfterStart) - restoreFirewalld := client.Name() == lifecycle.ProviderFirewalld && - (operation == lifecycle.OperationStart || operation == lifecycle.OperationRestart) - if operation != lifecycle.OperationStart && operation != lifecycle.OperationRestart { - return operationErr - } - if operationErr != nil { - var completedErr *lifecycle.CompletedOperationError - var dockerRestartErr *lifecycle.DockerRestartError - if !errors.As(operationErr, &completedErr) && !errors.As(operationErr, &dockerRestartErr) { - return operationErr - } - if global.LOG != nil { - global.LOG.Warnf("firewall %s completed with post-start recovery errors: %v", operation, operationErr) - } - } - if restoreFirewalld { - restoreErr := s.restoreFirewalldRuntimeDependents(context.Background(), operation) - if restoreErr != nil && global.LOG != nil { - global.LOG.Errorf("restore firewalld runtime dependents after %s failed: %v", operation, restoreErr) - } - return nil - } - ReconcileDockerPortGuardBestEffort(context.Background()) - return nil -} - -func (s *FirewallService) UpdatePanelPort(ctx context.Context, oldPort, port uint) error { - if oldPort == 0 || oldPort > 65535 || port == 0 || port > 65535 { - return fmt.Errorf("invalid panel port transition %d -> %d", oldPort, port) - } - if LoadPanelPort() != strconv.Itoa(int(oldPort)) { - return fmt.Errorf("panel port changed before firewall update") - } - if oldPort == port { - return nil - } - return updateSystemAccessPortWhitelist(ctx, firewall.PortWhitelistTypePanel, []string{strconv.Itoa(int(port))}) -} - -func (s *FirewallService) restoreFirewalldRuntimeDependents(ctx context.Context, operation lifecycle.Operation) error { - restoreForwarding := s.restoreForwarding - if restoreForwarding == nil { - restoreForwarding = func(ctx context.Context) error { return newForwardingService().Restore(ctx) } - } - restoreDockerGuard := s.restoreDockerGuard - if restoreDockerGuard == nil { - restoreDockerGuard = ReconcileDockerPortGuard - } - dockerActive := s.dockerActive - if dockerActive == nil { - dockerActive = firewallDockerActive - } - - active, err := dockerActive() - restoreErr := restoreFirewalldDependents( - ctx, fmt.Sprintf("after firewalld %s", operation), err == nil && active, restoreForwarding, restoreDockerGuard, - ) - if err != nil { - return errors.Join(fmt.Errorf("check Docker status after firewalld %s: %w", operation, err), restoreErr) - } - return restoreErr -} - func (s *FirewallService) OperateFilterChain(request dto.FilterChainOperation) error { - provider, err := selectedSystemFirewallProvider() + client, err := NewSelectedSystemFirewallClient() if err != nil { return err } + provider := client.Name() if err := s.operateFilterChainBase(provider, request); err != nil { return err } @@ -440,17 +237,16 @@ func (s *FirewallService) OperateFilterChain(request dto.FilterChainOperation) e return errors.Join(rulesErr, whitelistErr) } -func (s *FirewallService) QueueFilterChainInitialization( - request dto.FilterChainOperation, -) (dto.FilterChainOperationResponse, error) { +func (s *FirewallService) QueueFilterChainInitialization(request dto.FilterChainOperation) (dto.FilterChainOperationResponse, error) { if request.Operate != string(firewall.BaseOperationInit) { return dto.FilterChainOperationResponse{}, fmt.Errorf("only filter chain initialization can be queued") } - provider, err := selectedSystemFirewallProvider() + client, err := NewSelectedSystemFirewallClient() if err != nil { return dto.FilterChainOperationResponse{}, err } - if !supportsManagedFilterChains(provider) { + provider := client.Name() + if provider != constant.FirewallProviderIptables && provider != constant.FirewallProviderNftables { return dto.FilterChainOperationResponse{}, fmt.Errorf("filter chain operations are not supported for %s", provider) } if err := task.CheckScopeTaskIsExecuting(task.TaskScopeFirewall, 0); err != nil { @@ -483,34 +279,7 @@ func (s *FirewallService) QueueFilterChainInitialization( return dto.FilterChainOperationResponse{TaskID: taskItem.TaskID, Queued: true}, nil } -func (s *FirewallService) operateFilterChainBase(provider string, request dto.FilterChainOperation) error { - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - return s.operateFilterChainBaseLocked(provider, request) -} - -func (s *FirewallService) operateFilterChainBaseLocked(provider string, request dto.FilterChainOperation) error { - filterruntime.InvalidateInventory() - defer filterruntime.InvalidateInventory() - if err := s.checkSelectedProvider(context.Background(), filter.Provider(provider)); err != nil { - return err - } - if !supportsManagedFilterChains(provider) { - return fmt.Errorf("filter chain operations are not supported for %s", provider) - } - if provider == constant.FirewallProviderNftables { - if err := newNftablesHelperManager().Operate(firewall.BaseOperation(request.Operate)); err != nil { - return err - } - } else if err := s.iptablesHelper.Operate(firewall.BaseOperation(request.Operate)); err != nil { - return err - } - return nil -} - func (s *FirewallService) Reset(ctx context.Context, request dto.FirewallRuleReset) (dto.FirewallRuleResetResponse, error) { - filterruntime.InvalidateInventory() - defer filterruntime.InvalidateInventory() if err := lockFirewallLifecycleIdle(); err != nil { return dto.FirewallRuleResetResponse{}, err } @@ -518,29 +287,26 @@ func (s *FirewallService) Reset(ctx context.Context, request dto.FirewallRuleRes firewallRuleMutationMu.Lock() defer firewallRuleMutationMu.Unlock() + selected, err := s.selectedProvider(ctx) + if err != nil { + return dto.FirewallRuleResetResponse{}, err + } provider := request.Provider - selected := provider if provider == "" { - var err error - selected, err = s.selectedProvider(ctx) - if err != nil { - return dto.FirewallRuleResetResponse{}, err - } provider = selected - } else if isDirectFirewallProvider(provider) { - var err error - selected, err = s.selectedProvider(ctx) - if err != nil { - return dto.FirewallRuleResetResponse{}, err - } } stored, err := s.rules.List(ctx) if err != nil { return dto.FirewallRuleResetResponse{}, err } + if selected == provider && len(stored) > 0 { + if err := s.saveFirewallResetOrder(ctx, provider, stored); err != nil { + return dto.FirewallRuleResetResponse{}, err + } + } if provider == filter.ProviderIptables || provider == filter.ProviderNftables { cleanup := s.cleanupBackend - if isDirectFirewallProvider(selected) && selected != provider { + if (selected == filter.ProviderIptables || selected == filter.ProviderNftables) && selected != provider { cleanup = s.cleanupInactiveBackend if cleanup == nil { cleanup = cleanupInactiveSystemBackend @@ -574,7 +340,7 @@ func (s *FirewallService) Reset(ctx context.Context, request dto.FirewallRuleRes } resetErr := reset(string(provider), restartDocker) if resetErr != nil { - var dockerRestartErr *lifecycle.DockerRestartError + var dockerRestartErr *firewallDockerRestartError if provider != filter.ProviderFirewalld || !errors.As(resetErr, &dockerRestartErr) { return dto.FirewallRuleResetResponse{}, resetErr } @@ -598,61 +364,6 @@ func (s *FirewallService) Reset(ctx context.Context, request dto.FirewallRuleRes return dto.FirewallRuleResetResponse{Removed: len(stored), Disabled: true}, nil } -func isDirectFirewallProvider(provider filter.Provider) bool { - return provider == filter.ProviderIptables || provider == filter.ProviderNftables -} - -func resetServiceFirewallBackend(provider string, withDockerRestart bool) error { - client, err := lifecycle.NewClientFor(provider) - if err != nil { - return err - } - return resetServiceFirewallClient(client, withDockerRestart, func( - client lifecycle.Client, - restartDocker bool, - prepareStop func() error, - ) error { - return lifecycle.NewOperator(client).StopWithPrepare(restartDocker, prepareStop) - }) -} - -func resetServiceFirewallClient( - client lifecycle.Client, - withDockerRestart bool, - stop func(lifecycle.Client, bool, func() error) error, -) error { - resetter, ok := client.(lifecycle.Resetter) - if !ok { - return fmt.Errorf("firewall provider %s does not support reset", client.Name()) - } - if resetBeforeStop, ok := client.(lifecycle.PreStopResetter); ok { - if err := stop(client, withDockerRestart, resetBeforeStop.ResetBeforeStop); err != nil { - return err - } - return nil - } - return resetter.Reset() -} - -func restoreFirewalldDependents( - ctx context.Context, - reason string, - restoreDocker bool, - restoreForwarding func(context.Context) error, - restoreDockerGuard func(context.Context) error, -) error { - var errs []error - if err := restoreForwarding(ctx); err != nil { - errs = append(errs, fmt.Errorf("restore port forwarding %s: %w", reason, err)) - } - if restoreDocker { - if err := restoreDockerGuard(ctx); err != nil { - errs = append(errs, fmt.Errorf("restore Docker port guard %s: %w", reason, err)) - } - } - return errors.Join(errs...) -} - func (s *FirewallService) Inventory(ctx context.Context, request dto.FirewallRuleInventory) (dto.FirewallRuleInventoryResponse, error) { requestedScopes := request.Scopes if len(requestedScopes) == 0 && request.Scope.Provider != "" { @@ -665,12 +376,13 @@ func (s *FirewallService) Inventory(ctx context.Context, request dto.FirewallRul for index, requested := range requestedScopes { scopes[index] = requested.Normalize() } - if len(scopes) == 1 && isCombinedUFWInventoryScope(scopes[0]) { + if len(scopes) == 1 && scopes[0].Provider == filter.ProviderUFW && scopes[0].Family == filter.FamilyInet && scopes[0].Table == "" && + scopes[0].Zone == "" && scopes[0].Chain == filter.UFWInputChain && scopes[0].Direction == filter.DirectionInput { scope := scopes[0] if err := s.checkSelectedProvider(ctx, scope.Provider); err != nil { return dto.FirewallRuleInventoryResponse{}, err } - runtime, err := s.adapters.Resolve(scope.Provider) + runtime, err := s.firewallAdapter(scope.Provider) if err != nil { return dto.FirewallRuleInventoryResponse{}, err } @@ -694,7 +406,7 @@ func (s *FirewallService) Inventory(ctx context.Context, request dto.FirewallRul if err := s.checkSelectedProvider(ctx, provider); err != nil { return dto.FirewallRuleInventoryResponse{}, err } - runtime, err := s.adapters.Resolve(provider) + runtime, err := s.firewallAdapter(provider) if err != nil { return dto.FirewallRuleInventoryResponse{}, err } @@ -702,245 +414,54 @@ func (s *FirewallService) Inventory(ctx context.Context, request dto.FirewallRul if err != nil { return dto.FirewallRuleInventoryResponse{}, err } - desiredByScope, failures := s.desiredFirewallRulesByScope(ctx, stored, provider) + desiredByScope, failures := s.desiredFirewallRulesByScope(ctx, stored, runtime) response := dto.FirewallRuleInventoryResponse{Items: failures} - runtime = runtime.NewObservationSession() unavailable := make(map[filter.Family]error) - for _, scope := range scopes { - var snapshot filter.Snapshot - err := unavailable[scope.Family] - if err == nil { - snapshot, err = runtime.ObserveInventory(ctx, scope, request.Refresh) - } + for _, group := range firewallScopeReadGroups(scopes) { + snapshots, err := readFirewallRuleScopes(runtime, ctx, group) if errors.Is(err, filter.ErrFamilyUnavailable) { - if unavailable[scope.Family] == nil { - response.Notices = append(response.Notices, filter.ScopeNotice{Code: filter.ScopeNoticeFamilyUnavailable, Values: []string{string(scope.Family), err.Error()}}) - unavailable[scope.Family] = err - } - for _, desired := range desiredByScope[scope.Key()] { - response.Items = append(response.Items, filter.InventoryItem{ - Rule: desired.Rule, Desired: &desired, State: filter.InventoryStateDrifted, - Match: filter.InventoryMatchNone, Error: err.Error(), - }) + for _, scope := range group { + if unavailable[scope.Family] == nil { + response.Notices = append(response.Notices, filter.ScopeNotice{Code: filter.ScopeNoticeFamilyUnavailable, Values: []string{string(scope.Family), err.Error()}}) + unavailable[scope.Family] = err + } + for _, desired := range desiredByScope[scope.Key()] { + response.Items = append(response.Items, filter.InventoryItem{ + Rule: desired.Rule, Desired: &desired, State: filter.InventoryStateDrifted, + Match: filter.InventoryMatchNone, Error: err.Error(), + }) + } } continue } if err != nil { return dto.FirewallRuleInventoryResponse{}, err } - desired := desiredByScope[scope.Key()] - items, err := filter.MergeInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desired}) - if err != nil { - return dto.FirewallRuleInventoryResponse{}, err + for _, snapshot := range snapshots { + scope := snapshot.Scope + desired := desiredByScope[scope.Key()] + items, err := mergeFirewallInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desired}) + if err != nil { + return dto.FirewallRuleInventoryResponse{}, err + } + response.Items = append(response.Items, items...) + response.Notices = append(response.Notices, snapshot.Notices...) } - response.Items = append(response.Items, items...) - response.Notices = append(response.Notices, snapshot.Notices...) } return finalizeFirewallInventory(response, request), nil } -func finalizeFirewallInventory( - response dto.FirewallRuleInventoryResponse, - request dto.FirewallRuleInventory, -) dto.FirewallRuleInventoryResponse { - provider := request.Scope.Provider - if len(request.Scopes) > 0 { - provider = request.Scopes[0].Provider - } - response.IPv4Range, response.IPv6Range = filter.InventoryPositionRanges(provider, response.Items) - response.AllTotal = int64(len(response.Items)) - for _, item := range response.Items { - if isDeletableManagedInventoryItem(item) { - response.ManagedTotal++ +func (s *FirewallService) LoadFirewallNativeDetail(ctx context.Context, request dto.FirewallNativeDetail) (string, error) { + provider := filter.Provider(strings.ToLower(strings.TrimSpace(string(request.Provider)))) + nativeKind := filter.NativeKind(strings.ToLower(strings.TrimSpace(string(request.NativeKind)))) + switch provider { + case filter.ProviderFirewalld: + if nativeKind != filter.NativeKindZoneService { + return "", fmt.Errorf("%w: firewalld detail kind %q", filter.ErrInvalidRule, nativeKind) } - } - filtered := make([]filter.InventoryItem, 0, len(response.Items)) - for _, item := range response.Items { - if matchesFirewallInventoryRequest(item, request) { - filtered = append(filtered, item) - } - } - response.Total = int64(len(filtered)) - if request.All { - response.Items = filtered - return response - } - page, pageSize := max(1, request.Page), max(1, request.PageSize) - start := (page - 1) * pageSize - if start >= len(filtered) { - response.Items = make([]filter.InventoryItem, 0) - return response - } - end := min(start+pageSize, len(filtered)) - response.Items = filtered[start:end] - return response -} - -func matchesFirewallInventoryRequest(item filter.InventoryItem, request dto.FirewallRuleInventory) bool { - if slicesContains(request.ExcludeChains, item.Rule.Scope.Chain) { - return false - } - if len(request.Families) > 0 && !matchesFirewallInventoryFamily(item.Rule, request.Families) { - return false - } - if len(request.Actions) > 0 && !matchesFirewallInventoryAction(item.Rule.Action, request.Actions) { - return false - } - if len(request.States) > 0 && !slicesContains(request.States, item.State) { - return false - } - keyword := strings.ToLower(strings.TrimSpace(request.Info)) - if keyword == "" { - return true - } - rule := item.Rule - values := []string{ - firewallInventoryProtocol(rule), rule.SourceAddress, rule.SourcePort, rule.DestinationAddress, - rule.DestinationPort, rule.Description, string(rule.Action), string(item.State), - } - if item.Observed != nil { - values = append(values, item.Observed.Rule.Description) - } - if item.Desired != nil { - values = append(values, item.Desired.Rule.Description) - } - for _, value := range values { - if strings.Contains(strings.ToLower(value), keyword) { - return true - } - } - return false -} - -func matchesFirewallInventoryFamily(rule filter.FirewallRule, families []filter.Family) bool { - for _, family := range families { - if rule.Scope.Family != filter.FamilyInet && rule.Scope.Family == family { - return true - } - if rule.Scope.Family == filter.FamilyInet && - (rule.SourceAddress == "" || (family == filter.FamilyIPv6) == strings.Contains(rule.SourceAddress, ":")) { - return true - } - } - return false -} - -func matchesFirewallInventoryAction(action filter.Action, actions []string) bool { - for _, requested := range actions { - if requested == "accept" && action == filter.ActionAccept { - return true - } - if requested == "deny" && (action == filter.ActionDrop || action == filter.ActionReject) { - return true - } - } - return false -} - -func firewallInventoryProtocol(rule filter.FirewallRule) string { - if rule.NativeKind == filter.NativeKindZoneService { - return "service" - } - if rule.NativeKind == filter.NativeKindUFWApplication && rule.Protocol == "" { - return "app" - } - if rule.Scope.Provider == filter.ProviderUFW && rule.Protocol == "all" && rule.DestinationPort != "" { - return "tcp/udp" - } - return rule.Protocol -} - -func slicesContains[T comparable](values []T, target T) bool { - for _, value := range values { - if value == target { - return true - } - } - return false -} - -func isDeletableManagedInventoryItem(item filter.InventoryItem) bool { - if item.Desired == nil || item.Desired.Protected || item.State == filter.InventoryStateProtected { - return false - } - if item.Desired.Origin != filter.RuleOriginCreated && item.Desired.Origin != filter.RuleOriginAdopted { - return false - } - if isIptablesSystemPresetInventoryScope(item.Rule.Scope) { - return false - } - return item.State != filter.InventoryStateDrifted || - (item.Match == filter.InventoryMatchMissing && item.Observed == nil) -} - -func isIptablesSystemPresetInventoryScope(scope filter.Scope) bool { - return (scope.Provider == filter.ProviderIptables || scope.Provider == filter.ProviderNftables) && - (scope.Chain == filter.BasicBeforeChain || scope.Chain == filter.BasicAfterChain) -} - -func isCombinedUFWInventoryScope(scope filter.Scope) bool { - scope = scope.Normalize() - return scope.Provider == filter.ProviderUFW && scope.Family == filter.FamilyInet && scope.Table == "" && - scope.Zone == "" && scope.Chain == filter.UFWInputChain && scope.Direction == filter.DirectionInput -} - -func (s *FirewallService) combinedUFWInventory( - ctx context.Context, - runtime *filterruntime.Engine, - scope filter.Scope, - refresh bool, -) (dto.FirewallRuleInventoryResponse, error) { - scopes := []filter.Scope{scope, scope} - scopes[0].Family = filter.FamilyIPv4 - scopes[1].Family = filter.FamilyIPv6 - snapshots, err := runtime.ObserveInventoryScopes(ctx, scopes, refresh) - if err != nil { - return dto.FirewallRuleInventoryResponse{}, err - } - if len(snapshots) != len(scopes) { - return dto.FirewallRuleInventoryResponse{}, fmt.Errorf("%w: incomplete UFW multi-family inventory", filter.ErrAdapterUnavailable) - } - - stored, err := s.rules.List(ctx) - if err != nil { - return dto.FirewallRuleInventoryResponse{}, err - } - desiredByScope, failures := s.desiredFirewallRulesByScope(ctx, stored, scope.Provider) - response := dto.FirewallRuleInventoryResponse{Items: failures} - seenNotices := make(map[string]struct{}) - for index, snapshot := range snapshots { - if snapshot.Scope.Key() != scopes[index].Key() { - return dto.FirewallRuleInventoryResponse{}, fmt.Errorf("%w: unexpected UFW inventory scope %q", filter.ErrInvalidScope, snapshot.Scope.Key()) - } - desired := desiredByScope[snapshot.Scope.Key()] - items, err := filter.MergeInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desired}) - if err != nil { - return dto.FirewallRuleInventoryResponse{}, err - } - response.Items = append(response.Items, items...) - for _, notice := range snapshot.Notices { - key := string(notice.Code) + "\x00" + strings.Join(notice.Values, "\x00") - if _, exists := seenNotices[key]; exists { - continue - } - seenNotices[key] = struct{}{} - response.Notices = append(response.Notices, notice) - } - } - return response, nil -} - -func (s *FirewallService) LoadFirewallNativeDetail(ctx context.Context, request dto.FirewallNativeDetail) (string, error) { - provider := filter.Provider(strings.ToLower(strings.TrimSpace(string(request.Provider)))) - nativeKind := filter.NativeKind(strings.ToLower(strings.TrimSpace(string(request.NativeKind)))) - switch provider { - case filter.ProviderFirewalld: - if nativeKind != filter.NativeKindZoneService { - return "", fmt.Errorf("%w: firewalld detail kind %q", filter.ErrInvalidRule, nativeKind) - } - case filter.ProviderUFW: - if nativeKind != filter.NativeKindUFWApplication { - return "", fmt.Errorf("%w: UFW detail kind %q", filter.ErrInvalidRule, nativeKind) + case filter.ProviderUFW: + if nativeKind != filter.NativeKindUFWApplication { + return "", fmt.Errorf("%w: UFW detail kind %q", filter.ErrInvalidRule, nativeKind) } default: return "", fmt.Errorf("%w: native details for %s", filter.ErrUnsupportedScope, provider) @@ -948,779 +469,123 @@ func (s *FirewallService) LoadFirewallNativeDetail(ctx context.Context, request if err := s.checkSelectedProvider(ctx, provider); err != nil { return "", err } - runtime, err := s.adapters.Resolve(provider) + runtime, err := s.firewallAdapter(provider) if err != nil { return "", err } - return runtime.NativeDetail(ctx, request.Name, request.Permanent) -} - -func applySelectedProviderScopeDefaults(rule filter.FirewallRule, selected filter.Provider) filter.FirewallRule { - scope := rule.Scope - if scope.Provider == "" { - scope.Provider = selected - } - if scope.Provider != selected { - return rule - } - if scope.Direction == "" { - scope.Direction = filter.DirectionInput - } - if scope.Family == "" { - scope.Family = defaultFirewallRuleFamily(rule, selected) - } - switch selected { - case filter.ProviderIptables, filter.ProviderNftables: - if scope.Table == "" { - scope.Table = "filter" - } - if scope.Chain == "" { - scope.Chain = filter.IptablesInputChain - } - case filter.ProviderFirewalld: - if scope.Zone == "" { - scope.Zone = filter.FirewalldInputZone - } - case filter.ProviderUFW: - if scope.Chain == "" { - scope.Chain = filter.UFWInputChain - } - } - rule.Scope = scope - return rule -} - -func defaultFirewallRuleFamily(rule filter.FirewallRule, provider filter.Provider) filter.Family { - if strings.EqualFold(strings.TrimSpace(rule.Protocol), "icmpv6") || - strings.Contains(rule.SourceAddress, ":") || strings.Contains(rule.DestinationAddress, ":") { - return filter.FamilyIPv6 - } - if provider == filter.ProviderFirewalld { - return filter.FamilyInet - } - return filter.FamilyIPv4 -} - -type preparedFirewallRuleCreate struct { - request dto.FirewallRuleCreateItem - runtime *filterruntime.Engine -} - -func (s *FirewallService) Create( - ctx context.Context, - request dto.FirewallRuleCreate, -) (dto.FirewallRuleCreateResponse, error) { - taskItem, err := task.NewTask(firewallTaskName(task.TaskCreate, firewallTaskHost, ""), task.TaskCreate, task.TaskScopeFirewall, "", 0) - if err != nil { - return dto.FirewallRuleCreateResponse{}, err - } - taskItem.AddSubTaskWithOps(i18n.GetMsgByKey("FirewallCreateRulesStep"), func(t *task.Task) error { - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - _, err := s.createRules(t.TaskCtx, request, t) - return err - }, nil, 0, 0) - if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { - taskItem.LogFailedWithErr(taskItem.Name, err) - closeUnstartedFirewallTask(taskItem) - return dto.FirewallRuleCreateResponse{}, fmt.Errorf("save firewall creation task: %w", err) - } - go func() { - if err := taskItem.Execute(); err != nil && global.LOG != nil { - global.LOG.Errorf("firewall creation task %s failed: %v", taskItem.TaskID, err) - } - }() - return dto.FirewallRuleCreateResponse{TaskID: taskItem.TaskID, Queued: true}, nil -} - -func (s *FirewallService) createRules(ctx context.Context, request dto.FirewallRuleCreate, t *task.Task) (result dto.FirewallRuleCreateResponse, taskErr error) { - var firstFailure error - defer func() { - if t != nil { - t.Log(i18n.GetMsgWithMap("FirewallCreateRulesResult", map[string]interface{}{ - "succeeded": result.Succeeded, "failed": result.Failed, "skipped": result.Skipped, - })) - } - if taskErr == nil { - taskErr = firstFailure - } - }() - type itemOrigin struct { - index, part, count int - rule filter.FirewallRule - } - describe := func(rule filter.FirewallRule) string { - return fmt.Sprintf("%s %s %s:%s -> %s:%s %s", rule.Scope.Family, rule.Protocol, - rule.SourceAddress, rule.SourcePort, rule.DestinationAddress, rule.DestinationPort, rule.Action) - } - record := func(origin itemOrigin, status string, err error) { - rule := origin.rule - label := fmt.Sprintf("[%d/%d]", origin.index+1, len(request.Items)) - if origin.count > 1 { - label += fmt.Sprintf("[%d/%d]", origin.part+1, origin.count) - } - label += fmt.Sprintf(" %s %s", rule.Scope.Provider, describe(rule)) - switch status { - case "succeeded": - result.Succeeded++ - if t != nil { - t.LogSuccess(label) - } - case "failed": - if firstFailure == nil { - firstFailure = err - } - result.Failed++ - if t != nil { - t.LogFailedWithErr(label, err) - } - case "skipped": - result.Skipped++ - if t != nil { - t.Logf("%s %s: %v", label, i18n.GetMsgByKey("FirewallCreateRuleSkipped"), err) - } - } - if err != nil { - result.Errors = append(result.Errors, dto.FirewallRuleCreateFailure{ - Index: origin.index, Status: status, Rule: rule, Error: err.Error(), - }) - } - } - selected, err := s.selectedProvider(ctx) - if err != nil { - for index := range request.Items { - record(itemOrigin{index: index, rule: request.Items[index].Rule}, "skipped", err) - } - return result, err - } - var stop error - var prepared []preparedFirewallRuleCreate - var origins []itemOrigin - flush := func() { - if len(prepared) == 0 { - return - } - defer func() { prepared, origins = nil, nil }() - if stop == nil { - stop = ctx.Err() - } - if stop != nil { - for _, origin := range origins { - record(origin, "skipped", stop) - } - return - } - runtime := prepared[0].runtime - scope := prepared[0].request.Rule.Scope - snapshot, err := runtime.ObserveMutation(ctx, scope) - if err != nil { - for _, origin := range origins { - record(origin, "failed", err) - } - if !errors.Is(err, filter.ErrFamilyUnavailable) { - stop = err - } - return - } - stored, err := s.rules.List(ctx) - if err != nil { - for _, origin := range origins { - record(origin, "failed", err) - } - stop = err - return - } - identities, err := firewallRuleCollisions(stored, runtime.Provider(), "") - if err != nil { - for _, origin := range origins { - record(origin, "failed", err) - } - stop = err - return - } - observedIdentities, err := filter.ObservedRuleCollisionIndex(snapshot) - if err != nil { - for _, origin := range origins { - record(origin, "failed", err) - } - stop = err - return - } - valid := prepared[:0] - validOrigins := origins[:0] - for index, entry := range prepared { - rule := entry.request.Rule - checkErr := identities.Check(rule) - if checkErr == nil { - checkErr = observedIdentities.Check(rule) - } - if checkErr == nil { - checkErr = identities.Add(rule) - } - if checkErr != nil { - record(origins[index], "failed", checkErr) - continue - } - valid = append(valid, entry) - validOrigins = append(validOrigins, origins[index]) - } - if len(valid) == 0 { - return - } - if len(valid) > 1 && t != nil { - t.Log(i18n.GetMsgWithMap("FirewallCreateBatchStep", map[string]interface{}{ - "backend": runtime.Provider(), "count": len(valid), - })) - } - itemErrors := s.applyCreateRules(ctx, runtime, snapshot, stored, valid) - for offset, origin := range validOrigins { - if err := itemErrors[offset]; err != nil { - record(origin, "failed", err) - if firewallCreateUnavailable(err) { - stop = err - } - } else { - record(origin, "succeeded", nil) - } - } - } - for index, item := range request.Items { - if stop == nil { - stop = ctx.Err() - } - origin := itemOrigin{index: index, rule: item.Rule} - if stop != nil { - record(origin, "skipped", stop) - continue - } - rules := []filter.FirewallRule{item.Rule} - if item.SourceKind == constant.FirewallRuleSourceImported { - rules, err = convertImportedFirewallRule(item.Rule, selected) - if err != nil { - record(origin, "failed", err) - continue - } - if t != nil { - t.Log(i18n.GetMsgWithMap("FirewallImportRuleConversion", map[string]interface{}{ - "index": index + 1, "total": len(request.Items), "source": item.Rule.Scope.Provider, - "target": selected, "rule": describe(item.Rule), "count": len(rules), - })) - } - } else if selected == filter.ProviderUFW && strings.TrimSpace(item.Rule.DestinationPort) != "" { - protocol := strings.ToLower(strings.TrimSpace(item.Rule.Protocol)) - if protocol == "" || protocol == "all" || protocol == "any" { - rules, err = filter.ExpandAtomicRules(applySelectedProviderScopeDefaults(item.Rule, selected)) - if err != nil { - record(origin, "failed", err) - continue - } - } - } - for part, rule := range rules { - origin := itemOrigin{index: index, part: part, count: len(rules), rule: rule} - scope := applySelectedProviderScopeDefaults(rule, selected).Scope.Normalize() - if len(prepared) > 0 && (prepared[0].request.Rule.Scope != scope || rule.OrderIndex != nil) { - flush() - } - if stop == nil { - stop = ctx.Err() - } - if stop != nil { - record(origin, "skipped", stop) - continue - } - child := item - child.Rule = rule - entry, prepareErr := s.prepareCreate(ctx, selected, child) - if prepareErr != nil { - record(origin, "failed", prepareErr) - if firewallCreateUnavailable(prepareErr) { - stop = prepareErr - } - continue - } - origin.rule = entry.request.Rule - prepared = append(prepared, entry) - origins = append(origins, origin) - if !supportsNativeRuleBatch(selected) || rule.OrderIndex != nil { - flush() - } - } - } - flush() - return result, stop -} - -func convertImportedFirewallRule(rule filter.FirewallRule, selected filter.Provider) ([]filter.FirewallRule, error) { - source := rule.Scope.Normalize().Provider - if source == "" { - source = selected - } - rules, err := filter.ExpandAtomicRules(applySelectedProviderScopeDefaults(rule, source)) - if err != nil { - return nil, err - } - var converted []filter.FirewallRule - for _, sourceRule := range rules { - policy, err := model.FirewallRuleFromDomain(sourceRule) - if err != nil { - return nil, err - } - targetRules, err := policy.RulesForProvider(selected) - if err != nil { - return nil, err - } - converted = append(converted, targetRules...) - } - return converted, nil -} - -func firewallCreateUnavailable(err error) bool { - return errors.Is(err, filter.ErrProviderUnavailable) || errors.Is(err, filter.ErrAdapterUnavailable) || - errors.Is(err, filter.ErrInventoryUnavailable) || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) -} - -func (s *FirewallService) prepareCreate(ctx context.Context, selected filter.Provider, request dto.FirewallRuleCreateItem) (preparedFirewallRuleCreate, error) { - rule, err := filter.NormalizeRule(applySelectedProviderScopeDefaults(request.Rule, selected)) - if err != nil { - return preparedFirewallRuleCreate{}, err - } - if rule.Scope.Provider != selected { - return preparedFirewallRuleCreate{}, fmt.Errorf("%w: selected provider is %s", filter.ErrInvalidRule, selected) - } - runtime, err := s.adapters.Resolve(selected) - if err != nil { - return preparedFirewallRuleCreate{}, err - } - rule, err = runtime.Prepare(rule) - if err != nil { - return preparedFirewallRuleCreate{}, err - } - if err := runtime.CheckRule(ctx, rule); err != nil { - return preparedFirewallRuleCreate{}, err - } - rule.UUID = "" - request.Rule = rule - if request.SourceKind == "" { - request.SourceKind = constant.FirewallRuleSourceUser - } - return preparedFirewallRuleCreate{request: request, runtime: runtime}, nil -} - -func (s *FirewallService) applyCreateRules(ctx context.Context, runtime *filterruntime.Engine, snapshot filter.Snapshot, stored []model.FirewallRule, prepared []preparedFirewallRuleCreate) []error { - results := make([]error, len(prepared)) - failAll := func(err error) []error { - for index := range results { - results[index] = err - } - return results - } - var maximumSequence int64 - for _, record := range stored { - if record.Sequence != nil && *record.Sequence > maximumSequence { - maximumSequence = *record.Sequence - } - } - records := make([]model.FirewallRule, 0, len(prepared)) - changes := make([]filter.DesiredChange, 0, len(prepared)) - for _, entry := range prepared { - rule := entry.request.Rule - appendRule := false - if rule.Scope.Provider == filter.ProviderUFW && rule.OrderIndex == nil { - position, err := runtime.AppendPosition(ctx, snapshot, rule) - if err != nil { - return failAll(err) - } - rule.OrderIndex, appendRule = &position, true - } else if rule.OrderIndex != nil { - maximum, err := runtime.MaxPosition(ctx, snapshot, rule) - if err != nil { - return failAll(err) - } - if *rule.OrderIndex < 1 || *rule.OrderIndex > maximum+1 { - return failAll(fmt.Errorf("%w: create target position %d is out of range 1-%d", filter.ErrInvalidRule, *rule.OrderIndex, maximum+1)) - } - appendRule = rule.Scope.Provider == filter.ProviderUFW && *rule.OrderIndex == maximum+1 - } - record, err := firewallRuleModelForCreate(rule, entry.request, constant.FirewallRuleOriginCreated) - if err != nil { - return failAll(err) - } - if rule.Scope.Provider != filter.ProviderFirewalld { - maximumSequence += model.FirewallRuleSequenceStep - sequence := maximumSequence - if rule.OrderIndex != nil { - sequence, err = s.sequenceForCreatedFirewallRule(ctx, snapshot, rule) - if err != nil { - return failAll(err) - } - } - record.Sequence = &sequence - } - record.UUID = uuid.NewString() - rule.UUID = record.UUID - records = append(records, record) - changes = append(changes, filter.DesiredChange{Operation: filter.ChangeCreate, After: &rule, Append: appendRule}) - } - if err := runtime.ExecuteCreate(ctx, snapshot, changes); err != nil { - return failAll(firewallCreateExecutionError(err)) - } - for index := range records { - results[index] = s.saveFirewallRule(ctx, &records[index]) - } - return results -} - -func firewallCreateExecutionError(err error) error { - return fmt.Errorf("%s: %w", i18n.GetMsgByKey("FirewallCreateRuleExecutionFailed"), err) -} - -func (s *FirewallService) saveFirewallRule(ctx context.Context, record *model.FirewallRule) error { - if err := s.rules.Create(ctx, record); err != nil { - message := "FirewallCreateRulePersistenceFailed" - if record.Origin == constant.FirewallRuleOriginAdopted { - message = "FirewallAdoptRulePersistenceFailed" - } - return fmt.Errorf("%s: %w", i18n.GetMsgByKey(message), err) - } - return nil -} - -type preparedFirewallRuleDelete struct { - index int - stored model.FirewallRule - desired filter.DesiredRule - runtime *filterruntime.Engine - compiled int -} - -func (s *FirewallService) Delete(ctx context.Context, request dto.FirewallRuleDelete) (dto.FirewallRuleDeleteResponse, error) { - if err := ctx.Err(); err != nil { - return dto.FirewallRuleDeleteResponse{}, err - } - if len(request.UUIDs) == 0 && len(request.BeforeRules) == 0 { - return dto.FirewallRuleDeleteResponse{}, fmt.Errorf("%w: rules are required", filter.ErrInvalidRule) - } - request.UUIDs = append([]string(nil), request.UUIDs...) - request.BeforeRules = append([]dto.FirewallRuleDeleteTarget(nil), request.BeforeRules...) - taskItem, err := task.NewTask(firewallTaskName(task.TaskDelete, firewallTaskHost, ""), task.TaskDelete, task.TaskScopeFirewall, "", 0) - if err != nil { - return dto.FirewallRuleDeleteResponse{}, err - } - taskItem.AddSubTaskWithOps(taskItem.Name, func(t *task.Task) error { - t.Logf("rules=%d", len(request.UUIDs)+len(request.BeforeRules)) - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - if err := t.TaskCtx.Err(); err != nil { - return err - } - _, err := s.deleteRules(t.TaskCtx, request, t) - return err - }, nil, 0, 0) - if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { - taskItem.LogFailedWithErr(taskItem.Name, err) - closeUnstartedFirewallTask(taskItem) - return dto.FirewallRuleDeleteResponse{}, fmt.Errorf("save firewall deletion task: %w", err) - } - go func() { - if err := taskItem.Execute(); err != nil && global.LOG != nil { - global.LOG.Errorf("firewall deletion task %s failed: %v", taskItem.TaskID, err) - } - }() - return dto.FirewallRuleDeleteResponse{TaskID: taskItem.TaskID, Queued: true}, nil -} - -func (s *FirewallService) deleteRules(ctx context.Context, request dto.FirewallRuleDelete, t *task.Task) (result dto.FirewallRuleDeleteResponse, taskErr error) { - var firstFailure error - defer func() { - sort.SliceStable(result.Errors, func(i, j int) bool { return result.Errors[i].Index < result.Errors[j].Index }) - if t != nil { - t.Log(i18n.GetMsgWithMap("FirewallRuleOperationResult", map[string]interface{}{ - "succeeded": result.Succeeded, "failed": result.Failed, - })) - } - if taskErr == nil { - taskErr = firstFailure - } - }() - record := func(index int, ruleUUID string, err error) { - if err != nil { - result.Failed++ - result.Errors = append(result.Errors, dto.FirewallRuleDeleteFailure{Index: index, UUID: ruleUUID, Error: err.Error()}) - if firstFailure == nil { - firstFailure = err - } - } else { - result.Succeeded++ - } - if t != nil { - label := fmt.Sprintf("[%d/%d] %s", result.Succeeded+result.Failed, len(request.UUIDs)+len(request.BeforeRules), ruleUUID) - t.LogWithStatus(label, err) - } - } - selectedProvider, err := s.selectedProvider(ctx) - if err != nil { - return result, err - } - type beforeGroup struct { - targets []dto.FirewallRuleDeleteTarget - indexes []int - } - beforeGroups := make(map[string]*beforeGroup) - for index, target := range request.BeforeRules { - key := target.Scope.Normalize().Key() - if beforeGroups[key] == nil { - beforeGroups[key] = &beforeGroup{} - } - group := beforeGroups[key] - group.targets = append(group.targets, target) - group.indexes = append(group.indexes, len(request.UUIDs)+index) - } - for _, group := range beforeGroups { - err := s.deleteBeforeRules(ctx, selectedProvider, group.targets) - for index, target := range group.targets { - record(group.indexes[index], target.InstanceKey, err) - } - } - type deleteGroup struct{ items []preparedFirewallRuleDelete } - groups := make([]deleteGroup, 0) - groupIndexes := make(map[string]int) - seen := make(map[string]struct{}, len(request.UUIDs)) - for index, value := range request.UUIDs { - ruleUUID := strings.TrimSpace(value) - if err := ctx.Err(); err != nil { - record(index, ruleUUID, err) - continue - } - if _, exists := seen[ruleUUID]; exists { - record(index, ruleUUID, fmt.Errorf("duplicate firewall rule UUID")) - continue - } - seen[ruleUUID] = struct{}{} - prepared, err := s.prepareDelete(ctx, index, ruleUUID, selectedProvider) - if err != nil { - record(index, ruleUUID, err) - continue - } - groupKey := string(prepared.desired.Rule.Scope.Provider) + ":" + prepared.desired.Rule.Scope.Key() - if prepared.compiled != 1 { - groupKey += ":" + prepared.stored.UUID - } - groupIndex, exists := groupIndexes[groupKey] - if !exists { - groupIndex = len(groups) - groupIndexes[groupKey] = groupIndex - groups = append(groups, deleteGroup{}) - } - groups[groupIndex].items = append(groups[groupIndex].items, prepared) - } - for _, group := range groups { - if len(group.items) > 1 && group.items[0].compiled == 1 && supportsNativeRuleBatch(group.items[0].desired.Rule.Scope.Provider) { - err := ctx.Err() - if err == nil { - if t != nil { - t.Logf("%s: rules=%d", i18n.GetMsgByKey(task.TaskDelete), len(group.items)) - } - err = s.deleteNativeRuleBatch(ctx, group.items) - } - for _, item := range group.items { - record(item.index, item.stored.UUID, err) - } - continue - } - for _, item := range group.items { - err := ctx.Err() - if err == nil { - if t != nil { - t.Logf("%s %s", i18n.GetMsgByKey(task.TaskDelete), item.stored.UUID) - } - err = s.deleteRule(ctx, item.stored.UUID) - } - record(item.index, item.stored.UUID, err) - } + reader, ok := runtime.(filter.NativeDetailReader) + if !ok { + return "", fmt.Errorf("%w: native details for %s", filter.ErrAdapterUnavailable, runtime.Provider()) } - return result, nil + return reader.NativeDetail(ctx, request.Name, request.Permanent) } -func (s *FirewallService) deleteBeforeRules(ctx context.Context, provider filter.Provider, targets []dto.FirewallRuleDeleteTarget) error { - scope := targets[0].Scope.Normalize() - if scope.Provider != provider || !isDirectFirewallProvider(provider) || scope.Chain != filter.BasicBeforeChain { - return fmt.Errorf("%w: native deletion only supports the selected firewall before chain", filter.ErrUnsupportedScope) - } - runtime, err := s.adapters.Resolve(provider) - if err != nil { +func (s *FirewallService) Adopt(ctx context.Context, request dto.FirewallRuleAdopt) error { + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + if err := s.checkSelectedProvider(ctx, request.Scope.Provider); err != nil { return err } - snapshot, err := runtime.ObserveMutation(ctx, scope) + runtime, err := s.firewallAdapter(request.Scope.Provider) if err != nil { return err } - changes := make([]filter.DesiredChange, 0, len(targets)) - seen := make(map[string]bool, len(targets)) - for _, target := range targets { - if target.Scope.Normalize().Key() != scope.Key() || seen[target.InstanceKey] { - return fmt.Errorf("%w: duplicate or mismatched before rule", filter.ErrInvalidRule) + if runtime.Provider() != filter.ProviderNftables { + if request.Rule == nil { + return fmt.Errorf("%w: original rule is required for adoption", filter.ErrInvalidRule) } - seen[target.InstanceKey] = true - observed, err := filter.FindCandidate(snapshot.Rules, target.InstanceKey) + rule, err := filter.NormalizeRule(*request.Rule) if err != nil { - return filter.ErrRuleStale - } - if err := filter.GuardMutation(observed); err != nil { return err } - if observed.ParseStatus != filter.ParseStatusSupported || observed.Locator.Position == nil { - return fmt.Errorf("%w: before rule cannot be deleted", filter.ErrUnsupportedScope) - } - rule := firewallsync.ObservedRule(observed) - if rule.UUID == "" { - rule.UUID = uuid.NewString() - } - locator := observed.Locator - changes = append(changes, filter.DesiredChange{ - Operation: filter.ChangeDelete, Before: &rule, Locator: &locator, UnmarkedAdopted: observed.Marker == "", - }) + if rule.Scope.Key() != request.Scope.Key() { + return fmt.Errorf("%w: adoption rule scope mismatch", filter.ErrInvalidScope) + } + observed := filter.ObservedRule{Rule: rule, Marker: request.Marker, ParseStatus: filter.ParseStatusSupported} + return s.adoptRule(ctx, runtime, filter.RuleSet{Scope: rule.Scope}, observed, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}) } - sort.Slice(changes, func(i, j int) bool { return *changes[i].Locator.Position > *changes[j].Locator.Position }) - plan, verification, err := runtime.Execute(ctx, snapshot, changes) + if request.InstanceKey == "" { + return fmt.Errorf("%w: instance key is required for nftables adoption", filter.ErrInvalidRule) + } + snapshot, err := readMutableFirewallRules(runtime, ctx, request.Scope) if err != nil { return err } - if !verification.Matched { - return filter.ErrVerificationFailed - } - if len(verification.Snapshot.Rules) != len(snapshot.Rules)-len(changes) { - return rollbackFirewallPlan(ctx, runtime, plan, filter.ErrVerificationFailed) + observed, err := filter.FindCandidate(snapshot.Rules, request.InstanceKey) + if err != nil { + return filter.ErrRuleStale } - return nil + return s.adoptRule(ctx, runtime, snapshot, observed, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}) } -func (s *FirewallService) prepareDelete( - ctx context.Context, - index int, - ruleUUID string, - selectedProvider filter.Provider, -) (preparedFirewallRuleDelete, error) { - if ruleUUID == "" { - return preparedFirewallRuleDelete{}, fmt.Errorf("%w: rule UUID is required", repo.ErrFirewallPersistenceInvalid) +func (s *FirewallService) Create(ctx context.Context, request dto.FirewallRuleCreate) (dto.FirewallRuleCreateResponse, error) { + if len(request.Items) > filter.MaxAtomicExpansion { + return dto.FirewallRuleCreateResponse{}, fmt.Errorf("create or import at most %d rules per batch (after expansion)", filter.MaxAtomicExpansion) } - stored, err := s.rules.GetByUUID(ctx, ruleUUID) + provider, err := s.selectedProvider(ctx) if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return preparedFirewallRuleDelete{}, fmt.Errorf("%w: managed rule %q was not found", filter.ErrInvalidRule, ruleUUID) - } - return preparedFirewallRuleDelete{}, err - } - if err := checkFirewallRuleWhitelistProtection(selectedProvider, stored); err != nil { - return preparedFirewallRuleDelete{}, err + return dto.FirewallRuleCreateResponse{}, err } - if stored.Origin != constant.FirewallRuleOriginCreated && stored.Origin != constant.FirewallRuleOriginAdopted { - return preparedFirewallRuleDelete{}, fmt.Errorf("%w: only created or adopted rules can be deleted", filter.ErrInvalidRule) + if err := validateFirewallCreateBatch(request, provider); err != nil { + return dto.FirewallRuleCreateResponse{}, err } - desiredRules, err := s.compileStoredFirewallRules(ctx, stored, selectedProvider) + taskItem, err := task.NewTask(firewallTaskName(task.TaskCreate, firewallTaskHost, ""), task.TaskCreate, task.TaskScopeFirewall, "", 0) if err != nil { - return preparedFirewallRuleDelete{}, err - } - if len(desiredRules) == 0 { - return preparedFirewallRuleDelete{}, fmt.Errorf("%w: policy %q has no compiled target rules", filter.ErrInvalidRule, ruleUUID) + return dto.FirewallRuleCreateResponse{}, err } - desired := desiredRules[0] - runtime, err := s.adapters.Resolve(desired.Rule.Scope.Provider) - if err != nil { - return preparedFirewallRuleDelete{}, err + taskItem.AddSubTaskWithOps(i18n.GetMsgByKey("FirewallCreateRulesStep"), func(t *task.Task) error { + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + _, err := s.createRules(t.TaskCtx, request, t) + return err + }, nil, 0, 0) + if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { + taskItem.LogFailedWithErr(taskItem.Name, err) + closeUnstartedFirewallTask(taskItem) + return dto.FirewallRuleCreateResponse{}, fmt.Errorf("save firewall creation task: %w", err) } - return preparedFirewallRuleDelete{ - index: index, stored: stored, desired: desired, runtime: runtime, compiled: len(desiredRules), - }, nil + go func() { + if err := taskItem.Execute(); err != nil && global.LOG != nil { + global.LOG.Errorf("firewall creation task %s failed: %v", taskItem.TaskID, err) + } + }() + return dto.FirewallRuleCreateResponse{TaskID: taskItem.TaskID, Queued: true}, nil } -func (s *FirewallService) deleteNativeRuleBatch(ctx context.Context, prepared []preparedFirewallRuleDelete) error { - runtime := prepared[0].runtime - snapshot, err := runtime.ObserveMutation(ctx, prepared[0].desired.Rule.Scope) - if err != nil { - return err +func (s *FirewallService) Delete(ctx context.Context, request dto.FirewallRuleDelete) (dto.FirewallRuleDeleteResponse, error) { + if err := ctx.Err(); err != nil { + return dto.FirewallRuleDeleteResponse{}, err } - type positionedDelete struct { - position int - change filter.DesiredChange - } - positioned := make([]positionedDelete, 0, len(prepared)) - for _, item := range prepared { - observed, observeErr := filter.ManagedObserved(snapshot, item.desired) - if observeErr != nil { - if errors.Is(observeErr, filter.ErrRuleStale) { - missing, mergeErr := managedFirewallRuleMissing(snapshot, item.desired) - if mergeErr != nil { - return mergeErr - } - if missing { - continue - } - } - return observeErr - } - if observed.Locator.Position == nil { - return fmt.Errorf("%w: managed native firewall rule has no position", filter.ErrRuleStale) - } - before := item.desired.Rule - locator := observed.Locator - positioned = append(positioned, positionedDelete{ - position: *observed.Locator.Position, - change: filter.DesiredChange{ - Operation: filter.ChangeDelete, Before: &before, Locator: &locator, - }, - }) + if len(request.UUIDs) == 0 && len(request.BeforeRules) == 0 { + return dto.FirewallRuleDeleteResponse{}, fmt.Errorf("%w: rules are required", filter.ErrInvalidRule) } - sort.Slice(positioned, func(i, j int) bool { return positioned[i].position > positioned[j].position }) - changes := make([]filter.DesiredChange, 0, len(positioned)) - for _, item := range positioned { - changes = append(changes, item.change) + request.UUIDs = append([]string(nil), request.UUIDs...) + request.BeforeRules = append([]dto.FirewallRuleDeleteTarget(nil), request.BeforeRules...) + taskItem, err := task.NewTask(firewallTaskName(task.TaskDelete, firewallTaskHost, ""), task.TaskDelete, task.TaskScopeFirewall, "", 0) + if err != nil { + return dto.FirewallRuleDeleteResponse{}, err } - var backendPlan filter.BackendPlan - if len(changes) > 0 { - var verification filter.VerifyResult - backendPlan, verification, err = runtime.Execute(ctx, snapshot, changes) - if err != nil { + taskItem.AddSubTaskWithOps(taskItem.Name, func(t *task.Task) error { + t.Logf("rules=%d", len(request.UUIDs)+len(request.BeforeRules)) + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + if err := t.TaskCtx.Err(); err != nil { return err } - if !verification.Matched { - return filter.ErrVerificationFailed - } + _, err := s.deleteRules(t.TaskCtx, request, t) + return err + }, nil, 0, 0) + if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { + taskItem.LogFailedWithErr(taskItem.Name, err) + closeUnstartedFirewallTask(taskItem) + return dto.FirewallRuleDeleteResponse{}, fmt.Errorf("save firewall deletion task: %w", err) } - - deleted := make([]model.FirewallRule, 0, len(prepared)) - for _, item := range prepared { - if err = s.rules.DeleteWithRevision(ctx, item.stored.UUID, item.stored.Revision); err != nil { - if len(changes) > 0 { - err = rollbackFirewallPlan(ctx, runtime, backendPlan, err) - } - return s.restoreDeletedFirewallRecords(ctx, deleted, err) + go func() { + if err := taskItem.Execute(); err != nil && global.LOG != nil { + global.LOG.Errorf("firewall deletion task %s failed: %v", taskItem.TaskID, err) } - deleted = append(deleted, item.stored) - } - return nil -} - -func supportsNativeRuleBatch(provider filter.Provider) bool { - return provider == filter.ProviderIptables || provider == filter.ProviderNftables -} - -func (s *FirewallService) restoreDeletedFirewallRecords( - ctx context.Context, - deleted []model.FirewallRule, - cause error, -) error { - restoreErrors := make([]error, 0) - for index := range deleted { - record := deleted[index] - if err := s.rules.Create(ctx, &record); err != nil { - restoreErrors = append(restoreErrors, fmt.Errorf("restore deleted firewall rule %q: %w", record.UUID, err)) - } - } - if len(restoreErrors) == 0 { - return cause - } - return errors.Join(append([]error{cause}, restoreErrors...)...) + }() + return dto.FirewallRuleDeleteResponse{TaskID: taskItem.TaskID, Queued: true}, nil } func (s *FirewallService) Update(ctx context.Context, clientIP string, request dto.FirewallRuleUpdate) error { @@ -1741,1114 +606,986 @@ func (s *FirewallService) Update(ctx context.Context, clientIP string, request d return s.updateRuleDescription(ctx, request.UUID, *request.Description) } -func (s *FirewallService) updateRuleDescription(ctx context.Context, ruleUUID, description string) error { - stored, err := s.rules.GetByUUID(ctx, ruleUUID) - if err != nil { - return err - } - selected, err := s.selectedProviderForStoredRule(ctx, stored) - if err != nil { - return err - } - if err := checkFirewallRuleWhitelistProtection(selected, stored); err != nil { - return err - } - if stored.Origin != constant.FirewallRuleOriginCreated && stored.Origin != constant.FirewallRuleOriginAdopted { - return fmt.Errorf("%w: only created or adopted rules can be changed", filter.ErrInvalidRule) - } - description = strings.TrimSpace(description) - if stored.Description == description { - return nil - } - return s.rules.UpdateWithRevision(ctx, stored.UUID, stored.Revision, map[string]interface{}{"description": description}) -} - func (s *FirewallService) Reorder(ctx context.Context, clientIP string, request dto.FirewallRuleReorder) error { firewallRuleMutationMu.Lock() defer firewallRuleMutationMu.Unlock() return s.updateRuleOrder(ctx, request.UUID, request.TargetPosition, request.Priority, nil) } -func (s *FirewallService) checkSelectedProvider(ctx context.Context, requested filter.Provider) error { - selected, err := s.selectedProvider(ctx) - if err != nil { - return err - } - if selected != requested { - return fmt.Errorf("%w: selected provider is %s, requested %s", filter.ErrProviderUnavailable, selected, requested) - } - return nil +func NewIFirewallService() IFirewallService { + return newFirewallService() } -func (s *FirewallService) Adopt(ctx context.Context, request dto.FirewallRuleAdopt) error { - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - if err := s.checkSelectedProvider(ctx, request.Scope.Provider); err != nil { - return err +func currentFirewallLifecycleTaskID() string { + firewallLifecycleTaskMu.Lock() + defer firewallLifecycleTaskMu.Unlock() + return firewallLifecycleTaskID +} + +func (s *FirewallService) OperateFirewall(request dto.FirewallLifecycleOperation) error { + switch request.Operation { + case "disableBanPing": + if err := firewall.UpdatePingStatus("0"); err != nil { + return err + } + return settingRepo.Update(constant.FirewallPingStatusKey, constant.StatusDisable) + case "enableBanPing": + if err := firewall.UpdatePingStatus("1"); err != nil { + return err + } + return settingRepo.Update(constant.FirewallPingStatusKey, constant.StatusEnable) } - runtime, err := s.adapters.Resolve(request.Scope.Provider) - if err != nil { - return err + baseClient := s.baseClient + if baseClient == nil { + baseClient = NewSelectedSystemFirewallClient } - snapshot, err := runtime.ObserveMutation(ctx, request.Scope) + client, err := baseClient() if err != nil { return err } - observed, err := filter.FindCandidate(snapshot.Rules, request.InstanceKey) - if err != nil { - return filter.ErrRuleStale + operation := request.Operation + operationErr := operateFirewallLifecycle(firewallLifecycleClient{client}, operation, request.WithDockerRestart, s.restoreFirewallAfterStart, nil) + restoreFirewalld := client.Name() == lifecycle.ProviderFirewalld && + (operation == string(lifecycle.OperationStart) || operation == string(lifecycle.OperationRestart)) + if operation != string(lifecycle.OperationStart) && operation != string(lifecycle.OperationRestart) { + return operationErr } - return s.adoptRule(ctx, runtime, snapshot, observed, dto.FirewallRuleCreateItem{SourceKind: constant.FirewallRuleSourceUser}) + if operationErr != nil { + var completedErr *firewallCompletedOperationError + var dockerRestartErr *firewallDockerRestartError + if !errors.As(operationErr, &completedErr) && !errors.As(operationErr, &dockerRestartErr) { + return operationErr + } + if global.LOG != nil { + global.LOG.Warnf("firewall %s completed with post-start recovery errors: %v", operation, operationErr) + } + } + if restoreFirewalld { + restoreErr := s.restoreFirewalldRuntimeDependents(context.Background(), operation) + if restoreErr != nil && global.LOG != nil { + global.LOG.Errorf("restore firewalld runtime dependents after %s failed: %v", operation, restoreErr) + } + return nil + } + ReconcileDockerPortGuardBestEffort(context.Background()) + return nil } -func (s *FirewallService) adoptRule(ctx context.Context, runtime *filterruntime.Engine, snapshot filter.Snapshot, observed filter.ObservedRule, source dto.FirewallRuleCreateItem) error { - if isIptablesSystemPresetInventoryScope(observed.Rule.Scope) { - return fmt.Errorf("%w: system preset chains cannot be adopted", filter.ErrUnsupportedScope) +func (s *FirewallService) restoreFirewallAfterStart(client lifecycle.Client) error { + ctx := context.Background() + provider := filter.Provider(client.Name()) + var recoveryErrors []error + recordFailure := func(stage string, err error) { + if err == nil { + return + } + wrapped := fmt.Errorf("%s for %s: %w", stage, provider, err) + recoveryErrors = append(recoveryErrors, wrapped) + if global.LOG != nil { + global.LOG.Errorf("firewall post-start recovery failed: %v", wrapped) + } } - if observed.Protected { - return filter.ErrProtectedRule + if provider == filter.ProviderIptables || provider == filter.ProviderNftables { + isInit, _, err := loadFirewallInitStatus(string(provider), "base") + if err != nil { + recordFailure("load managed chain status", err) + return errors.Join(recoveryErrors...) + } + if !isInit { + return nil + } } - if observed.ParseStatus != filter.ParseStatusSupported || - (observed.Persistence != "" && observed.Persistence != filter.PersistenceStatusConverged) { - return fmt.Errorf("%w: rule cannot be managed", filter.ErrRuleOperation) + if err := s.restoreStoredFirewallRules(ctx, provider, nil); err != nil { + recordFailure("restore stored firewall rules", err) } - rule, err := runtime.Prepare(observed.Rule) - if err != nil { - return err + recordFailure("restore whitelist allowances", s.SyncPortWhitelist(ctx)) + return errors.Join(recoveryErrors...) +} + +func (s *FirewallService) restoreFirewalldRuntimeDependents(ctx context.Context, operation string) error { + restoreForwarding := s.restoreForwarding + if restoreForwarding == nil { + restoreForwarding = func(ctx context.Context) error { return newForwardingService().Restore(ctx) } } - if err := runtime.CheckRule(ctx, rule); err != nil { - return err + restoreDockerGuard := s.restoreDockerGuard + if restoreDockerGuard == nil { + restoreDockerGuard = ReconcileDockerPortGuard } - if rule.Scope.Provider != filter.ProviderFirewalld && observed.Locator.Position != nil { - position := int64(*observed.Locator.Position) - rule.OrderIndex = &position + dockerActive := s.dockerActive + if dockerActive == nil { + dockerActive = firewallDockerActive } - record, err := firewallRuleModelForCreate(rule, source, constant.FirewallRuleOriginAdopted) + + active, err := dockerActive() + restoreErr := restoreFirewalldDependents( + ctx, fmt.Sprintf("after firewalld %s", operation), err == nil && active, restoreForwarding, restoreDockerGuard, + ) if err != nil { - return err + return errors.Join(fmt.Errorf("check Docker status after firewalld %s: %w", operation, err), restoreErr) } - stored, err := s.rules.List(ctx) - if err != nil { - return err + return restoreErr +} + +func (s *FirewallService) runFirewallLifecycleTask(t *task.Task, client lifecycle.Client, request dto.FirewallLifecycleOperation) error { + ctx := t.TaskCtx + provider := filter.Provider(client.Name()) + operationErr := operateFirewallLifecycle(firewallLifecycleClient{client}, request.Operation, request.WithDockerRestart, func(lifecycle.Client) error { + rulesErr := runFirewallLifecycleAction(t, i18n.GetWithName("FirewallRestoreRulesStep", client.Name()), func() error { + return s.restoreStoredFirewallRules(ctx, provider, t) + }) + whitelistErr := runFirewallLifecycleAction(t, i18n.GetMsgByKey("FirewallSyncWhitelistStep"), func() error { + return s.SyncPortWhitelist(ctx) + }) + return errors.Join(rulesErr, whitelistErr) + }, t) + if request.Operation == string(lifecycle.OperationStop) { + return operationErr } - if err := filter.CheckAdoptDuplicates(snapshot, rule); err != nil { - return err + var recoveryErr *firewallCompletedOperationError + if operationErr != nil && !errors.As(operationErr, &recoveryErr) { + return operationErr + } + var forwardingErr error + if provider == filter.ProviderFirewalld { + forwardingErr = runFirewallLifecycleAction(t, i18n.GetMsgByKey("FirewallRestoreForwardingRulesStep"), func() error { + if s.restoreForwarding != nil { + return s.restoreForwarding(ctx) + } + return newForwardingService().Restore(ctx) + }) } - identities, err := firewallRuleCollisions(stored, rule.Scope.Provider, "") + dockerErr := runFirewallLifecycleAction(t, i18n.GetMsgByKey("FirewallInspectDockerGuardStep"), func() error { + if provider == filter.ProviderFirewalld { + active := s.dockerActive + if active == nil { + active = firewallDockerActive + } + running, err := active() + if err != nil || !running { + return err + } + } + if s.restoreDockerGuard != nil { + return s.restoreDockerGuard(ctx) + } + return ReconcileDockerPortGuard(ctx) + }) + return errors.Join(operationErr, forwardingErr, dockerErr) +} + +func (s *FirewallService) saveFirewallResetOrder(ctx context.Context, provider filter.Provider, stored []model.FirewallRule) error { + runtime, err := s.firewallAdapter(provider) if err != nil { return err } - for _, existing := range stored { - marker := "1panel-rule:" + existing.UUID - if existing.UUID != "" && (observed.Marker == marker || strings.HasPrefix(observed.Marker, marker+"-")) { - return fmt.Errorf("%w: rule is already managed", filter.ErrRuleOperation) + desiredByScope := make(map[string][]filter.DesiredRule) + byUUID := make(map[string]model.FirewallRule, len(stored)) + for _, record := range stored { + if record.Origin != constant.FirewallRuleOriginCreated && record.Origin != constant.FirewallRuleOriginAdopted { + continue } - } - if err := identities.CheckDuplicate(rule); err != nil { - if errors.Is(err, filter.ErrRuleOperation) { - return filter.ErrDuplicateAdoption + desired, err := compileStoredFirewallRules(ctx, record, runtime) + if isFirewallPolicyIncompatible(err) { + continue } - return err - } - if rule.Scope.Provider != filter.ProviderFirewalld { - sequence, err := s.sequenceForCreatedFirewallRule(ctx, snapshot, rule) if err != nil { return err } - record.Sequence = &sequence - } - record.UUID = uuid.NewString() - rule.UUID = record.UUID - plan, verification, err := runtime.Execute(ctx, snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeAdopt, After: &rule, Locator: &observed.Locator, PreviousMarker: observed.Marker, - }}) - if err != nil { - return err - } - if !verification.Matched { - return filter.ErrVerificationFailed - } - if _, err := filter.FindCommittedObserved(verification.Snapshot, rule, plan); err != nil { - return rollbackFirewallPlan(ctx, runtime, plan, err) - } - return s.saveFirewallRule(ctx, &record) -} - -func (s *FirewallService) deleteRule(ctx context.Context, ruleUUID string) error { - if ruleUUID == "" { - return fmt.Errorf("%w: rule UUID is required", repo.ErrFirewallPersistenceInvalid) - } - stored, err := s.rules.GetByUUID(ctx, ruleUUID) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return fmt.Errorf("%w: managed rule %q was not found", filter.ErrInvalidRule, ruleUUID) - } - return err - } - if stored.Origin != constant.FirewallRuleOriginCreated && stored.Origin != constant.FirewallRuleOriginAdopted { - return fmt.Errorf("%w: only created or adopted rules can be deleted", filter.ErrInvalidRule) - } - selected, err := s.selectedProviderForStoredRule(ctx, stored) - if err != nil { - return err - } - desiredRules, err := s.compileStoredFirewallRules(ctx, stored, selected) - if err != nil { - return err - } - if err := checkFirewallRuleWhitelistProtection(selected, stored); err != nil { - return err - } - type appliedDelete struct { - runtime *filterruntime.Engine - plan filter.BackendPlan - } - applied := make([]appliedDelete, 0, len(desiredRules)) - rollback := func(cause error) error { - for index := len(applied) - 1; index >= 0; index-- { - cause = rollbackFirewallPlan(ctx, applied[index].runtime, applied[index].plan, cause) + byUUID[record.UUID] = record + for _, rule := range desired { + key := rule.Rule.Scope.Key() + desiredByScope[key] = append(desiredByScope[key], rule) } - return cause } - for _, desired := range desiredRules { - runtime, runtimeErr := s.resolveRuntime(ctx, desired.Rule.Scope.Provider) - if runtimeErr != nil { - return rollback(runtimeErr) + captured := make(map[string]model.FirewallRule) + for _, group := range firewallScopeReadGroups(filter.ManagedInputScopes(provider)) { + snapshots, err := listFirewallRuleScopes(runtime, ctx, group) + if errors.Is(err, filter.ErrFamilyUnavailable) { + continue } - snapshot, observeErr := runtime.ObserveMutation(ctx, desired.Rule.Scope) - if observeErr != nil { - return rollback(observeErr) + if err != nil { + return err } - observed, managedErr := filter.ManagedObserved(snapshot, desired) - if managedErr != nil { - if errors.Is(managedErr, filter.ErrRuleStale) { - missing, mergeErr := managedFirewallRuleMissing(snapshot, desired) - if mergeErr != nil { - return rollback(mergeErr) - } - if missing { + for _, snapshot := range snapshots { + items, err := mergeFirewallInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desiredByScope[snapshot.Scope.Key()]}) + if err != nil { + return err + } + for _, item := range items { + if item.Desired == nil || item.Observed == nil { continue } + record := byUUID[item.Desired.UUID] + if provider == filter.ProviderFirewalld { + record.Priority = item.Observed.Rule.Priority + record.Sequence = nil + } else { + if item.Observed.Locator.Position == nil { + return filter.ErrVerificationFailed + } + position := int64(*item.Observed.Locator.Position) + if previous, exists := captured[record.UUID]; exists && *previous.Sequence <= position { + continue + } + record.Sequence, record.Priority = &position, nil + } + captured[record.UUID] = record } - return rollback(managedErr) - } - restoreAtEnd := false - if desired.Rule.Scope.Provider == filter.ProviderUFW && observed.Locator.Position != nil { - maxPosition := maxObservedFirewallPosition(snapshot) - restoreAtEnd = int64(*observed.Locator.Position) == maxPosition } - locator := observed.Locator - before := desired.Rule - backendPlan, verification, executeErr := runtime.Execute(ctx, snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeDelete, Before: &before, Locator: &locator, RestoreAtEnd: restoreAtEnd, - }}) - if executeErr != nil { - return rollback(executeErr) - } - if !verification.Matched { - return rollback(filter.ErrVerificationFailed) - } - applied = append(applied, appliedDelete{runtime: runtime, plan: backendPlan}) } - if err := s.rules.DeleteWithRevision(ctx, stored.UUID, stored.Revision); err != nil { - return rollback(err) + orders := make([]model.FirewallRule, 0, len(captured)) + for _, record := range stored { + if snapshot, exists := captured[record.UUID]; exists { + orders = append(orders, snapshot) + } } - return nil + return s.rules.SaveResetOrder(ctx, orders) } -func managedFirewallRuleMissing(snapshot filter.Snapshot, desired filter.DesiredRule) (bool, error) { - items, err := filter.MergeInventory(filter.InventoryMergeInput{ - Observed: snapshot.Rules, - Desired: []filter.DesiredRule{desired}, - }) +func (s *FirewallService) combinedUFWInventory(ctx context.Context, runtime filter.Adapter, scope filter.Scope, refresh bool) (dto.FirewallRuleInventoryResponse, error) { + scopes := []filter.Scope{scope, scope} + scopes[0].Family = filter.FamilyIPv4 + scopes[1].Family = filter.FamilyIPv6 + snapshots, err := readFirewallRuleScopes(runtime, ctx, scopes) if err != nil { - return false, err + return dto.FirewallRuleInventoryResponse{}, err } - for _, item := range items { - if item.Desired != nil && item.Desired.UUID == desired.UUID { - return item.Match == filter.InventoryMatchMissing && item.Observed == nil, nil - } + if len(snapshots) != len(scopes) { + return dto.FirewallRuleInventoryResponse{}, fmt.Errorf("%w: incomplete UFW multi-family inventory", filter.ErrAdapterUnavailable) } - return false, nil -} -func (s *FirewallService) updateRule(ctx context.Context, clientIP, ruleUUID string, requestedRule filter.FirewallRule) error { - requestedRule, err := filter.NormalizeRule(requestedRule) - if err != nil { - return err - } - stored, err := s.rules.GetByUUID(ctx, ruleUUID) + stored, err := s.rules.List(ctx) if err != nil { - return err + return dto.FirewallRuleInventoryResponse{}, err } - previousRules, compileErr := stored.RulesForProvider(requestedRule.Scope.Provider) - if compileErr == nil && len(previousRules) == 1 { - sameContent, err := filter.SameRuleContent(previousRules[0], requestedRule) + desiredByScope, failures := s.desiredFirewallRulesByScope(ctx, stored, runtime) + response := dto.FirewallRuleInventoryResponse{Items: failures} + seenNotices := make(map[string]struct{}) + for index, snapshot := range snapshots { + if snapshot.Scope.Key() != scopes[index].Key() { + return dto.FirewallRuleInventoryResponse{}, fmt.Errorf("%w: unexpected UFW inventory scope %q", filter.ErrInvalidScope, snapshot.Scope.Key()) + } + desired := desiredByScope[snapshot.Scope.Key()] + items, err := mergeFirewallInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desired}) if err != nil { - return err + return dto.FirewallRuleInventoryResponse{}, err } - if sameContent { - if requestedRule.Scope.Provider == filter.ProviderFirewalld { - beforePriority, afterPriority := 0, 0 - if stored.Priority != nil { - beforePriority = *stored.Priority - } - if requestedRule.Priority != nil { - afterPriority = *requestedRule.Priority - } - if beforePriority != afterPriority { - return s.updateRuleOrder(ctx, ruleUUID, nil, &afterPriority, &requestedRule.Description) - } - } else if requestedRule.OrderIndex != nil { - return s.updateRuleOrder(ctx, ruleUUID, requestedRule.OrderIndex, nil, &requestedRule.Description) + response.Items = append(response.Items, items...) + for _, notice := range snapshot.Notices { + key := string(notice.Code) + "\x00" + strings.Join(notice.Values, "\x00") + if _, exists := seenNotices[key]; exists { + continue } - return s.updateRuleDescription(ctx, ruleUUID, requestedRule.Description) + seenNotices[key] = struct{}{} + response.Notices = append(response.Notices, notice) } } - prepared, err := s.prepareManagedUpdate(ctx, clientIP, ruleUUID, requestedRule) - if err != nil { - return err + return response, nil +} + +func finalizeFirewallInventory(response dto.FirewallRuleInventoryResponse, request dto.FirewallRuleInventory) dto.FirewallRuleInventoryResponse { + provider := request.Scope.Provider + if len(request.Scopes) > 0 { + provider = request.Scopes[0].Provider } - metadataOnly, err := isFirewallMetadataOnlyUpdate(prepared.Before.Rule, prepared.After, prepared.Observed.Locator) - if err != nil { - return err + response.IPv4Range, response.IPv6Range = firewallInventoryPositionRanges(provider, response.Items) + response.AllTotal = int64(len(response.Items)) + for _, item := range response.Items { + if isDeletableManagedInventoryItem(item) { + response.ManagedTotal++ + } } - if metadataOnly { - if prepared.After.Description == prepared.Before.Rule.Description { - return nil + filtered := make([]filter.InventoryItem, 0, len(response.Items)) + for _, item := range response.Items { + if matchesFirewallInventoryRequest(item, request) { + filtered = append(filtered, item) } - return s.rules.UpdateWithRevision(ctx, prepared.Stored.UUID, prepared.Stored.Revision, map[string]interface{}{ - "description": prepared.After.Description, - }) } - return s.executeManagedMutation(ctx, managedMutationRequest{ - Stored: prepared.Stored, Before: prepared.Before.Rule, After: prepared.After, - Snapshot: prepared.Snapshot, Locator: prepared.Observed.Locator, - AdapterOperation: filter.ChangeUpdate, Runtime: prepared.Runtime, - }) + response.Total = int64(len(filtered)) + if request.All { + response.Items = filtered + return response + } + page, pageSize := max(1, request.Page), max(1, request.PageSize) + start := (page - 1) * pageSize + if start >= len(filtered) { + response.Items = make([]filter.InventoryItem, 0) + return response + } + end := min(start+pageSize, len(filtered)) + response.Items = append([]filter.InventoryItem(nil), filtered[start:end]...) + return response } -func isFirewallMetadataOnlyUpdate(before, after filter.FirewallRule, locator filter.Locator) (bool, error) { - beforeKey, err := filter.RuleKey(before) - if err != nil { - return false, err - } - afterKey, err := filter.RuleKey(after) - if err != nil { - return false, err +func isDeletableManagedInventoryItem(item filter.InventoryItem) bool { + if item.Desired == nil || item.Desired.Protected || item.State == filter.InventoryStateProtected { + return false } - if beforeKey != afterKey { - return false, nil + if item.Desired.Origin != filter.RuleOriginCreated && item.Desired.Origin != filter.RuleOriginAdopted { + return false } - if after.Scope.Provider == filter.ProviderFirewalld { - return true, nil + if (item.Rule.Scope.Provider == filter.ProviderIptables || item.Rule.Scope.Provider == filter.ProviderNftables) && + (item.Rule.Scope.Chain == filter.BasicBeforeChain || item.Rule.Scope.Chain == filter.BasicAfterChain) { + return false } - return locator.Position != nil && after.OrderIndex != nil && *after.OrderIndex == int64(*locator.Position), nil + return item.State != filter.InventoryStateDrifted || + (item.Match == filter.InventoryMatchMissing && item.Observed == nil) } -func (s *FirewallService) updateRuleOrder(ctx context.Context, ruleUUID string, targetPosition *int64, priority *int, description *string) error { - if (targetPosition == nil) == (priority == nil) { - return fmt.Errorf("%w: provide either position or priority", filter.ErrInvalidRule) - } - if ruleUUID == "" { - return fmt.Errorf("%w: rule UUID is required", repo.ErrFirewallPersistenceInvalid) - } - stored, before, snapshot, observed, runtime, err := s.loadManagedMutation(ctx, ruleUUID) - if err != nil { - return err +func matchesFirewallInventoryRequest(item filter.InventoryItem, request dto.FirewallRuleInventory) bool { + if slices.Contains(request.ExcludeChains, item.Rule.Scope.Chain) { + return false } - capabilities, err := runtime.Capabilities(ctx) - if err != nil { - return err + if len(request.Families) > 0 && !matchesFirewallInventoryFamily(item.Rule, request.Families) { + return false } - after := before.Rule - adapterOperation := filter.ChangeReorder - switch { - case capabilities.ExplicitPosition || capabilities.OwnedChains: - if targetPosition == nil || *targetPosition < 1 { - return fmt.Errorf("%w: target position is required", filter.ErrInvalidRule) - } - if err := runtime.ValidatePosition(ctx, snapshot, before.Rule, *targetPosition); err != nil { - return err - } - after.OrderIndex = targetPosition - case capabilities.ExplicitPriority: - if before.Rule.NativeKind != filter.NativeKindRichRule { - return fmt.Errorf("%w: only rich rules support explicit priority", filter.ErrUnsupportedScope) - } - if priority == nil { - return fmt.Errorf("%w: priority is required", filter.ErrInvalidRule) - } - after.Priority = priority - adapterOperation = filter.ChangeUpdate - default: - return fmt.Errorf("%w: provider does not support rule reordering", filter.ErrUnsupportedScope) + if len(request.Actions) > 0 && !matchesFirewallInventoryAction(item.Rule.Action, request.Actions) { + return false } - if description != nil { - after.Description = strings.TrimSpace(*description) + if len(request.States) > 0 && !slices.Contains(request.States, item.State) { + return false } - after, err = runtime.Prepare(after) - if err != nil { - return err + keyword := strings.ToLower(strings.TrimSpace(request.Info)) + if keyword == "" { + return true } - if err := runtime.CheckRule(ctx, after); err != nil { - return err + rule := item.Rule + values := []string{ + firewallInventoryProtocol(rule), rule.SourceAddress, rule.SourcePort, rule.DestinationAddress, + rule.DestinationPort, rule.Description, string(rule.Action), string(item.State), } - metadataOnly, err := isFirewallMetadataOnlyUpdate(before.Rule, after, observed.Locator) - if err != nil { - return err + if item.Observed != nil { + values = append(values, item.Observed.Rule.Description) } - if metadataOnly { - return s.updateRuleDescription(ctx, stored.UUID, after.Description) + if item.Desired != nil { + values = append(values, item.Desired.Rule.Description) } - if err := filter.GuardMutation(observed); err != nil { - return err + for _, value := range values { + if strings.Contains(strings.ToLower(value), keyword) { + return true + } } - return s.executeManagedMutation(ctx, managedMutationRequest{ - Stored: stored, Before: before.Rule, After: after, Snapshot: snapshot, Locator: observed.Locator, - AdapterOperation: adapterOperation, Runtime: runtime, - }) -} - -type managedMutationRequest struct { - Stored model.FirewallRule - Before filter.FirewallRule - After filter.FirewallRule - Snapshot filter.Snapshot - Locator filter.Locator - AdapterOperation filter.ChangeOperation - Runtime *filterruntime.Engine -} - -type preparedManagedUpdate struct { - Stored model.FirewallRule - Before filter.DesiredRule - After filter.FirewallRule - Snapshot filter.Snapshot - Observed filter.ObservedRule - Runtime *filterruntime.Engine + return false } -func (s *FirewallService) prepareManagedUpdate( - ctx context.Context, - clientIP string, - ruleUUID string, - requestedRule filter.FirewallRule, -) (preparedManagedUpdate, error) { - ruleUUID = strings.TrimSpace(ruleUUID) - if ruleUUID == "" { - return preparedManagedUpdate{}, fmt.Errorf("%w: rule UUID is required", repo.ErrFirewallPersistenceInvalid) - } - stored, before, snapshot, observed, runtime, err := s.loadManagedMutation(ctx, ruleUUID) - if err != nil { - return preparedManagedUpdate{}, err - } - after, err := filter.NormalizeRule(requestedRule) - if err != nil { - return preparedManagedUpdate{}, err - } - after.UUID = stored.UUID - after, err = runtime.Prepare(after) - if err != nil { - return preparedManagedUpdate{}, err - } - if err := runtime.CheckRule(ctx, after); err != nil { - return preparedManagedUpdate{}, err - } - if after.Scope.Key() != before.Rule.Scope.Key() { - return preparedManagedUpdate{}, filter.ErrManagedScopeChange - } - if !supportsManagedNativeKindTransition(before.Rule, after) { - return preparedManagedUpdate{}, fmt.Errorf("%w: native rule conversion requires an explicit workflow", filter.ErrUnsupportedScope) - } - capabilities, err := runtime.Capabilities(ctx) - if err != nil { - return preparedManagedUpdate{}, err - } - if capabilities.ExplicitPosition || capabilities.OwnedChains { - if observed.Locator.Position == nil { - return preparedManagedUpdate{}, fmt.Errorf("%w: managed rule has no positional locator", filter.ErrInvalidRule) +func matchesFirewallInventoryFamily(rule filter.FirewallRule, families []filter.Family) bool { + for _, family := range families { + if rule.Scope.Family != filter.FamilyInet && rule.Scope.Family == family { + return true } - currentPosition := int64(*observed.Locator.Position) - if after.OrderIndex == nil { - after.OrderIndex = ¤tPosition - } else if *after.OrderIndex != currentPosition { - if err := runtime.ValidatePosition(ctx, snapshot, before.Rule, *after.OrderIndex); err != nil { - return preparedManagedUpdate{}, err - } + if rule.Scope.Family == filter.FamilyInet && + (rule.SourceAddress == "" || (family == filter.FamilyIPv6) == strings.Contains(rule.SourceAddress, ":")) { + return true } } - if err := filter.GuardMutation(observed); err != nil { - return preparedManagedUpdate{}, err - } - if err := s.checkManagedMutationCollisions(ctx, before.Rule, after, snapshot, observed.Locator, stored.UUID); err != nil { - return preparedManagedUpdate{}, err - } - return preparedManagedUpdate{ - Stored: stored, Before: before, After: after, Snapshot: snapshot, Observed: observed, Runtime: runtime, - }, nil + return false } -func supportsManagedNativeKindTransition(before, after filter.FirewallRule) bool { - if before.NativeKind == after.NativeKind { - return true - } - if before.Scope.Key() != after.Scope.Key() || before.Scope.Provider != filter.ProviderFirewalld { - return false +func matchesFirewallInventoryAction(action filter.Action, actions []string) bool { + for _, requested := range actions { + if requested == "accept" && action == filter.ActionAccept { + return true + } + if requested == "deny" && (action == filter.ActionDrop || action == filter.ActionReject) { + return true + } } - return before.NativeKind == filter.NativeKindZonePort && after.NativeKind == filter.NativeKindRichRule || - before.NativeKind == filter.NativeKindRichRule && after.NativeKind == filter.NativeKindZonePort + return false } -func (s *FirewallService) loadManagedMutation( - ctx context.Context, - ruleUUID string, -) (model.FirewallRule, filter.DesiredRule, filter.Snapshot, filter.ObservedRule, *filterruntime.Engine, error) { - stored, err := s.rules.GetByUUID(ctx, ruleUUID) - if err != nil { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, err - } - if stored.Origin != constant.FirewallRuleOriginCreated && stored.Origin != constant.FirewallRuleOriginAdopted { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, - fmt.Errorf("%w: only created or adopted rules can be changed", filter.ErrInvalidRule) - } - selected, err := s.selectedProviderForStoredRule(ctx, stored) - if err != nil { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, err - } - if err := checkFirewallRuleWhitelistProtection(selected, stored); err != nil { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, err - } - desiredRules, err := s.compileStoredFirewallRules(ctx, stored, selected) - if err != nil { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, err +func firewallInventoryProtocol(rule filter.FirewallRule) string { + if rule.NativeKind == filter.NativeKindZoneService { + return "service" } - if len(desiredRules) != 1 { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, - fmt.Errorf("%w: policy %q expands to %d target rules and cannot be edited atomically", filter.ErrUnsupportedScope, ruleUUID, len(desiredRules)) + if rule.NativeKind == filter.NativeKindUFWApplication && rule.Protocol == "" { + return "app" } - desired := desiredRules[0] - runtime, err := s.resolveRuntime(ctx, desired.Rule.Scope.Provider) - if err != nil { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, err + if rule.Scope.Provider == filter.ProviderUFW && rule.Protocol == "all" && rule.DestinationPort != "" { + return "tcp/udp" } - snapshot, err := runtime.ObserveMutation(ctx, desired.Rule.Scope) - if err != nil { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, err + return rule.Protocol +} + +func (s *FirewallService) deleteRules(ctx context.Context, request dto.FirewallRuleDelete, t *task.Task) (result dto.FirewallRuleDeleteResponse, taskErr error) { + var firstFailure error + defer func() { + sort.SliceStable(result.Errors, func(i, j int) bool { return result.Errors[i].Index < result.Errors[j].Index }) + if t != nil { + t.Log(i18n.GetMsgWithMap("FirewallRuleOperationResult", map[string]interface{}{ + "succeeded": result.Succeeded, "failed": result.Failed, + })) + } + if taskErr == nil { + taskErr = firstFailure + } + }() + record := func(index int, ruleUUID string, err error) { + if err != nil { + result.Failed++ + result.Errors = append(result.Errors, dto.FirewallRuleDeleteFailure{Index: index, UUID: ruleUUID, Error: err.Error()}) + if firstFailure == nil { + firstFailure = err + } + } else { + result.Succeeded++ + } + if t != nil { + label := fmt.Sprintf("[%d/%d] %s", result.Succeeded+result.Failed, len(request.UUIDs)+len(request.BeforeRules), ruleUUID) + t.LogWithStatus(label, err) + } } - observed, err := filter.ManagedObserved(snapshot, desired) + selectedProvider, err := s.selectedProvider(ctx) if err != nil { - return model.FirewallRule{}, filter.DesiredRule{}, filter.Snapshot{}, filter.ObservedRule{}, nil, err + return result, err } - return stored, desired, snapshot, observed, runtime, nil -} - -func (s *FirewallService) selectedProviderForStoredRule( - ctx context.Context, - _ model.FirewallRule, -) (filter.Provider, error) { - if s.selectedProvider != nil { - return s.selectedProvider(ctx) + type beforeGroup struct { + targets []dto.FirewallRuleDeleteTarget + indexes []int } - if s.adapters != nil { - providers := s.adapters.Providers() - if len(providers) == 1 { - return providers[0], nil + beforeGroups := make(map[string]*beforeGroup) + for index, target := range request.BeforeRules { + key := target.Scope.Normalize().Key() + if beforeGroups[key] == nil { + beforeGroups[key] = &beforeGroup{} } + group := beforeGroups[key] + group.targets = append(group.targets, target) + group.indexes = append(group.indexes, len(request.UUIDs)+index) } - return "", fmt.Errorf("%w: selected provider is unavailable", filter.ErrProviderUnavailable) -} - -func (s *FirewallService) executeManagedMutation(ctx context.Context, request managedMutationRequest) error { - before, after := request.Before, request.After - appendRule, restoreAtEnd := false, false - if after.Scope.Provider == filter.ProviderUFW && (request.AdapterOperation == filter.ChangeUpdate || request.AdapterOperation == filter.ChangeReorder) { - maxPosition := maxObservedFirewallPosition(request.Snapshot) - appendRule = after.OrderIndex != nil && *after.OrderIndex == maxPosition - restoreAtEnd = request.Locator.Position != nil && int64(*request.Locator.Position) == maxPosition - } - backendPlan, verification, err := request.Runtime.Execute(ctx, request.Snapshot, []filter.DesiredChange{{ - Operation: request.AdapterOperation, - Before: &before, - After: &after, - Locator: &request.Locator, - Append: appendRule, - RestoreAtEnd: restoreAtEnd, - }}) - if err != nil { - return err + for _, group := range beforeGroups { + failures, batchErr := s.deleteBeforeRules(ctx, selectedProvider, group.targets) + for index, target := range group.targets { + err := batchErr + if failures != nil && failures[index] != nil { + err = failures[index] + } + record(group.indexes[index], target.InstanceKey, err) + } } - if !verification.Matched { - return filter.ErrVerificationFailed + if len(request.UUIDs) == 0 { + return result, nil } - _, err = filter.FindCommittedObserved(verification.Snapshot, request.After, backendPlan) - if err != nil { - return rollbackFirewallPlan(ctx, request.Runtime, backendPlan, err) + ports, err := loadFirewallPortWhiteList() + var runtime filter.Adapter + if err == nil { + runtime, err = s.firewallAdapter(selectedProvider) } - updates, err := firewallRuleSemanticUpdates(request.After) if err != nil { - return rollbackFirewallPlan(ctx, request.Runtime, backendPlan, err) + for index, value := range request.UUIDs { + record(index, strings.TrimSpace(value), err) + } + return result, err } - if request.After.Scope.Provider == filter.ProviderFirewalld { - updates["sequence"] = nil - } else { - position, positionErr := firewallRuleMarkerPosition(verification.Snapshot, request.Stored.UUID) - if positionErr != nil { - return rollbackFirewallPlan(ctx, request.Runtime, backendPlan, positionErr) + whitelist := filter.NewPortWhitelistIndex(ports) + groups := make([][]firewallRuleDeleteItem, 0) + groupIndexes := make(map[string]int) + seen := make(map[string]bool, len(request.UUIDs)) + for index, value := range request.UUIDs { + ruleUUID := strings.TrimSpace(value) + if err := ctx.Err(); err != nil { + record(index, ruleUUID, err) + continue } - sequence, sequenceErr := s.sequenceForFirewallRulePosition( - ctx, verification.Snapshot, position, request.Stored.UUID, request.Stored.Sequence, - ) - if sequenceErr != nil { - return rollbackFirewallPlan(ctx, request.Runtime, backendPlan, sequenceErr) + if ruleUUID == "" { + record(index, ruleUUID, fmt.Errorf("%w: rule UUID is required", repo.ErrFirewallPersistenceInvalid)) + continue } - updates["sequence"] = sequence - } - if err := s.rules.UpdateWithRevision(ctx, request.Stored.UUID, request.Stored.Revision, updates); err != nil { - return rollbackFirewallPlan(ctx, request.Runtime, backendPlan, err) + if seen[ruleUUID] { + record(index, ruleUUID, fmt.Errorf("duplicate firewall rule UUID")) + continue + } + seen[ruleUUID] = true + stored, err := s.rules.GetByUUID(ctx, ruleUUID) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + err = fmt.Errorf("%w: managed rule %q was not found", filter.ErrInvalidRule, ruleUUID) + } + record(index, ruleUUID, err) + continue + } + if stored.Origin != constant.FirewallRuleOriginCreated && stored.Origin != constant.FirewallRuleOriginAdopted { + record(index, ruleUUID, fmt.Errorf("%w: only created or adopted rules can be deleted", filter.ErrInvalidRule)) + continue + } + rules, err := compileStoredFirewallRules(ctx, stored, runtime) + if err == nil && len(rules) == 0 { + err = fmt.Errorf("%w: policy %q has no compiled target rules", filter.ErrInvalidRule, ruleUUID) + } + if err == nil { + for _, rule := range rules { + if whitelist.Matches(rule.Rule) { + err = filter.ErrProtectedRule + break + } + } + } + if err != nil { + record(index, ruleUUID, err) + continue + } + key := rules[0].Rule.Scope.Key() + if selectedProvider == filter.ProviderUFW { + key = string(selectedProvider) + } + groupIndex, exists := groupIndexes[key] + if !exists { + groupIndex = len(groups) + groupIndexes[key] = groupIndex + groups = append(groups, nil) + } + groups[groupIndex] = append(groups[groupIndex], firewallRuleDeleteItem{index: index, stored: stored, rules: rules}) } - return nil -} - -func maxObservedFirewallPosition(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) + for _, group := range groups { + failures, batchErr := s.deleteFirewallRuleBatch(ctx, runtime, group) + for index, item := range group { + err := batchErr + if failures != nil && failures[index] != nil { + err = failures[index] + } + record(item.index, item.stored.UUID, err) } } - return maximum -} -func (s *FirewallService) checkManagedMutationCollisions( - ctx context.Context, - before, after filter.FirewallRule, - snapshot filter.Snapshot, - locator filter.Locator, - excludedUUID string, -) error { - sameContent, err := filter.SameRuleContent(before, after) - if err != nil { - return err - } - if sameContent { - return nil - } - if err := filter.CheckObservedRuleCollisions(snapshot, after, &locator); err != nil { - return err - } - return s.ensureFirewallRuleIdentityAvailable(ctx, after, excludedUUID) + return result, nil } -func (s *FirewallService) ensureFirewallRuleIdentityAvailable(ctx context.Context, requested filter.FirewallRule, excludedUUID string) error { - stored, err := s.rules.List(ctx) +func (s *FirewallService) deleteBeforeRules(ctx context.Context, provider filter.Provider, targets []dto.FirewallRuleDeleteTarget) ([]error, error) { + scope := targets[0].Scope.Normalize() + if scope.Provider != provider || (provider != filter.ProviderIptables && provider != filter.ProviderNftables) || scope.Chain != filter.BasicBeforeChain { + return nil, fmt.Errorf("%w: native deletion only supports the selected firewall before chain", filter.ErrUnsupportedScope) + } + runtime, err := s.firewallAdapter(provider) if err != nil { - return err + return nil, err } - identities, err := firewallRuleCollisions(stored, requested.Scope.Provider, excludedUUID) + snapshot, err := readMutableFirewallRules(runtime, ctx, scope) if err != nil { - return err + return nil, err } - return identities.Check(requested) -} - -func firewallRuleCollisions(stored []model.FirewallRule, provider filter.Provider, excludedUUID string) (filter.RuleCollisionIndex, error) { - identities := make(filter.RuleCollisionIndex, len(stored)) - for _, candidate := range stored { - if candidate.UUID == excludedUUID { + failures := make([]error, len(targets)) + changes := make([]filter.RuleChange, 0, len(targets)) + seen := make(map[string]bool, len(targets)) + byInstance := make(map[string][]int, len(snapshot.Rules)) + for index, observed := range snapshot.Rules { + key, err := filter.InstanceKey(observed) + if err == nil { + byInstance[key] = append(byInstance[key], index) + } + } + for index, target := range targets { + if target.Scope.Normalize().Key() != scope.Key() || seen[target.InstanceKey] { + failures[index] = fmt.Errorf("%w: duplicate or mismatched before rule", filter.ErrInvalidRule) continue } - rules, err := candidate.RulesForProvider(provider) - if err != nil { + seen[target.InstanceKey] = true + matches := byInstance[target.InstanceKey] + if len(matches) != 1 { + failures[index] = filter.ErrRuleStale continue } - for _, rule := range rules { - if err := identities.Add(rule); err != nil { - return nil, err - } + observed := snapshot.Rules[matches[0]] + if err := filter.GuardMutation(observed); err != nil { + failures[index] = err + continue } - } - return identities, nil -} - -func (s *FirewallService) resolveRuntime(ctx context.Context, provider filter.Provider) (*filterruntime.Engine, error) { - if s.selectedProvider != nil { - selected, err := s.selectedProvider(ctx) - if err != nil { - return nil, err + if observed.ParseStatus != filter.ParseStatusSupported || observed.Locator.Position == nil { + failures[index] = fmt.Errorf("%w: before rule cannot be deleted", filter.ErrUnsupportedScope) + continue } - if selected != provider { - return nil, fmt.Errorf("%w: selected provider is %s, requested %s", filter.ErrProviderUnavailable, selected, provider) + rule := observed.Rule + if rule.UUID == "" && strings.HasPrefix(observed.Marker, "1panel-rule:") { + rule.UUID = strings.TrimSpace(strings.TrimPrefix(observed.Marker, "1panel-rule:")) + } + if rule.UUID == "" { + rule.UUID = uuid.NewString() } + locator := observed.Locator + changes = append(changes, filter.RuleChange{ + Operation: filter.ChangeDelete, Before: &rule, Locator: &locator, UnmarkedAdopted: observed.Marker == "", CommandOnly: true, + }) } - if s.adapters == nil { - return nil, filter.ErrAdapterUnavailable + if len(changes) == 0 { + return failures, nil } - return s.adapters.Resolve(provider) -} - -func (s *FirewallService) ensureSystemPort(ctx context.Context, port dto.FirewallSystemPort) error { - firewallRuleMutationMu.Lock() - err := s.ensureSystemPortLocked(ctx, port) - firewallRuleMutationMu.Unlock() - if err == nil { - return nil + sort.Slice(changes, func(i, j int) bool { return *changes[i].Locator.Position > *changes[j].Locator.Position }) + if err := ctx.Err(); err != nil { + return failures, err } - if errors.Is(err, filter.ErrInventoryUnavailable) { - err = s.appendUFWSystemPortUnverified(ctx, port, err) + plan, err := runtime.BuildCommands(snapshot, changes) + if err == nil { + plan.CommandOnly = true + err = runtime.RunCommands(ctx, plan) } - if port.Family == constant.FirewallFamilyIPv6 && filterufw.IsIPv6Unavailable(err) { - if global.LOG != nil { - global.LOG.Warnf("skip accepted UFW IPv6 port %s/%s: %v", port.Port, port.Protocol, err) - } - return nil + if err == nil { + err = persistFirewallRules(ctx, runtime, plan) } - return err + return failures, err } -func (s *FirewallService) appendUFWSystemPortUnverified( - ctx context.Context, - port dto.FirewallSystemPort, - cause error, -) error { - if s.selectedProvider == nil || s.adapters == nil { - return cause - } - provider, providerErr := s.selectedProvider(ctx) - if providerErr != nil { - return errors.Join(cause, providerErr) - } - if provider != filter.ProviderUFW { - return cause - } - if global.LOG != nil { - global.LOG.Warnf( - "UFW inventory is unavailable while restoring accepted port %s/%s; attempting a restricted direct allow: %v", - port.Port, port.Protocol, cause, - ) +func (s *FirewallService) deleteFirewallRuleBatch(ctx context.Context, runtime filter.Adapter, items []firewallRuleDeleteItem) ([]error, error) { + byScope := make(map[string]filter.RuleSet) + byMarker := make(map[string]map[string][]filter.ObservedRule) + if runtime.Provider() == filter.ProviderNftables { + scopes := make([]filter.Scope, 0) + for _, item := range items { + for _, desired := range item.rules { + scopes = append(scopes, desired.Rule.Scope) + } + } + snapshots, err := readMutableFirewallRuleScopes(runtime, ctx, scopes) + if err != nil { + return nil, err + } + for _, snapshot := range snapshots { + byScope[snapshot.Scope.Key()] = snapshot + byMarker[snapshot.Scope.Key()] = firewallRulesByMarker(snapshot.Rules) + } } - runtime, resolveErr := s.adapters.Resolve(provider) - if resolveErr != nil { - return errors.Join(cause, resolveErr) + failures := make([]error, len(items)) + changes := make([]firewallRuleBatchItem, 0, len(items)) + indexes := make([]int, 0, len(items)) + for index, item := range items { + for _, desired := range item.rules { + if desired.Protected { + failures[index] = filter.ErrProtectedRule + break + } + snapshot := filter.RuleSet{Scope: desired.Rule.Scope} + change := filter.RuleChange{Operation: filter.ChangeDelete, Before: &desired.Rule, CommandOnly: true} + if runtime.Provider() == filter.ProviderNftables { + snapshot = byScope[desired.Rule.Scope.Key()] + candidates := snapshot + if desired.Marker != "" { + candidates.Rules = byMarker[snapshot.Scope.Key()][desired.Marker] + } + observed, err := managedFirewallObserved(candidates, desired) + if errors.Is(err, filter.ErrRuleStale) { + missing, missingErr := managedFirewallRuleMissing(snapshot, desired) + if missingErr != nil { + err = missingErr + } else if missing { + continue + } + } + if err == nil { + change, err = firewallDeleteChange(observed, desired) + } + if err != nil { + failures[index] = err + break + } + } + changes = append(changes, firewallRuleBatchItem{snapshot: snapshot, change: change}) + indexes = append(indexes, index) + } } - comment := "1panel-system-port:" + systemPortKey(port) - if appendErr := runtime.AppendUnverified(ctx, systemPortRule(provider, port), comment); appendErr != nil { - if global.LOG != nil { - global.LOG.Errorf( - "restore accepted UFW port %s/%s without rule inventory failed: %v; original error: %v", - port.Port, port.Protocol, appendErr, cause, - ) + validChanges, validIndexes := changes[:0], indexes[:0] + for index, change := range changes { + if failures[indexes[index]] == nil { + validChanges = append(validChanges, change) + validIndexes = append(validIndexes, indexes[index]) } - return errors.Join(cause, fmt.Errorf("append accepted UFW port without rule inventory: %w", appendErr)) } - if global.LOG != nil { - global.LOG.Warnf( - "restored accepted UFW port %s/%s without rule inventory; normal rule management failed: %v", - port.Port, port.Protocol, cause, - ) + executeFirewallRuleBatches(ctx, runtime, validChanges, nil, func(index int, failure error) { + if failure != nil { + failures[validIndexes[index]] = failure + } + }) + deleted := make([]model.FirewallRule, 0, len(items)) + for index, item := range items { + if failures[index] == nil { + deleted = append(deleted, item.stored) + } } - return nil + persistFailures := s.rules.DeleteBatchWithRevision(context.WithoutCancel(ctx), deleted) + for index, item := range items { + if failures[index] == nil { + failures[index] = persistFailures[item.stored.UUID] + } + } + return failures, nil } -func (s *FirewallService) ensureSystemPortLocked(ctx context.Context, port dto.FirewallSystemPort) error { - provider, err := s.selectedProvider(ctx) - if err != nil { - return err - } - source := dto.FirewallRuleCreateItem{ - Rule: systemPortRule(provider, port), SourceKind: constant.FirewallRuleSourceSecurity, - SourceID: systemPortSourceID(port), - } - prepared, err := s.prepareCreate(ctx, provider, source) - if err != nil { - return err - } - rule, runtime := prepared.request.Rule, prepared.runtime - snapshot, err := runtime.ObserveMutation(ctx, rule.Scope) +func managedFirewallRuleMissing(snapshot filter.RuleSet, desired filter.DesiredRule) (bool, error) { + wanted, err := filter.RuleKey(desired.Rule) if err != nil { - return err + return false, err } - stored, err := s.rules.List(ctx) - if err != nil { - return err + if desired.RuleKey != "" && desired.RuleKey != wanted { + return false, filter.ErrInvalidRule } - matchKey, err := filter.RuleMatchKey(rule) + wanted, err = firewallInventoryRuleKey(desired.Rule) if err != nil { - return err + return false, err } - for _, record := range stored { - candidates, err := record.RulesForProvider(provider) - if err != nil { - continue - } - matched := false - for _, candidate := range candidates { - key, err := filter.RuleMatchKey(candidate) - if err != nil { - return err + for _, observed := range snapshot.Rules { + if desired.Marker != "" { + if observed.Marker == desired.Marker { + return false, nil } - if key == matchKey { - matched = true - break + adopted := desired.Origin == filter.RuleOriginAdopted && observed.Marker == "" + legacy := observed.Marker == "1panel-rule:"+desired.UUID && observed.Marker != desired.Marker + if (adopted || legacy) && filter.ObservedRuleMatchesExpected(observed, desired.Rule) { + return false, nil } + continue } - if !matched { + if desired.ObservedInstanceKey != "" { + key, err := filter.InstanceKey(observed) + if err == nil && key == desired.ObservedInstanceKey { + return false, nil + } continue } - if record.Action != string(rule.Action) { - return filter.ErrRuleConflict + if observed.ParseStatus != filter.ParseStatusSupported { + continue } - desired, err := s.compileStoredFirewallRules(ctx, record, provider) + key, err := firewallInventoryRuleKey(observed.Rule) if err != nil { - return err + return false, err } - items, err := filter.MergeInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desired}) + if key == wanted { + return false, nil + } + } + return true, nil +} + +func (s *FirewallService) updateRule(ctx context.Context, clientIP, ruleUUID string, requestedRule filter.FirewallRule) error { + requestedRule, err := filter.NormalizeRule(requestedRule) + if err != nil { + return err + } + stored, err := s.rules.GetByUUID(ctx, ruleUUID) + if err != nil { + return err + } + previousRules, compileErr := expandStoredFirewallRule(stored, requestedRule.Scope.Provider) + if compileErr == nil && len(previousRules) == 1 { + sameContent, err := filter.SameRuleContent(previousRules[0], requestedRule) if err != nil { return err } - for _, item := range items { - if item.Desired != nil && item.Match == filter.InventoryMatchExact && item.State != filter.InventoryStateDrifted && item.Observed != nil { - return nil + if sameContent { + if requestedRule.Scope.Provider == filter.ProviderFirewalld { + if requestedRule.Priority != nil { + return s.updateRuleOrder(ctx, ruleUUID, nil, requestedRule.Priority, &requestedRule.Description) + } + } else if requestedRule.OrderIndex != nil { + return s.updateRuleOrder(ctx, ruleUUID, requestedRule.OrderIndex, nil, &requestedRule.Description) } + return s.updateRuleDescription(ctx, ruleUUID, requestedRule.Description) } - return filter.ErrRuleStale } - matches, err := filter.MatchObservedByRuleKey(snapshot.Rules, rule) + prepared, err := s.prepareManagedUpdate(ctx, clientIP, ruleUUID, requestedRule) if err != nil { return err } - if len(matches) > 0 { - if matches[0].Protected { - if matches[0].Persistence != "" && matches[0].Persistence != filter.PersistenceStatusConverged { - return filter.ErrRuleStale - } + metadataOnly, err := isFirewallMetadataOnlyUpdate(prepared.Before.Rule, prepared.After, prepared.Observed.Locator) + if err != nil { + return err + } + if metadataOnly { + if prepared.After.Description == prepared.Before.Rule.Description { return nil } - return s.adoptRule(ctx, runtime, snapshot, matches[0], source) + return s.rules.UpdateWithRevision(ctx, prepared.Stored.UUID, prepared.Stored.Revision, map[string]interface{}{ + "description": prepared.After.Description, + }) } - _, err = s.createRules(ctx, dto.FirewallRuleCreate{Items: []dto.FirewallRuleCreateItem{source}}, nil) - return err -} - -func hasSystemFirewallRuleOwner(rule model.FirewallRule) bool { - acceptedPrefix := model.FirewallRuleOwner( - constant.FirewallRuleSourceSecurity, - constant.FirewallSystemAcceptedPortSourcePrefix, - ) - return strings.HasPrefix(rule.Owner, acceptedPrefix) -} - -func systemPortRule(provider filter.Provider, port dto.FirewallSystemPort) filter.FirewallRule { - return firewall.RuleForSystemPort(provider, firewall.SystemPort(port)) -} - -func systemPortKey(port dto.FirewallSystemPort) string { - return firewall.SystemPortKey(firewall.SystemPort(port)) -} - -func systemPortSourceID(port dto.FirewallSystemPort) string { - return constant.FirewallSystemAcceptedPortSourcePrefix + systemPortKey(port) + if prepared.Runtime.Provider() != filter.ProviderNftables { + return s.replaceManagedRule(ctx, prepared) + } + return s.executeManagedMutation(ctx, managedMutationRequest{ + Stored: prepared.Stored, Before: prepared.Before.Rule, After: prepared.After, + RuleSet: prepared.RuleSet, Locator: prepared.Observed.Locator, + AdapterOperation: filter.ChangeUpdate, Runtime: prepared.Runtime, + }) } -func firewallRuleModelForCreate(rule filter.FirewallRule, request dto.FirewallRuleCreateItem, origin string) (model.FirewallRule, error) { - record, err := model.FirewallRuleFromDomain(rule) +func (s *FirewallService) prepareManagedUpdate(ctx context.Context, clientIP string, ruleUUID string, requestedRule filter.FirewallRule) (preparedManagedUpdate, error) { + ruleUUID = strings.TrimSpace(ruleUUID) + if ruleUUID == "" { + return preparedManagedUpdate{}, fmt.Errorf("%w: rule UUID is required", repo.ErrFirewallPersistenceInvalid) + } + stored, before, runtime, err := s.loadManagedRule(ctx, ruleUUID) if err != nil { - return model.FirewallRule{}, err + return preparedManagedUpdate{}, err } - record.Origin = origin - record.Owner = model.FirewallRuleOwner(request.SourceKind, request.SourceID) - return record, nil -} - -func firewallRuleSemanticUpdates(rule filter.FirewallRule) (map[string]interface{}, error) { - record, err := model.FirewallRuleFromDomain(rule) + after, err := filter.NormalizeRule(requestedRule) if err != nil { - return nil, err + return preparedManagedUpdate{}, err } - return map[string]interface{}{ - "family": record.Family, "protocol": record.Protocol, - "source_address": record.SourceAddress, "source_port": record.SourcePort, - "destination_address": record.DestinationAddress, "destination_port": record.DestinationPort, - "interface": record.Interface, "connection_states": record.ConnectionStates, "action": record.Action, - "description": record.Description, "compatibility_error": "", "priority": record.Priority, - }, nil -} - -func (s *FirewallService) nextFirewallRuleSequence(ctx context.Context) (int64, error) { - stored, err := s.rules.List(ctx) + after.UUID = stored.UUID + after, err = prepareFirewallBackendRule(ctx, runtime, after) if err != nil { - return 0, err + return preparedManagedUpdate{}, err } - var maximum int64 - for _, record := range stored { - if record.Sequence != nil && *record.Sequence > maximum { - maximum = *record.Sequence + if after.Scope.Key() != before.Rule.Scope.Key() { + return preparedManagedUpdate{}, buserr.New("ErrFirewallRuleScopeChange") + } + if !supportsManagedNativeKindTransition(before.Rule, after) { + return preparedManagedUpdate{}, fmt.Errorf("%w: native rule conversion requires an explicit workflow", filter.ErrUnsupportedScope) + } + if runtime.Provider() != filter.ProviderNftables { + if after.OrderIndex != nil && *after.OrderIndex < 1 { + return preparedManagedUpdate{}, fmt.Errorf("%w: target position must be positive", filter.ErrInvalidRule) } + return preparedManagedUpdate{Stored: stored, Before: before, After: after, Runtime: runtime}, nil } - return maximum + model.FirewallRuleSequenceStep, nil -} - -func (s *FirewallService) sequenceForCreatedFirewallRule( - ctx context.Context, - snapshot filter.Snapshot, - rule filter.FirewallRule, -) (int64, error) { - if rule.OrderIndex == nil { - return s.nextFirewallRuleSequence(ctx) - } - return s.sequenceForFirewallRulePosition(ctx, snapshot, int(*rule.OrderIndex), "", nil) -} - -func (s *FirewallService) sequenceForFirewallRulePosition( - ctx context.Context, - snapshot filter.Snapshot, - targetPosition int, - excludedUUID string, - current *int64, -) (int64, error) { - stored, err := s.rules.List(ctx) + snapshot, err := readMutableFirewallRules(runtime, ctx, before.Rule.Scope) if err != nil { - return 0, err + return preparedManagedUpdate{}, err } - byUUID := make(map[string]model.FirewallRule, len(stored)) - for _, record := range stored { - byUUID[record.UUID] = record - compiled, err := s.compileStoredFirewallRules(ctx, record, snapshot.Scope.Provider) - if err != nil { - if isFirewallPolicyIncompatible(err) { - continue - } - return 0, err - } - for _, desired := range compiled { - if desired.Rule.Scope.Key() == snapshot.Scope.Key() { - byUUID[desired.Rule.UUID] = record - } - } + observed, err := managedFirewallObserved(snapshot, before) + if err != nil { + return preparedManagedUpdate{}, err } - var previous, next *model.FirewallRule - for _, observed := range snapshot.Rules { - uuid := strings.TrimPrefix(observed.Marker, "1panel-rule:") - if observed.Marker == uuid || uuid == excludedUUID || observed.Locator.Position == nil { - continue - } - record, exists := byUUID[uuid] - if !exists || record.UUID == excludedUUID { - continue - } - position := *observed.Locator.Position - if position < targetPosition { - copy := record - previous = © - } else if position > targetPosition || excludedUUID == "" { - copy := record - next = © - break - } - } - if previous != nil && previous.Sequence == nil || next != nil && next.Sequence == nil { - return s.rebalanceFirewallRuleSequences(ctx, snapshot, targetPosition, excludedUUID, byUUID) - } - if current != nil && - (previous == nil || *previous.Sequence < *current) && (next == nil || *current < *next.Sequence) { - return *current, nil - } - switch { - case previous == nil && next == nil: - return model.FirewallRuleSequenceStep, nil - case previous == nil: - return *next.Sequence - model.FirewallRuleSequenceStep, nil - case next == nil: - return *previous.Sequence + model.FirewallRuleSequenceStep, nil - case *next.Sequence-*previous.Sequence > 1: - return *previous.Sequence + (*next.Sequence-*previous.Sequence)/2, nil - default: - return s.rebalanceFirewallRuleSequences(ctx, snapshot, targetPosition, excludedUUID, byUUID) + capabilities, err := runtime.Capabilities(ctx) + if err != nil { + return preparedManagedUpdate{}, err } -} - -func (s *FirewallService) rebalanceFirewallRuleSequences( - ctx context.Context, - snapshot filter.Snapshot, - targetPosition int, - excludedUUID string, - byUUID map[string]model.FirewallRule, -) (int64, error) { - targetSequence := int64(targetPosition) * model.FirewallRuleSequenceStep - updated := make(map[string]bool) - for _, observed := range snapshot.Rules { + if capabilities.ExplicitPosition || capabilities.OwnedChains { if observed.Locator.Position == nil { - continue - } - uuid := strings.TrimPrefix(observed.Marker, "1panel-rule:") - if observed.Marker == uuid || uuid == excludedUUID { - continue - } - record, exists := byUUID[uuid] - if !exists || record.UUID == excludedUUID || updated[record.UUID] { - continue - } - updated[record.UUID] = true - position := *observed.Locator.Position - if excludedUUID == "" && position >= targetPosition { - position++ - } - sequence := int64(position) * model.FirewallRuleSequenceStep - if record.Sequence != nil && *record.Sequence == sequence { - continue + return preparedManagedUpdate{}, fmt.Errorf("%w: managed rule has no positional locator", filter.ErrInvalidRule) } - if err := s.rules.UpdateWithRevision(ctx, record.UUID, record.Revision, map[string]interface{}{ - "sequence": sequence, - }); err != nil { - return 0, err + currentPosition := int64(*observed.Locator.Position) + if after.OrderIndex == nil { + after.OrderIndex = ¤tPosition + } else if *after.OrderIndex != currentPosition { + if err := validateFirewallRulePosition(snapshot, before.Rule, *after.OrderIndex); err != nil { + return preparedManagedUpdate{}, err + } } } - return targetSequence, nil -} - -func firewallRuleMarkerPosition(snapshot filter.Snapshot, ruleUUID string) (int, error) { - marker := "1panel-rule:" + ruleUUID - for _, observed := range snapshot.Rules { - if observed.Marker == marker && observed.Locator.Position != nil { - return *observed.Locator.Position, nil - } + if err := filter.GuardMutation(observed); err != nil { + return preparedManagedUpdate{}, err } - return 0, fmt.Errorf("%w: committed firewall rule %q has no position", filter.ErrVerificationFailed, ruleUUID) + return preparedManagedUpdate{ + Stored: stored, Before: before, After: after, RuleSet: snapshot, Observed: observed, Runtime: runtime, + }, nil } -func (s *FirewallService) compileStoredFirewallRules( - ctx context.Context, - stored model.FirewallRule, - target filter.Provider, -) ([]filter.DesiredRule, error) { - rules, err := stored.RulesForProvider(target) - if err != nil { - return nil, err +func supportsManagedNativeKindTransition(before, after filter.FirewallRule) bool { + if before.NativeKind == after.NativeKind { + return true } - runtime, err := s.adapters.Resolve(target) - if err != nil { - return nil, err + if before.Scope.Key() != after.Scope.Key() || before.Scope.Provider != filter.ProviderFirewalld { + return false } - return runtime.CompileDesired(ctx, stored.UUID, filter.RuleOrigin(stored.Origin), rules) + return before.NativeKind == filter.NativeKindZonePort && after.NativeKind == filter.NativeKindRichRule || + before.NativeKind == filter.NativeKindRichRule && after.NativeKind == filter.NativeKindZonePort } -func isFirewallPolicyIncompatible(err error) bool { - return errors.Is(err, filter.ErrInvalidRule) || errors.Is(err, filter.ErrUnsupportedScope) || - errors.Is(err, filter.ErrInvalidScope) || errors.Is(err, filter.ErrCompositeRule) -} +func (s *FirewallService) ensureWebsitePorts(ctx context.Context, ports map[string]firewall.SystemPort) error { + if len(ports) == 0 { + return nil + } + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() -func (s *FirewallService) compileRestorableFirewallRules( - ctx context.Context, - stored model.FirewallRule, - provider filter.Provider, -) (restorable, preserved []filter.DesiredRule, err error) { - compiled, err := s.compileStoredFirewallRules(ctx, stored, provider) + provider, err := s.selectedProvider(ctx) if err != nil { - return nil, nil, err - } - if !supportsManagedFilterChains(string(provider)) || !hasSystemFirewallRuleOwner(stored) { - return compiled, nil, nil - } - loadRequired := s.requiredPorts - if loadRequired == nil { - loadRequired = LoadRequiredFirewallPortWhiteList + return err } - required, err := loadRequired() + runtime, err := s.firewallAdapter(provider) if err != nil { - return nil, nil, err + return err } - requiredPorts := firewall.ExpandPortWhitelist(required) - for _, desired := range compiled { - covered := false - for _, port := range requiredPorts { - covered, err = filter.SameRuleContent(desired.Rule, systemPortRule(provider, port)) - if err != nil { - return nil, nil, err + const comment = "1panel website" + rules := make([]filter.FirewallRule, 0, len(ports)) + var scopes []filter.Scope + seenScopes := make(map[string]bool) + for _, key := range firewall.SortedSystemPortKeys(ports) { + prepared, err := prepareFirewallCreateRule(ctx, runtime, dto.FirewallRuleCreateItem{Rule: firewall.RuleForSystemPort(provider, ports[key])}) + if err != nil { + return err + } + rule := prepared.Rule + rule.Description = comment + rules = append(rules, rule) + if !seenScopes[rule.Scope.Key()] { + seenScopes[rule.Scope.Key()] = true + scopes = append(scopes, rule.Scope) + } + } + external, supportsExternal := runtime.(filter.ExternalRuleAdapter) + existing := make(map[string]bool) + if provider != filter.ProviderFirewalld { + if !supportsExternal { + return fmt.Errorf("%w: %s does not support external rules", filter.ErrAdapterUnavailable, provider) + } + observed, err := external.ListRulesByComment(ctx, scopes, comment) + if err != nil { + return err + } + for _, candidate := range observed { + if candidate.ParseStatus != filter.ParseStatusSupported || candidate.Rule.Description != comment || candidate.Rule.Action != filter.ActionAccept { + continue } - if covered { - break + key, err := filter.RuleMatchKey(candidate.Rule) + if err != nil { + return err } + existing[key] = true } - if covered { - preserved = append(preserved, desired) + } + var failures []error + changedScopes := make(map[string]filter.Scope) + for _, rule := range rules { + key, err := filter.RuleMatchKey(rule) + if err != nil { + return err + } + if existing[key] { + continue + } + if provider == filter.ProviderFirewalld { + rule.UUID = uuid.NewString() + plan, buildErr := runtime.BuildCommands(filter.RuleSet{Scope: rule.Scope}, []filter.RuleChange{{Operation: filter.ChangeCreate, After: &rule, CommandOnly: true, Append: true}}) + if buildErr != nil { + return buildErr + } + plan.CommandOnly = true + err = runtime.RunCommands(ctx, plan) + if errors.Is(err, filterfirewalld.ErrAlreadyEnabled) { + err = nil + } } else { - restorable = append(restorable, desired) + err = external.AppendUnverified(ctx, rule, comment) } - } - return restorable, preserved, nil -} - -func (s *FirewallService) desiredFirewallRulesByScope( - ctx context.Context, - stored []model.FirewallRule, - provider filter.Provider, -) (map[string][]filter.DesiredRule, []filter.InventoryItem) { - model.SortFirewallRules(stored, provider) - desired := make(map[string][]filter.DesiredRule) - var failures []filter.InventoryItem - ports, protectionErr := loadFirewallPortWhiteList() - for _, record := range stored { - compiled, _, err := s.compileRestorableFirewallRules(ctx, record, provider) - if err == nil { - err = protectionErr + if rule.Scope.Family == filter.FamilyIPv6 && filterufw.IsIPv6Unavailable(err) { + continue } if err != nil { - rule := filter.FirewallRule{ - UUID: record.UUID, - Scope: filter.Scope{Provider: provider, Family: filter.Family(record.Family), Direction: filter.DirectionInput}.Normalize(), - Protocol: record.Protocol, SourceAddress: record.SourceAddress, SourcePort: record.SourcePort, - DestinationAddress: record.DestinationAddress, DestinationPort: record.DestinationPort, - Interface: record.Interface, ConnectionStates: strings.FieldsFunc(record.ConnectionStates, func(r rune) bool { return r == ',' }), - Action: filter.Action(record.Action), Description: record.Description, Priority: record.Priority, - } - failures = append(failures, filter.InventoryItem{ - Incompatible: isFirewallPolicyIncompatible(err), - Rule: rule, State: filter.InventoryStateDrifted, Match: filter.InventoryMatchNone, - Desired: &filter.DesiredRule{UUID: record.UUID, Rule: rule, Origin: filter.RuleOrigin(record.Origin), Protected: protectionErr != nil || filter.RuleMatchesPortWhitelist(rule, ports)}, - Error: fmt.Sprintf("policy %s: %v", record.UUID, err), - }) + failures = append(failures, fmt.Errorf("allow website port %s/%s (%s): %w", rule.DestinationPort, rule.Protocol, rule.Scope.Family, err)) continue } - for _, rule := range compiled { - rule.Protected = filter.RuleMatchesPortWhitelist(rule.Rule, ports) - rule.Expanded = len(compiled) > 1 - key := rule.Rule.Scope.Key() - desired[key] = append(desired[key], rule) + existing[key] = true + saveKey := rule.Scope.Key() + if provider == filter.ProviderNftables { + saveKey = string(provider) } + changedScopes[saveKey] = rule.Scope } - return desired, failures -} - -func firewallRuleSnapshotPolicy(ctx context.Context, snapshot filter.Snapshot) (filter.Snapshot, error) { - ports, err := loadFirewallPortWhiteList() - if err != nil { - return filter.Snapshot{}, err - } - return filter.ProtectSnapshot(snapshot, ports) -} - -func firewallRuleSelectedProvider(context.Context) (filter.Provider, error) { - provider, err := selectedSystemFirewallProvider() - if err != nil { - return "", fmt.Errorf("%w: %v", filter.ErrProviderUnavailable, err) - } - return filter.Provider(provider), nil -} - -func rollbackFirewallPlan(ctx context.Context, runtime *filterruntime.Engine, 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)) + if saver, ok := runtime.(filter.RuleSaver); ok { + for _, scope := range changedScopes { + saveCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + err := saver.SaveRules(saveCtx, scope) + cancel() + if err != nil { + failures = append(failures, err) + } + } } - return cause + return errors.Join(failures...) } func ensureFirewallPorts(ports []int) error { - filterruntime.InvalidateInventory() - defer filterruntime.InvalidateInventory() - client, err := selectedSystemFirewallClient() + if len(ports) == 0 { + return nil + } + client, err := NewSelectedSystemFirewallClient() if err != nil { return err } @@ -2877,106 +1614,7 @@ func ensureFirewallPorts(ports []int) error { if err != nil { return err } - service := newFirewallService() - var failures []error - for _, key := range firewall.SortedSystemPortKeys(normalized) { - if err := service.ensureSystemPort(context.Background(), normalized[key]); err != nil { - wrapped := fmt.Errorf("allow firewall port %s: %w", key, err) - if supportsNativeRuleBatch(filter.Provider(state.Name)) { - return wrapped - } - failures = append(failures, wrapped) - if global.LOG != nil { - global.LOG.Errorf("%v", wrapped) - } - } - } - return errors.Join(failures...) -} - -func LoadPanelPort() string { - if !global.IsMaster { - return global.CONF.Base.Port - } - var portSetting model.Setting - _ = global.CoreDB.Where("key = ?", "ServerPort").First(&portSetting).Error - return portSetting.Value -} - -func loadFirewallPortWhiteList() ([]firewall.PortWhitelist, error) { - ports, err := loadPortWhitelistSetting(global.DB) - if err != nil { - return nil, err - } - return firewall.ValidatePortWhitelist(ports) -} - -func LoadRequiredFirewallPortWhiteList() ([]firewall.PortWhitelist, error) { - ports, err := loadFirewallPortWhiteList() - if err != nil { - return nil, err - } - return firewall.RequiredPortWhitelist(ports) -} - -func newIptablesHelperManager() *iptables_helper.Manager { - return &iptables_helper.Manager{ - UpdateSetting: settingRepo.Update, - LoadRequiredPorts: LoadRequiredFirewallPortWhiteList, - } -} - -func newNftablesHelperManager() *nftables_helper.Manager { - return &nftables_helper.Manager{ - UpdateSetting: settingRepo.Update, - LoadRequiredPorts: LoadRequiredFirewallPortWhiteList, - } -} - -func loadFirewallInitStatus(provider, tab string) (bool, bool, error) { - switch provider { - case constant.FirewallProviderNftables: - return nftables_helper.LoadInitStatus(tab) - case constant.FirewallProviderIptables: - return iptables_helper.LoadInitStatus(tab) - default: - return false, false, fmt.Errorf("unsupported firewall provider: %s", provider) - } -} - -func supportsManagedFilterChains(provider string) bool { - return provider == constant.FirewallProviderIptables || provider == constant.FirewallProviderNftables -} - -func (s *FirewallService) restoreFirewallAfterStart(client lifecycle.Client) error { - ctx := context.Background() - provider := filter.Provider(client.Name()) - var recoveryErrors []error - recordFailure := func(stage string, err error) { - if err == nil { - return - } - wrapped := fmt.Errorf("%s for %s: %w", stage, provider, err) - recoveryErrors = append(recoveryErrors, wrapped) - if global.LOG != nil { - global.LOG.Errorf("firewall post-start recovery failed: %v", wrapped) - } - } - if provider == filter.ProviderIptables || provider == filter.ProviderNftables { - isInit, _, err := loadFirewallInitStatus(string(provider), "base") - if err != nil { - recordFailure("load managed chain status", err) - return errors.Join(recoveryErrors...) - } - if !isInit { - return nil - } - } - if err := s.restoreStoredFirewallRules(ctx, provider, nil); err != nil { - recordFailure("restore stored firewall rules", err) - } - recordFailure("restore whitelist allowances", s.SyncPortWhitelist(ctx)) - return errors.Join(recoveryErrors...) + return newFirewallService().ensureWebsitePorts(context.Background(), normalized) } func AdoptLegacyHostFirewallRuleOwnership(ctx context.Context) error { @@ -2994,7 +1632,7 @@ func (s *FirewallService) adoptLegacyHostFirewallRuleOwnership(ctx context.Conte if selected != filter.ProviderIptables && selected != filter.ProviderUFW { return fmt.Errorf("%w: selected provider %s does not require legacy ownership transfer", filter.ErrProviderUnavailable, selected) } - runtime, err := s.adapters.Resolve(selected) + runtime, err := s.firewallAdapter(selected) if err != nil { return err } @@ -3002,7 +1640,7 @@ func (s *FirewallService) adoptLegacyHostFirewallRuleOwnership(ctx context.Conte if err != nil { return err } - desiredByScope, failures := s.desiredFirewallRulesByScope(ctx, stored, selected) + desiredByScope, failures := s.desiredFirewallRulesByScope(ctx, stored, runtime) if len(failures) > 0 { return errors.New(failures[0].Error) } @@ -3011,42 +1649,43 @@ func (s *FirewallService) adoptLegacyHostFirewallRuleOwnership(ctx context.Conte if len(desired) == 0 { continue } - snapshot, err := runtime.ObserveMutation(ctx, scope) + snapshot, err := readMutableFirewallRules(runtime, ctx, scope) if err != nil { return err } - for { - items, err := filter.MergeInventory(filter.InventoryMergeInput{ - Observed: snapshot.Rules, - Desired: desired, - }) - if err != nil { + items, err := mergeFirewallInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desired}) + if err != nil { + return err + } + for _, item := range items { + if item.Match != filter.InventoryMatchChanged || item.Desired == nil || item.Observed == nil || + item.Desired.Origin != filter.RuleOriginAdopted || strings.TrimSpace(item.Desired.Marker) == "" || + strings.TrimSpace(item.Observed.Marker) != "" || item.Observed.Protected || + !filter.ObservedRuleMatchesExpected(*item.Observed, item.Desired.Rule) { + continue + } + if err := ctx.Err(); err != nil { return err } - var candidate *filter.InventoryItem - for index := range items { - item := &items[index] - if item.Match != filter.InventoryMatchChanged || item.Desired == nil || item.Observed == nil || - item.Desired.Origin != filter.RuleOriginAdopted || strings.TrimSpace(item.Desired.Marker) == "" || - strings.TrimSpace(item.Observed.Marker) != "" || item.Observed.Protected || - !filter.ObservedRuleMatchesExpected(*item.Observed, item.Desired.Rule) { + matches := make([]filter.ObservedRule, 0, 1) + for _, observed := range snapshot.Rules { + if observed.Marker != "" || observed.Rule.SourceAddress != item.Observed.Rule.SourceAddress || + observed.Rule.DestinationAddress != item.Observed.Rule.DestinationAddress || + observed.Rule.SourcePort != item.Observed.Rule.SourcePort || observed.Rule.DestinationPort != item.Observed.Rule.DestinationPort { continue } - candidate = item - break + if filter.ObservedRuleMatchesExpected(observed, item.Desired.Rule) { + matches = append(matches, observed) + } } - if candidate == nil { - break + if len(matches) != 1 { + return filter.ErrRuleStale } - after := candidate.Desired.Rule - before := firewallsync.ObservedRule(*candidate.Observed) - locator := candidate.Observed.Locator - _, verification, err := runtime.Execute(ctx, snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeAdopt, - Before: &before, - After: &after, - Locator: &locator, - PreviousMarker: candidate.Observed.Marker, + after := item.Desired.Rule + before := matches[0].Rule + locator := matches[0].Locator + _, verification, err := applyFirewallChanges(runtime, ctx, snapshot, []filter.RuleChange{{ + Operation: filter.ChangeAdopt, Before: &before, After: &after, Locator: &locator, PreviousMarker: matches[0].Marker, }}) if err != nil { return err @@ -3054,7 +1693,7 @@ func (s *FirewallService) adoptLegacyHostFirewallRuleOwnership(ctx context.Conte if !verification.Matched { return filter.ErrVerificationFailed } - snapshot = verification.Snapshot + snapshot = verification.RuleSet } } if selected == filter.ProviderIptables { diff --git a/agent/app/service/firewall_docker.go b/agent/app/service/firewall_docker.go index bb0e912ddf13..82ca60ff5dd9 100644 --- a/agent/app/service/firewall_docker.go +++ b/agent/app/service/firewall_docker.go @@ -5,14 +5,10 @@ import ( "encoding/json" "errors" "fmt" - "net/netip" "os" - "path/filepath" "sort" - "strconv" "strings" "sync" - "time" "github.com/1Panel-dev/1Panel/agent/app/dto" "github.com/1Panel-dev/1Panel/agent/app/model" @@ -20,17 +16,11 @@ import ( "github.com/1Panel-dev/1Panel/agent/app/task" "github.com/1Panel-dev/1Panel/agent/buserr" "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/cmd" - "github.com/1Panel-dev/1Panel/agent/utils/docker" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" - containertypes "github.com/docker/docker/api/types/container" - "github.com/docker/docker/api/types/system" + "github.com/1Panel-dev/1Panel/agent/i18n" + dockerfirewall "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/google/uuid" - "gorm.io/gorm" ) const ( @@ -48,41 +38,15 @@ const ( dockerReasonNoMatchingPath = "no_matching_path" ) -type dockerProxyEndpoint struct { - protocol string - hostIP string - hostPort uint16 -} - -type dockerForwardRules struct { - output string - inspected bool -} - -type dockerProxyEndpoints struct { - items []dockerProxyEndpoint - inspected bool -} - -type dockerGuardRuntime = docker_guard.Runtime - type DockerPortGuardService struct { policies repo.IDockerPortGuardRepo - runtime dockerGuardRuntime - runtimeForBackend func(string) dockerGuardRuntime + runtime dockerfirewall.Runtime + runtimeForBackend func(string) dockerfirewall.Runtime client func() (*client.Client, error) version func(string) string } -var ( - dockerPortGuardServiceMu sync.Mutex - dockerPortGuardSyncMu sync.RWMutex - dockerPortGuardSyncErr error - ErrDockerGuardInvalid = docker_guard.ErrInvalidPolicy - ErrDockerUnavailable = docker.ErrUnavailable - ErrDockerIptablesChainUnavailable = docker_guard.ErrDockerIptablesChainUnavailable - ErrDockerNftablesChainUnavailable = docker_guard.ErrDockerNftablesChainUnavailable -) +var dockerPortGuardServiceMu sync.Mutex type IDockerPortGuardService interface { LoadOverview(context.Context) (dto.DockerPortGuardList, error) @@ -94,56 +58,6 @@ type IDockerPortGuardService interface { Reconcile(context.Context) error } -func NewIDockerPortGuardService() IDockerPortGuardService { - return newDockerPortGuardService() -} - -func newDockerPortGuardService() *DockerPortGuardService { - return &DockerPortGuardService{ - policies: repo.NewIDockerPortGuardRepo(), - client: docker.NewDockerClient, - version: dockerFirewallVersion, - } -} - -func ReconcileDockerPortGuard(ctx context.Context) error { - if global.DB == nil { - return nil - } - return NewIDockerPortGuardService().Reconcile(ctx) -} - -func ReconcileDockerPortGuardBestEffort(ctx context.Context) { - if err := ReconcileDockerPortGuard(ctx); err != nil { - global.LOG.Warnf("reconcile Docker port guard failed, err: %v", err) - } -} - -func (s *DockerPortGuardService) LoadPublishedPorts(ctx context.Context) ([]dto.DockerPortGuardContainer, error) { - cli, err := s.client() - if err != nil { - return nil, fmt.Errorf("%w: %v", ErrDockerUnavailable, err) - } - defer cli.Close() - - if socketPath, local := strings.CutPrefix(cli.DaemonHost(), "unix://"); local { - if _, statErr := os.Stat(socketPath); errors.Is(statErr, os.ErrNotExist) { - return []dto.DockerPortGuardContainer{}, nil - } - } - - endpoints, err := discoverDockerEndpoints(ctx, cli, false) - if err != nil { - return nil, err - } - backend := selectedDockerFirewallBackend("") - if info, infoErr := cli.Info(ctx); infoErr == nil { - backend = dockerFirewallBackend(info) - } - annotateDockerEndpointManagement(endpoints, backend) - return groupDockerGuardContainers(endpoints), nil -} - func (s *DockerPortGuardService) LoadOverview(ctx context.Context) (dto.DockerPortGuardList, error) { policies, err := s.policies.ListManaged(ctx) if err != nil { @@ -152,8 +66,7 @@ func (s *DockerPortGuardService) LoadOverview(ctx context.Context) (dto.DockerPo unavailable := func() dto.DockerPortGuardList { backend := selectedDockerFirewallBackend("") base := s.runtimeStatus(s.guardRuntime(backend), backend) - base.Version = s.loadFirewallVersion(backend) - base.Message = agenti18n.Get("ErrDockerFailed") + base.Message = i18n.Get("ErrDockerFailed") return dto.DockerPortGuardList{Base: base, Containers: []dto.DockerPortGuardContainer{}, OrphanPolicies: dockerGuardPolicyEndpoints(policies)} } cli, err := s.client() @@ -168,10 +81,6 @@ func (s *DockerPortGuardService) LoadOverview(ctx context.Context) (dto.DockerPo detectedBackend := dockerFirewallBackend(info) backend := selectedDockerFirewallBackend(detectedBackend) base := s.runtimeStatus(s.guardRuntime(backend), backend) - base.Version = s.loadFirewallVersion(backend) - if reconcileErr := lastDockerPortGuardReconcileError(); reconcileErr != nil { - markDockerGuardReconcileFailure(&base, reconcileErr) - } endpoints, err := discoverDockerEndpoints(ctx, cli, true) if err != nil { return dto.DockerPortGuardList{}, err @@ -179,58 +88,37 @@ func (s *DockerPortGuardService) LoadOverview(ctx context.Context) (dto.DockerPo annotateDockerEndpointManagement(endpoints, detectedBackend) endpoints, orphanPolicies := matchDockerGuardPolicies(base, policies, endpoints) sort.Slice(endpoints, func(i, j int) bool { - return guardEndpointKey(endpoints[i].Family, endpoints[i].HostIP, endpoints[i].HostPort, endpoints[i].Protocol) < guardEndpointKey(endpoints[j].Family, endpoints[j].HostIP, endpoints[j].HostPort, endpoints[j].Protocol) + return fmt.Sprintf("%s|%s|%d|%s", endpoints[i].Family, endpoints[i].HostIP, endpoints[i].HostPort, endpoints[i].Protocol) < fmt.Sprintf("%s|%s|%d|%s", endpoints[j].Family, endpoints[j].HostIP, endpoints[j].HostPort, endpoints[j].Protocol) }) sort.Slice(orphanPolicies, func(i, j int) bool { - return guardEndpointKey(orphanPolicies[i].Family, orphanPolicies[i].HostIP, orphanPolicies[i].HostPort, orphanPolicies[i].Protocol) < guardEndpointKey(orphanPolicies[j].Family, orphanPolicies[j].HostIP, orphanPolicies[j].HostPort, orphanPolicies[j].Protocol) + return fmt.Sprintf("%s|%s|%d|%s", orphanPolicies[i].Family, orphanPolicies[i].HostIP, orphanPolicies[i].HostPort, orphanPolicies[i].Protocol) < fmt.Sprintf("%s|%s|%d|%s", orphanPolicies[j].Family, orphanPolicies[j].HostIP, orphanPolicies[j].HostPort, orphanPolicies[j].Protocol) }) return dto.DockerPortGuardList{Base: base, Containers: groupDockerGuardContainers(endpoints), OrphanPolicies: orphanPolicies}, nil } -func matchDockerGuardPolicies( - base dto.DockerPortGuardBase, - policies []model.DockerPortGuardPolicy, - endpoints []dto.DockerPortGuardEndpoint, -) ([]dto.DockerPortGuardEndpoint, []dto.DockerPortGuardEndpoint) { - byEndpoint := make(map[string]model.DockerPortGuardPolicy, len(policies)) - for _, policy := range policies { - byEndpoint[guardEndpointKey(policy.Family, policy.HostIP, policy.HostPort, policy.Protocol)] = policy +func (s *DockerPortGuardService) LoadPublishedPorts(ctx context.Context) ([]dto.DockerPortGuardContainer, error) { + cli, err := s.client() + if err != nil { + return nil, buserr.WithDetail("ErrDockerFailed", err.Error(), err) } - for i := range endpoints { - key := guardEndpointKey(endpoints[i].Family, endpoints[i].HostIP, endpoints[i].HostPort, endpoints[i].Protocol) - policy, ok := byEndpoint[key] - if !ok { - continue + defer cli.Close() + + if socketPath, local := strings.CutPrefix(cli.DaemonHost(), "unix://"); local { + if _, statErr := os.Stat(socketPath); errors.Is(statErr, os.ErrNotExist) { + return []dto.DockerPortGuardContainer{}, nil } - 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 = endpoints[i].ManagementTarget == dockerManagementContainerGuard && - ((policy.Family == docker_guard.FamilyIPv4 && base.IPv4.Effective) || (policy.Family == docker_guard.FamilyIPv6 && base.IPv6.Effective)) - delete(byEndpoint, key) - } - orphanPolicies := make([]dto.DockerPortGuardEndpoint, 0, len(byEndpoint)) - 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: docker_guard.DecodeSources(policy.Sources), Description: policy.Description, - TrafficPath: dockerTrafficPathUnknown, ManagementTarget: dockerManagementNeedsDiagnosis, - ManagementReason: dockerReasonNoMatchingPath, - }) } - return endpoints, orphanPolicies -} -func initializeDockerGuardRuntime(runtime dockerGuardRuntime, policies []docker_guard.Policy) error { - err := runtime.Initialize(policies) - if errors.Is(err, docker_guard.ErrDockerForwardPolicyDrop) { - family := "IPv4" - var familyErr *docker_guard.FamilyError - if errors.As(err, &familyErr) && familyErr.Family == docker_guard.FamilyIPv6 { - family = "IPv6" - } - return buserr.WithMap("ErrDockerForwardPolicyDrop", map[string]interface{}{"family": family}, err) + endpoints, err := discoverDockerEndpoints(ctx, cli, false) + if err != nil { + return nil, err + } + backend := selectedDockerFirewallBackend("") + if info, infoErr := cli.Info(ctx); infoErr == nil { + backend = dockerFirewallBackend(info) } - return err + annotateDockerEndpointManagement(endpoints, backend) + return groupDockerGuardContainers(endpoints), nil } func (s *DockerPortGuardService) Operate(ctx context.Context, request dto.DockerPortGuardOperation) error { @@ -246,8 +134,11 @@ func (s *DockerPortGuardService) Operate(ctx context.Context, request dto.Docker if err != nil { return err } - if err := initializeDockerGuardRuntime(runtime, policies); err != nil { - recordDockerPortGuardReconcileError(err) + inventory, err := runtime.ListPolicies() + if err != nil { + return err + } + if err := runtime.Initialize(policies, inventory); err != nil { return err } if err := settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, backend); err != nil { @@ -256,7 +147,6 @@ func (s *DockerPortGuardService) Operate(ctx context.Context, request dto.Docker if err := settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusEnable); err != nil { return err } - recordDockerPortGuardReconcileError(nil) return nil case "bind": runtime, _, err := s.runtimeForDocker(ctx) @@ -272,7 +162,7 @@ func (s *DockerPortGuardService) Operate(ctx context.Context, request dto.Docker if s.runtime != nil { err = s.runtime.Unbind() } else { - err = errors.Join(docker_guard.NewManager().Unbind(), docker_guard.NewNftablesManager().Unbind()) + err = errors.Join(dockerfirewall.NewIptables().Unbind(), dockerfirewall.NewNftables().Unbind()) } if err != nil { return err @@ -283,9 +173,7 @@ func (s *DockerPortGuardService) Operate(ctx context.Context, request dto.Docker } } -func (s *DockerPortGuardService) QueueInitialization( - request dto.DockerPortGuardOperation, -) (dto.FilterChainOperationResponse, error) { +func (s *DockerPortGuardService) QueueInitialization(request dto.DockerPortGuardOperation) (dto.FilterChainOperationResponse, error) { if request.Operation != "initialize" { return dto.FilterChainOperationResponse{}, fmt.Errorf("only Docker port guard initialization can be queued") } @@ -296,38 +184,38 @@ func (s *DockerPortGuardService) QueueInitialization( if err != nil { return dto.FilterChainOperationResponse{}, fmt.Errorf("create Docker port guard initialization task: %w", err) } - var runtime dockerGuardRuntime + var runtime dockerfirewall.Runtime var backend string - var policies []docker_guard.Policy - taskItem.AddSubTask(agenti18n.GetMsgByKey("FirewallInspectDockerGuardStep"), func(t *task.Task) error { + taskItem.AddSubTask(i18n.GetMsgByKey("FirewallInspectDockerGuardStep"), func(t *task.Task) error { var err error runtime, backend, err = s.runtimeForDocker(t.TaskCtx) if err != nil { return err } - policies, err = s.runtimePolicies(t.TaskCtx) - if err != nil { - return err - } t.Logf("backend=%s", backend) return nil }, nil) - taskItem.AddSubTask(agenti18n.GetWithName("FirewallInitializeDockerGuardStep", "Docker"), func(t *task.Task) error { + taskItem.AddSubTask(i18n.GetWithName("FirewallInitializeDockerGuardStep", "Docker"), func(t *task.Task) error { dockerPortGuardServiceMu.Lock() defer dockerPortGuardServiceMu.Unlock() + policies, err := s.runtimePolicies(t.TaskCtx) + if err != nil { + return err + } t.Logf("backend=%s", backend) - err := initializeDockerGuardRuntime(runtime, policies) - recordDockerPortGuardReconcileError(err) - return err + inventory, err := runtime.ListPolicies() + if err != nil { + return err + } + return runtime.Initialize(policies, inventory) }, nil) - taskItem.AddSubTask(agenti18n.GetMsgByKey("FirewallPersistDockerGuardStep"), func(t *task.Task) error { + taskItem.AddSubTask(i18n.GetMsgByKey("FirewallPersistDockerGuardStep"), func(t *task.Task) error { if err := settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, backend); err != nil { return err } if err := settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusEnable); err != nil { return err } - recordDockerPortGuardReconcileError(nil) return nil }, nil) if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { @@ -338,7 +226,7 @@ func (s *DockerPortGuardService) QueueInitialization( } func (s *DockerPortGuardService) DeletePolicies(request dto.DockerPortGuardPolicyBatchDelete) (dto.FilterChainOperationResponse, error) { - uuids, err := docker_guard.NormalizePolicyUUIDs(request.UUIDs) + uuids, err := normalizeDockerFirewallUUIDs(request.UUIDs) if err != nil { return dto.FilterChainOperationResponse{}, err } @@ -360,36 +248,51 @@ func (s *DockerPortGuardService) DeletePolicies(request dto.DockerPortGuardPolic } func (s *DockerPortGuardService) UpsertPolicies(request dto.DockerPortGuardPolicyBatch) (dto.FilterChainOperationResponse, error) { + if len(request.Policies) > filter.MaxAtomicExpansion { + return dto.FilterChainOperationResponse{}, fmt.Errorf("create or import at most %d rules per batch (after expansion)", filter.MaxAtomicExpansion) + } labels := make([]string, len(request.Policies)) + policies := make([]model.DockerPortGuardPolicy, 0, len(request.Policies)) + endpoints := make([]dto.DockerPortGuardEndpointIdentity, 0, len(request.Policies)) + count := 0 for i, policy := range request.Policies { labels[i] = fmt.Sprintf("[%d/%d] %s %s %s:%d %s", i+1, len(request.Policies), policy.Family, policy.Protocol, policy.HostIP, policy.HostPort, policy.Mode) + normalized, err := normalizeDockerFirewallPolicy(dockerfirewall.Policy{ + Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, + Protocol: policy.Protocol, Mode: policy.Mode, Sources: policy.Sources, + }) + if err != nil { + return dto.FilterChainOperationResponse{}, fmt.Errorf("%s: %w", labels[i], err) + } + if normalized.Mode == dockerfirewall.ModeAll { + count++ + } else { + count += len(normalized.Sources) + if normalized.Mode == dockerfirewall.ModeAllow { + count++ + } + } + if count > filter.MaxAtomicExpansion { + return dto.FilterChainOperationResponse{}, fmt.Errorf("create or import at most %d rules per batch (after expansion)", filter.MaxAtomicExpansion) + } + encoded, err := json.Marshal(normalized.Sources) + if err != nil { + return dto.FilterChainOperationResponse{}, fmt.Errorf("%s: %w", labels[i], err) + } + policies = append(policies, model.DockerPortGuardPolicy{ + UUID: uuid.NewString(), Family: normalized.Family, HostIP: normalized.HostIP, + HostPort: normalized.HostPort, Protocol: normalized.Protocol, Mode: normalized.Mode, + Sources: string(encoded), Description: strings.TrimSpace(policy.Description), + }) + endpoints = append(endpoints, dto.DockerPortGuardEndpointIdentity{ + Family: normalized.Family, HostIP: normalized.HostIP, HostPort: normalized.HostPort, Protocol: normalized.Protocol, + }) } return queueFirewallRuleTask(firewallTaskDocker, task.TaskUpdate, labels, func(ctx context.Context) error { dockerPortGuardServiceMu.Lock() defer dockerPortGuardServiceMu.Unlock() - policies := make([]model.DockerPortGuardPolicy, 0, len(request.Policies)) - endpoints := make([]dto.DockerPortGuardEndpointIdentity, 0, len(request.Policies)) - for i, policy := range request.Policies { - if err := ctx.Err(); err != nil { - return err - } - normalized, err := docker_guard.NormalizePolicy(docker_guard.Policy{ - Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, - Protocol: policy.Protocol, Mode: policy.Mode, Sources: policy.Sources, - }) - if err != nil { - return fmt.Errorf("%s: %w", labels[i], err) - } - encoded, err := json.Marshal(normalized.Sources) - if err != nil { - return fmt.Errorf("%s: %w", labels[i], err) - } - policies = append(policies, model.DockerPortGuardPolicy{ - UUID: uuid.NewString(), Family: normalized.Family, HostIP: normalized.HostIP, - HostPort: normalized.HostPort, Protocol: normalized.Protocol, Mode: normalized.Mode, - Sources: string(encoded), Description: strings.TrimSpace(policy.Description), - }) - endpoints = append(endpoints, policy.DockerPortGuardEndpointIdentity) + if err := ctx.Err(); err != nil { + return err } if err := s.rejectHostInputDockerGuardEndpoints(ctx, endpoints); err != nil { return err @@ -401,227 +304,26 @@ func (s *DockerPortGuardService) UpsertPolicies(request dto.DockerPortGuardPolic }) } -func (s *DockerPortGuardService) rejectHostInputDockerGuardEndpoints( - ctx context.Context, - requested []dto.DockerPortGuardEndpointIdentity, -) error { - if s.client == nil || len(requested) == 0 { - return nil - } - cli, err := s.client() - if err != nil { - return nil - } - defer cli.Close() - info, err := cli.Info(ctx) - if err != nil { - return nil - } - endpoints, err := discoverDockerEndpoints(ctx, cli, true) - if err != nil { - return nil - } - annotateDockerEndpointManagement(endpoints, dockerFirewallBackend(info)) - targets := make(map[string]string, len(endpoints)) - for _, endpoint := range endpoints { - targets[guardEndpointKey(endpoint.Family, endpoint.HostIP, endpoint.HostPort, endpoint.Protocol)] = endpoint.ManagementTarget - } - for _, endpoint := range requested { - target := targets[guardEndpointKey(endpoint.Family, endpoint.HostIP, endpoint.HostPort, endpoint.Protocol)] - if target == dockerManagementHostFirewall { - return fmt.Errorf("%w: endpoint traffic is handled by the host input firewall", ErrDockerGuardInvalid) - } - if target == dockerManagementNeedsDiagnosis { - return fmt.Errorf("%w: endpoint traffic management target requires diagnosis", ErrDockerGuardInvalid) - } - } - return nil -} - -func (s *DockerPortGuardService) Reconcile(ctx context.Context) error { - dockerPortGuardServiceMu.Lock() - defer dockerPortGuardServiceMu.Unlock() - return s.reconcileLocked(ctx) -} - -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: docker_guard.DecodeSources(policy.Sources), - } -} - -func dockerGuardReadOnlyPolicyUUID(policy docker_guard.ReadOnlyPolicy) string { - nativeRules, _ := json.Marshal(policy.NativeRules) - fingerprint := strings.Join([]string{ - policy.Policy.Family, policy.Policy.HostIP, strconv.Itoa(int(policy.Policy.HostPort)), - policy.Policy.Protocol, policy.Action, string(nativeRules), - }, "\x00") - return uuid.NewSHA1(uuid.NameSpaceOID, []byte(fingerprint)).String() -} - -func dockerGuardRuntimeReadOnlyModels(policies []docker_guard.ReadOnlyPolicy) ([]model.DockerPortGuardPolicy, error) { - result := make([]model.DockerPortGuardPolicy, 0, len(policies)) - for _, policy := range policies { - sources, err := json.Marshal(policy.Policy.Sources) - if err != nil { - return nil, err - } - nativeRules, err := json.Marshal(policy.NativeRules) - if err != nil { - return nil, err - } - result = append(result, model.DockerPortGuardPolicy{ - UUID: dockerGuardReadOnlyPolicyUUID(policy), - ReadOnly: true, - Family: policy.Policy.Family, HostIP: policy.Policy.HostIP, HostPort: policy.Policy.HostPort, - Protocol: policy.Policy.Protocol, Sources: string(sources), NativeAction: policy.Action, - NativeRules: string(nativeRules), Sequence: policy.Sequence, - }) - } - return result, nil -} - -func (s *DockerPortGuardService) replaceRuntimeReadOnlyPolicies( - ctx context.Context, - policies []docker_guard.ReadOnlyPolicy, -) error { - stored, err := dockerGuardRuntimeReadOnlyModels(policies) - if err != nil { - return err - } - return s.policies.ReplaceRuntimeReadOnly(ctx, stored) -} - -func (s *DockerPortGuardService) reconcileLocked(ctx context.Context) (err error) { - defer func() { recordDockerPortGuardReconcileError(err) }() - persistedEnabled, err := dockerPortGuardPersistedEnabled() - if err != nil { - return fmt.Errorf("load Docker port guard persisted status: %w", err) - } - initialized, err := s.anyRuntimeInitialized() - if err != nil { - return &docker_guard.FamilyError{Family: docker_guard.FamilyIPv4, Err: fmt.Errorf("inspect initialization: %w", err)} - } - if !initialized && !persistedEnabled { - return nil - } - runtime, _, err := s.runtimeForDocker(ctx) - if err != nil { - return err - } - initialized, err = runtime.Initialized(docker_guard.FamilyIPv4) - if err != nil { - return &docker_guard.FamilyError{Family: docker_guard.FamilyIPv4, Err: fmt.Errorf("inspect initialization: %w", err)} - } - if !initialized && !persistedEnabled { - return nil - } - policies, err := s.runtimePolicies(ctx) - if err != nil { - return err - } - inventory, err := runtime.ListPolicies() - if err != nil { - return err - } - if err := s.replaceRuntimeReadOnlyPolicies(ctx, inventory.ReadOnly); err != nil { - return err - } - if !initialized { - err = initializeDockerGuardRuntime(runtime, policies) - } else { - err = runtime.Reconcile(policies) - } - if err != nil { - return err - } - return docker_guard.Verify(runtime, policies, inventory.ReadOnly) -} - -func dockerPortGuardPersistedEnabled() (bool, error) { - status, err := settingRepo.GetValueByKey(constant.FirewallDockerPortGuardStatusKey) - if errors.Is(err, gorm.ErrRecordNotFound) { - return false, nil - } - return status == constant.StatusEnable, err -} - -func (s *DockerPortGuardService) anyRuntimeInitialized() (bool, error) { - if s.runtime != nil { - return s.runtime.Initialized(docker_guard.FamilyIPv4) - } - for _, runtime := range []dockerGuardRuntime{docker_guard.NewManager(), docker_guard.NewNftablesManager()} { - initialized, err := runtime.Initialized(docker_guard.FamilyIPv4) - if err != nil || initialized { - return initialized, err - } - } - return false, nil -} - -func recordDockerPortGuardReconcileError(err error) { - dockerPortGuardSyncMu.Lock() - dockerPortGuardSyncErr = err - dockerPortGuardSyncMu.Unlock() -} - -func lastDockerPortGuardReconcileError() error { - dockerPortGuardSyncMu.RLock() - defer dockerPortGuardSyncMu.RUnlock() - return dockerPortGuardSyncErr -} - -func markDockerGuardReconcileFailure(base *dto.DockerPortGuardBase, err error) { - var familyErr *docker_guard.FamilyError - if errors.As(err, &familyErr) { - markDockerGuardFamilyNotEffective(base, familyErr.Family) - if familyErr.Family == docker_guard.FamilyIPv4 { - markDockerGuardFamilyNotEffective(base, docker_guard.FamilyIPv6) - } - return - } - markDockerGuardFamilyNotEffective(base, docker_guard.FamilyIPv4) - markDockerGuardFamilyNotEffective(base, docker_guard.FamilyIPv6) -} - -func markDockerGuardFamilyNotEffective(base *dto.DockerPortGuardBase, family string) { - var status *dto.DockerPortGuardFamilyStatus - switch family { - case docker_guard.FamilyIPv4: - status = &base.IPv4 - case docker_guard.FamilyIPv6: - status = &base.IPv6 - default: - return - } - if !status.Initialized { - return - } - status.State = docker_guard.StatusNotEffective - status.Reason = docker_guard.ReasonInspectFailed - status.Effective = false +func NewIDockerPortGuardService() IDockerPortGuardService { + return newDockerPortGuardService() } -func (s *DockerPortGuardService) runtimePolicies(ctx context.Context) ([]docker_guard.Policy, error) { - stored, err := s.policies.ListManaged(ctx) - if err != nil { - return nil, err +func (s *DockerPortGuardService) runtimeStatus(runtime dockerfirewall.Runtime, backend string) dto.DockerPortGuardBase { + ipv4 := runtime.Status(dockerfirewall.FamilyIPv4) + ipv6 := runtime.Status(dockerfirewall.FamilyIPv6) + version := "-" + if s.version != nil { + version = s.version(backend) } - 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: docker_guard.DecodeSources(policy.Sources)}) + name := "iptables-docker" + if strings.EqualFold(strings.TrimSpace(backend), constant.FirewallProviderNftables) { + name = "nftables-docker" } - return policies, nil -} - -func (s *DockerPortGuardService) runtimeStatus(runtime dockerGuardRuntime, backend string) dto.DockerPortGuardBase { - ipv4 := runtime.Status(docker_guard.FamilyIPv4) - ipv6 := runtime.Status(docker_guard.FamilyIPv6) return dto.DockerPortGuardBase{ - Name: dockerFirewallDisplayName(backend), + Name: name, + Version: version, Backend: backend, - IsExist: ipv4.Reason != docker_guard.ReasonCommandMissing || ipv6.Reason != docker_guard.ReasonCommandMissing, + IsExist: ipv4.Reason != dockerfirewall.ReasonCommandMissing || ipv6.Reason != dockerfirewall.ReasonCommandMissing, Initialized: ipv4.Initialized || ipv6.Initialized, Bound: ipv4.Bound || ipv6.Bound, IPv4: dto.DockerPortGuardFamilyStatus{State: ipv4.State, Reason: ipv4.Reason, Initialized: ipv4.Initialized, Bound: ipv4.Bound, Effective: ipv4.Effective}, @@ -629,107 +331,14 @@ func (s *DockerPortGuardService) runtimeStatus(runtime dockerGuardRuntime, backe } } -func (s *DockerPortGuardService) guardRuntime(backend string) dockerGuardRuntime { - if s.runtimeForBackend != nil { - return s.runtimeForBackend(backend) - } - if s.runtime != nil { - return s.runtime - } - return docker_guard.NewRuntime(backend) -} - -func (s *DockerPortGuardService) runtimeForDocker(ctx context.Context) (dockerGuardRuntime, string, error) { - if s.runtime != nil { - return s.runtime, selectedDockerFirewallBackend(constant.FirewallProviderIptables), nil - } - cli, err := s.client() - if err != nil { - return nil, "", fmt.Errorf("%w: %v", ErrDockerUnavailable, err) - } - defer cli.Close() - info, err := cli.Info(ctx) - if err != nil { - return nil, "", fmt.Errorf("%w: %v", ErrDockerUnavailable, err) - } - backend := selectedDockerFirewallBackend(dockerFirewallBackend(info)) - if backend != constant.FirewallProviderIptables && backend != constant.FirewallProviderNftables { - return nil, backend, fmt.Errorf("Docker firewall backend %q is not supported", backend) - } - return s.guardRuntime(backend), backend, nil -} - -func dockerFirewallBackend(info system.Info) string { - if info.FirewallBackend == nil || info.FirewallBackend.Driver == "" { - return constant.FirewallProviderIptables - } - return strings.ToLower(info.FirewallBackend.Driver) -} - -func dockerFirewallDisplayName(backend string) string { - switch strings.ToLower(strings.TrimSpace(backend)) { - case constant.FirewallProviderNftables: - return "nftables-docker" - default: - return "iptables-docker" - } -} - -func (s *DockerPortGuardService) loadFirewallVersion(backend string) string { - if s.version == nil { - return "-" - } - return s.version(backend) -} - -func dockerFirewallVersion(backend string) string { - client, err := lifecycle.NewClientFor(backend) - if err != nil { - return "-" - } - version, err := client.Version() - if err != nil || strings.TrimSpace(version) == "" { - return "-" - } - return version -} - -func discoverDockerEndpoints(ctx context.Context, cli *client.Client, all bool) ([]dto.DockerPortGuardEndpoint, error) { - containers, err := cli.ContainerList(ctx, containertypes.ListOptions{All: all}) - if err != nil { - return nil, err - } - endpoints := make([]dto.DockerPortGuardEndpoint, 0) - for _, item := range containers { - name := strings.TrimPrefix(firstGuardString(item.Names), "/") - compose := item.Labels[dockerGuardComposeProjectLabel] - application := "" - if created, ok := item.Labels[dockerGuardComposeCreatedBy]; ok && created == "Apps" { - application = compose - } - for _, port := range item.Ports { - if port.PublicPort == 0 || (port.Type != "tcp" && port.Type != "udp") { - continue - } - family := docker_guard.FamilyIPv4 - hostIP := port.IP - if addr, err := netip.ParseAddr(hostIP); err == nil && addr.Is6() { - family = docker_guard.FamilyIPv6 - } else if hostIP == "" { - hostIP = "0.0.0.0" - } - endpoints = append(endpoints, dto.DockerPortGuardEndpoint{Family: family, HostIP: hostIP, HostPort: port.PublicPort, Protocol: port.Type, ContainerID: item.ID, ContainerName: name, ContainerState: item.State, ContainerPort: port.PrivatePort, Compose: compose, Application: application, Sources: []string{}}) - } - } - return endpoints, nil -} - func dockerGuardPolicyEndpoints(policies []model.DockerPortGuardPolicy) []dto.DockerPortGuardEndpoint { endpoints := make([]dto.DockerPortGuardEndpoint, 0, len(policies)) for _, policy := range policies { + sources := []string{} + _ = json.Unmarshal([]byte(policy.Sources), &sources) endpoints = append(endpoints, dto.DockerPortGuardEndpoint{ Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, - PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources), + PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: sources, Description: policy.Description, TrafficPath: dockerTrafficPathUnknown, ManagementTarget: dockerManagementNeedsDiagnosis, ManagementReason: dockerReasonNoMatchingPath, }) @@ -737,331 +346,69 @@ func dockerGuardPolicyEndpoints(policies []model.DockerPortGuardPolicy) []dto.Do return endpoints } -func groupDockerGuardContainers(endpoints []dto.DockerPortGuardEndpoint) []dto.DockerPortGuardContainer { - containers := make(map[string]*dto.DockerPortGuardContainer) - order := make([]string, 0) - for _, endpoint := range endpoints { - key := endpoint.ContainerID - if key == "" { - key = "__orphan__" - } - container, ok := containers[key] - if !ok { - container = &dto.DockerPortGuardContainer{ - Key: key, Name: endpoint.ContainerName, Compose: endpoint.Compose, - Application: endpoint.Application, Endpoints: []dto.DockerPortGuardEndpoint{}, - } - containers[key] = container - order = append(order, key) - } - container.Endpoints = append(container.Endpoints, endpoint) - } - sort.Slice(order, func(i, j int) bool { - return containers[order[i]].Name < containers[order[j]].Name - }) - - result := make([]dto.DockerPortGuardContainer, 0, len(order)) - for _, key := range order { - container := containers[key] - items := make([]docker.PortRangeItem, 0, len(container.Endpoints)) - for i, endpoint := range container.Endpoints { - sources := append([]string(nil), endpoint.Sources...) - sort.Strings(sources) - policyKey := fmt.Sprintf("%t|%s|%s|%t|%s|%s|%s", endpoint.PolicyUUID != "", endpoint.Mode, strings.Join(sources, ","), endpoint.Effective, endpoint.Description, endpoint.ManagementTarget, endpoint.ManagementReason) - items = append(items, docker.PortRangeItem{ - Key: endpoint.Family + "|" + endpoint.HostIP + "|" + endpoint.Protocol + "|" + policyKey, - PublicPort: endpoint.HostPort, PrivatePort: endpoint.ContainerPort, - HasPrivatePort: endpoint.ContainerPort != 0, Position: i, - }) - } - container.PortGroups = make([]dto.DockerPortGuardPortGroup, 0, len(items)) - for _, portRange := range docker.MergePortRanges(items) { - start := container.Endpoints[portRange.Start.Position] - address := start.HostIP - if strings.Contains(address, ":") { - address = "[" + address + "]" - } - ports := fmt.Sprintf("%d", portRange.Start.PublicPort) - if portRange.Start.PublicPort != portRange.End.PublicPort { - ports = fmt.Sprintf("%d-%d", portRange.Start.PublicPort, portRange.End.PublicPort) - } - container.PortGroups = append(container.PortGroups, dto.DockerPortGuardPortGroup{ - Key: fmt.Sprintf("%s|%d-%d", portRange.Start.Key, portRange.Start.PublicPort, portRange.End.PublicPort), - Label: fmt.Sprintf("%s:%s/%s", address, ports, start.Protocol), Endpoint: start, - Endpoints: func() []dto.DockerPortGuardEndpoint { - members := make([]dto.DockerPortGuardEndpoint, 0, len(portRange.Items)) - for _, item := range portRange.Items { - members = append(members, container.Endpoints[item.Position]) - } - return members - }(), - }) - } - result = append(result, *container) - } - return result -} - -func guardEndpointKey(family, hostIP string, hostPort uint16, protocol string) string { - return fmt.Sprintf("%s|%s|%d|%s", family, hostIP, hostPort, protocol) -} - -func firstGuardString(values []string) string { - if len(values) == 0 { - return "" - } - return values[0] -} - -func annotateDockerEndpointManagement(endpoints []dto.DockerPortGuardEndpoint, backend string) { - rules := map[string]dockerForwardRules{ - constant.FirewallFamilyIPv4: loadDockerDNATRules(backend, constant.FirewallFamilyIPv4), - constant.FirewallFamilyIPv6: loadDockerDNATRules(backend, constant.FirewallFamilyIPv6), +func matchDockerGuardPolicies(base dto.DockerPortGuardBase, policies []model.DockerPortGuardPolicy, endpoints []dto.DockerPortGuardEndpoint) ([]dto.DockerPortGuardEndpoint, []dto.DockerPortGuardEndpoint) { + byEndpoint := make(map[string]model.DockerPortGuardPolicy, len(policies)) + for _, policy := range policies { + byEndpoint[fmt.Sprintf("%s|%s|%d|%s", policy.Family, policy.HostIP, policy.HostPort, policy.Protocol)] = policy } - proxies := loadDockerProxyEndpoints() for i := range endpoints { - familyRules := rules[endpoints[i].Family] - endpoints[i].TrafficPath, endpoints[i].ManagementTarget, endpoints[i].ManagementReason = - dockerEndpointManagement(backend, familyRules, proxies, endpoints[i]) - } -} - -func dockerEndpointManagement( - backend string, - rules dockerForwardRules, - proxies dockerProxyEndpoints, - endpoint dto.DockerPortGuardEndpoint, -) (string, string, string) { - if !rules.inspected { - return dockerTrafficPathUnknown, dockerManagementNeedsDiagnosis, dockerReasonNATInspectFailed - } - dnatMatched := dockerDNATRuleMatches(backend, rules.output, endpoint) - if dnatMatched && dockerDNATIngressReachable(backend, rules.output) { - return dockerTrafficPathForward, dockerManagementContainerGuard, "" - } - if !proxies.inspected { - return dockerTrafficPathUnknown, dockerManagementNeedsDiagnosis, dockerReasonProxyInspectFailed - } - if dockerProxyEndpointMatches(proxies.items, endpoint) { - return dockerTrafficPathInput, dockerManagementHostFirewall, "" - } - if dnatMatched { - return dockerTrafficPathUnknown, dockerManagementNeedsDiagnosis, dockerReasonNATChainUnreachable - } - return dockerTrafficPathUnknown, dockerManagementNeedsDiagnosis, dockerReasonNoMatchingPath -} - -func loadDockerDNATRules(backend, family string) dockerForwardRules { - manager := cmd.NewCommandMgr(cmd.WithTimeout(10*time.Second), cmd.WithEnv("LC_ALL=C")) - if backend == constant.FirewallProviderNftables { - tableFamily := "ip" - if family == constant.FirewallFamilyIPv6 { - tableFamily = "ip6" - } - tables, err := manager.RunWithOptionalSudoAndStdout("nft", "list", "tables") - if err != nil { - return dockerForwardRules{} - } - if !strings.Contains(tables, "table "+tableFamily+" docker-bridges") { - return dockerForwardRules{inspected: true} - } - output, err := manager.RunWithOptionalSudoAndStdout("nft", "list", "table", tableFamily, "docker-bridges") - return dockerForwardRules{output: output, inspected: err == nil} - } - commands, err := lifecycle.ResolveIptablesCommands() - if err != nil { - return dockerForwardRules{} - } - executable := commands.IPv4 - if family == constant.FirewallFamilyIPv6 { - executable = commands.IPv6 - } - if executable == "" { - return dockerForwardRules{} - } - output, err := manager.RunWithOptionalSudoAndStdout(executable, "-w", "-t", "nat", "-S") - return dockerForwardRules{output: output, inspected: err == nil} -} - -func loadDockerProxyEndpoints() dockerProxyEndpoints { - manager := cmd.NewCommandMgr(cmd.WithTimeout(10*time.Second), cmd.WithEnv("LC_ALL=C")) - output, err := manager.RunWithStdout("ps", "-ww", "-eo", "args=") - if err != nil { - return dockerProxyEndpoints{} - } - return dockerProxyEndpoints{items: parseDockerProxyEndpoints(output), inspected: true} -} - -func parseDockerProxyEndpoints(output string) []dockerProxyEndpoint { - result := make([]dockerProxyEndpoint, 0) - for _, line := range strings.Split(output, "\n") { - fields := strings.Fields(line) - if len(fields) == 0 || !dockerProxyCommand(fields) { - continue - } - protocol := commandFlagValue(fields, "-proto") - hostIP := commandFlagValue(fields, "-host-ip") - hostPortValue := commandFlagValue(fields, "-host-port") - hostPort, err := strconv.ParseUint(hostPortValue, 10, 16) - if err != nil || (protocol != "tcp" && protocol != "udp") || hostIP == "" { + key := fmt.Sprintf("%s|%s|%d|%s", endpoints[i].Family, endpoints[i].HostIP, endpoints[i].HostPort, endpoints[i].Protocol) + policy, ok := byEndpoint[key] + if !ok { continue } - result = append(result, dockerProxyEndpoint{protocol: protocol, hostIP: canonicalAddress(hostIP), hostPort: uint16(hostPort)}) - } - return result -} - -func dockerProxyCommand(fields []string) bool { - for _, field := range fields { - if filepath.Base(field) == "docker-proxy" { - return true - } - } - return false -} - -func commandFlagValue(fields []string, name string) string { - for i := 0; i < len(fields); i++ { - if fields[i] == name && i+1 < len(fields) { - return fields[i+1] - } - if strings.HasPrefix(fields[i], name+"=") { - return strings.TrimPrefix(fields[i], name+"=") - } - } - return "" -} - -func dockerProxyEndpointMatches(proxies []dockerProxyEndpoint, endpoint dto.DockerPortGuardEndpoint) bool { - for _, proxy := range proxies { - if proxy.protocol == endpoint.Protocol && proxy.hostPort == endpoint.HostPort && hostAddressMatches(proxy.hostIP, endpoint.HostIP, endpoint.Family) { - return true - } - } - return false -} - -func dockerDNATRuleMatches(backend, output string, endpoint dto.DockerPortGuardEndpoint) bool { - if strings.TrimSpace(output) == "" { - return false - } - if backend == constant.FirewallProviderNftables { - return nftDNATRuleMatches(output, endpoint) - } - return iptablesDNATRuleMatches(output, endpoint) -} - -func dockerDNATIngressReachable(backend, output string) bool { - if backend == constant.FirewallProviderNftables { - return strings.Contains(output, "hook prerouting") + sources := []string{} + _ = json.Unmarshal([]byte(policy.Sources), &sources) + endpoints[i].PolicyUUID, endpoints[i].Mode, endpoints[i].Sources = policy.UUID, policy.Mode, sources + endpoints[i].Description = policy.Description + endpoints[i].Effective = endpoints[i].ManagementTarget == dockerManagementContainerGuard && + ((policy.Family == dockerfirewall.FamilyIPv4 && base.IPv4.Effective) || (policy.Family == dockerfirewall.FamilyIPv6 && base.IPv6.Effective)) + delete(byEndpoint, key) } - for _, line := range strings.Split(output, "\n") { - fields := strings.Fields(line) - if len(fields) >= 4 && fields[0] == "-A" && fields[1] == "PREROUTING" && commandFlagValue(fields, "-j") == "DOCKER" { - return true - } + orphanPolicies := make([]dto.DockerPortGuardEndpoint, 0, len(byEndpoint)) + for _, policy := range byEndpoint { + sources := []string{} + _ = json.Unmarshal([]byte(policy.Sources), &sources) + orphanPolicies = append(orphanPolicies, dto.DockerPortGuardEndpoint{ + Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, + PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: sources, Description: policy.Description, + TrafficPath: dockerTrafficPathUnknown, ManagementTarget: dockerManagementNeedsDiagnosis, + ManagementReason: dockerReasonNoMatchingPath, + }) } - return false + return endpoints, orphanPolicies } -func iptablesDNATRuleMatches(output string, endpoint dto.DockerPortGuardEndpoint) bool { - port := strconv.Itoa(int(endpoint.HostPort)) - for _, line := range strings.Split(output, "\n") { - fields := strings.Fields(line) - if commandFlagValue(fields, "-p") != endpoint.Protocol || commandFlagValue(fields, "--dport") != port || commandFlagValue(fields, "-j") != "DNAT" { - continue - } - if destinationAddressMatches(commandFlagValue(fields, "-d"), endpoint) { - return true - } +func (s *DockerPortGuardService) rejectHostInputDockerGuardEndpoints(ctx context.Context, requested []dto.DockerPortGuardEndpointIdentity) error { + if s.client == nil || len(requested) == 0 { + return nil } - return false -} - -func nftDNATRuleMatches(output string, endpoint dto.DockerPortGuardEndpoint) bool { - port := strconv.Itoa(int(endpoint.HostPort)) - for _, line := range strings.Split(output, "\n") { - fields := strings.Fields(strings.NewReplacer("{", " ", "}", " ", ",", " ", ";", " ").Replace(line)) - if !containsToken(fields, "dnat") || !nftProtocolPortMatches(fields, endpoint.Protocol, port) { - continue - } - destination := nftDestinationAddress(fields, endpoint.Family) - if destinationAddressMatches(destination, endpoint) { - return true - } + cli, err := s.client() + if err != nil { + return nil } - return false -} - -func nftProtocolPortMatches(fields []string, protocol, port string) bool { - for i := 0; i+2 < len(fields); i++ { - if fields[i] == protocol && fields[i+1] == "dport" && fields[i+2] == port { - return true - } - if fields[i] == "th" && fields[i+1] == "dport" && fields[i+2] == port && nftMetaProtocolMatches(fields, protocol) { - return true - } + defer cli.Close() + info, err := cli.Info(ctx) + if err != nil { + return nil } - return false -} - -func nftMetaProtocolMatches(fields []string, protocol string) bool { - for i := 0; i+2 < len(fields); i++ { - if fields[i] == "meta" && fields[i+1] == "l4proto" && fields[i+2] == protocol { - return true - } + endpoints, err := discoverDockerEndpoints(ctx, cli, true) + if err != nil { + return nil } - return false -} - -func nftDestinationAddress(fields []string, family string) string { - token := "ip" - if family == constant.FirewallFamilyIPv6 { - token = "ip6" + annotateDockerEndpointManagement(endpoints, dockerFirewallBackend(info)) + targets := make(map[string]string, len(endpoints)) + for _, endpoint := range endpoints { + targets[fmt.Sprintf("%s|%s|%d|%s", endpoint.Family, endpoint.HostIP, endpoint.HostPort, endpoint.Protocol)] = endpoint.ManagementTarget } - for i := 0; i+2 < len(fields); i++ { - if fields[i] == token && fields[i+1] == "daddr" { - return fields[i+2] + for _, endpoint := range requested { + target := targets[fmt.Sprintf("%s|%s|%d|%s", endpoint.Family, endpoint.HostIP, endpoint.HostPort, endpoint.Protocol)] + if target == dockerManagementHostFirewall { + return buserr.WithDetail("ErrInvalidParams", "endpoint traffic is handled by the host input firewall", nil) } - } - return "" -} - -func destinationAddressMatches(ruleAddress string, endpoint dto.DockerPortGuardEndpoint) bool { - ruleAddress = strings.TrimSpace(strings.Split(ruleAddress, "/")[0]) - if isWildcardHostAddress(endpoint.HostIP, endpoint.Family) { - return ruleAddress == "" - } - return ruleAddress == "" || canonicalAddress(ruleAddress) == canonicalAddress(endpoint.HostIP) -} - -func hostAddressMatches(left, right, family string) bool { - if isWildcardHostAddress(left, family) && isWildcardHostAddress(right, family) { - return true - } - return canonicalAddress(left) == canonicalAddress(right) -} - -func isWildcardHostAddress(value, family string) bool { - value = strings.TrimSpace(value) - if family == constant.FirewallFamilyIPv6 { - return value == "" || value == "::" - } - return value == "" || value == "0.0.0.0" -} - -func canonicalAddress(value string) string { - if address, err := netip.ParseAddr(strings.TrimSpace(value)); err == nil { - return address.String() - } - return strings.TrimSpace(value) -} - -func containsToken(fields []string, value string) bool { - for _, field := range fields { - if field == value { - return true + if target == dockerManagementNeedsDiagnosis { + return buserr.WithDetail("ErrInvalidParams", "endpoint traffic management target requires diagnosis", nil) } } - return false + return nil } diff --git a/agent/app/service/firewall_rule_task.go b/agent/app/service/firewall_rule_task.go deleted file mode 100644 index 443c858d05b7..000000000000 --- a/agent/app/service/firewall_rule_task.go +++ /dev/null @@ -1,76 +0,0 @@ -package service - -import ( - "context" - "fmt" - "io" - - "github.com/1Panel-dev/1Panel/agent/app/dto" - "github.com/1Panel-dev/1Panel/agent/app/repo" - "github.com/1Panel-dev/1Panel/agent/app/task" - "github.com/1Panel-dev/1Panel/agent/global" - "github.com/1Panel-dev/1Panel/agent/i18n" -) - -const ( - firewallTaskHost = "FirewallTaskHost" - firewallTaskForwarding = "FirewallTaskForwarding" - firewallTaskDocker = "FirewallTaskDocker" -) - -func firewallTaskName(operation, subsystem, backend string) string { - name := i18n.GetMsgByKey(subsystem) - if backend != "" { - name += " · " + backend - } - key := "FirewallRule" + operation - if operation == task.TaskExec { - key = "FirewallTaskInitialize" - } - return i18n.GetMsgWithMap(key, map[string]interface{}{"name": name}) -} - -func queueFirewallRuleTask(subsystem, operation string, labels []string, apply func(context.Context) error) (dto.FilterChainOperationResponse, error) { - taskItem, err := task.NewTask(firewallTaskName(operation, subsystem, ""), operation, task.TaskScopeFirewall, "", 0) - if err != nil { - return dto.FilterChainOperationResponse{}, err - } - taskItem.AddSubTaskWithOps(taskItem.Name, func(t *task.Task) error { - t.Logf("rules=%d", len(labels)) - err := t.TaskCtx.Err() - if err == nil { - err = apply(t.TaskCtx) - } - succeeded, failed := 0, 0 - for _, label := range labels { - if err != nil { - failed++ - t.LogFailedWithErr(label, err) - } else { - succeeded++ - t.LogSuccess(label) - } - } - t.Log(i18n.GetMsgWithMap("FirewallRuleOperationResult", map[string]interface{}{ - "succeeded": succeeded, "failed": failed, - })) - return err - }, nil, 0, 0) - if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { - taskItem.LogFailedWithErr(taskItem.Name, err) - closeUnstartedFirewallTask(taskItem) - return dto.FilterChainOperationResponse{}, fmt.Errorf("save firewall rule task: %w", err) - } - go func() { _ = taskItem.Execute() }() - return dto.FilterChainOperationResponse{TaskID: taskItem.TaskID, Queued: true}, nil -} - -func closeUnstartedFirewallTask(t *task.Task) { - if cancel, ok := global.LoadTaskCancel(t.TaskID); ok { - cancel() - } - global.RemoveTaskCancel(t.TaskID) - if closer, ok := t.Logger.Out.(io.Closer); ok { - _ = closer.Close() - } -} diff --git a/agent/app/service/firewall_selection.go b/agent/app/service/firewall_selection.go deleted file mode 100644 index dceb8a347537..000000000000 --- a/agent/app/service/firewall_selection.go +++ /dev/null @@ -1,65 +0,0 @@ -package service - -import ( - "strings" - - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/global" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" -) - -func selectedDockerFirewallBackend(fallback string) string { - selected := configuredDockerFirewallBackend() - if selected == constant.FirewallProviderIptables || selected == constant.FirewallProviderNftables { - return selected - } - fallback = strings.ToLower(strings.TrimSpace(fallback)) - if fallback == constant.FirewallProviderNftables { - return fallback - } - return constant.FirewallProviderIptables -} - -func configuredDockerFirewallBackend() string { - if global.DB == nil { - return "" - } - selected, _ := settingRepo.GetValueByKey(constant.FirewallDockerBackendKey) - selected = strings.ToLower(strings.TrimSpace(selected)) - if selected == constant.FirewallProviderIptables || selected == constant.FirewallProviderNftables { - return selected - } - return "" -} - -func selectedSystemFirewallClient() (lifecycle.Client, error) { - if provider := configuredSystemFirewallBackend(); provider != "" { - return lifecycle.NewClientFor(provider) - } - client, err := lifecycle.NewClient() - if err != nil { - return nil, err - } - _ = settingRepo.UpdateOrCreate(constant.FirewallSystemBackendKey, client.Name()) - return client, nil -} - -func configuredSystemFirewallBackend() string { - if global.DB == nil { - return "" - } - provider, _ := settingRepo.GetValueByKey(constant.FirewallSystemBackendKey) - return strings.TrimSpace(provider) -} - -func NewSelectedSystemFirewallClient() (lifecycle.Client, error) { - return selectedSystemFirewallClient() -} - -func selectedSystemFirewallProvider() (string, error) { - client, err := selectedSystemFirewallClient() - if err != nil { - return "", err - } - return client.Name(), nil -} diff --git a/agent/app/service/firewall_setting.go b/agent/app/service/firewall_setting.go index e8a5bcb5d072..dc1fbd1b89ba 100644 --- a/agent/app/service/firewall_setting.go +++ b/agent/app/service/firewall_setting.go @@ -5,23 +5,20 @@ import ( "encoding/json" "errors" "fmt" - "os" - "reflect" "slices" + "strings" "sync" "github.com/1Panel-dev/1Panel/agent/app/dto" "github.com/1Panel-dev/1Panel/agent/app/model" + "github.com/1Panel-dev/1Panel/agent/buserr" "github.com/1Panel-dev/1Panel/agent/constant" "github.com/1Panel-dev/1Panel/agent/global" "github.com/1Panel-dev/1Panel/agent/utils/cmd" "github.com/1Panel-dev/1Panel/agent/utils/firewall" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" + dockerfirewall "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - 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" "gorm.io/gorm" ) @@ -37,35 +34,61 @@ type FirewallSettingService struct{} var firewallWhitelistMu sync.Mutex -var ErrFirewallBackendCleanupRequired = errors.New("firewall backend cleanup required") - -func firewallBackendCleanupRequired(current, target string) error { - return fmt.Errorf( - "%w: current backend %s still contains 1Panel runtime rules; clean it up before switching to %s", - ErrFirewallBackendCleanupRequired, - current, - target, - ) -} - -func NewIFirewallSettingService() IFirewallSettingService { - return &FirewallSettingService{} -} - func (s *FirewallSettingService) CreatePortWhitelist(ctx context.Context, request dto.FirewallPortWhitelistCreate) error { - return savePortWhitelist(ctx, func(current []firewall.PortWhitelist) ([]firewall.PortWhitelist, error) { - return append(current, request.Rule), nil + firewallWhitelistMu.Lock() + defer firewallWhitelistMu.Unlock() + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + return global.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + current, err := loadPortWhitelistSetting(tx) + if err != nil { + return err + } + current = append(current, request.Rule) + current, err = firewall.ValidatePortWhitelist(current) + if err != nil { + return err + } + value, err := json.Marshal(current) + if err != nil { + return err + } + err = tx.Where("key = ?", constant.FirewallPortWhiteList).Assign(map[string]interface{}{"value": string(value)}).FirstOrCreate(&model.Setting{Key: constant.FirewallPortWhiteList}).Error + if err != nil { + return err + } + return nil }) } func (s *FirewallSettingService) UpdatePortWhitelist(ctx context.Context, request dto.FirewallPortWhitelistUpdate) error { - return savePortWhitelist(ctx, func(current []firewall.PortWhitelist) ([]firewall.PortWhitelist, error) { + firewallWhitelistMu.Lock() + defer firewallWhitelistMu.Unlock() + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + return global.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + current, err := loadPortWhitelistSetting(tx) + if err != nil { + return err + } index, err := findPortWhitelistRule(current, request.OldRule) if err != nil { - return nil, err + return err } current[index] = request.Rule - return current, nil + current, err = firewall.ValidatePortWhitelist(current) + if err != nil { + return err + } + value, err := json.Marshal(current) + if err != nil { + return err + } + err = tx.Where("key = ?", constant.FirewallPortWhiteList).Assign(map[string]interface{}{"value": string(value)}).FirstOrCreate(&model.Setting{Key: constant.FirewallPortWhiteList}).Error + if err != nil { + return err + } + return nil }) } @@ -73,156 +96,33 @@ func (s *FirewallSettingService) DeletePortWhitelist(ctx context.Context, reques if request.Rule == nil { return fmt.Errorf("select one firewall port whitelist rule to delete") } - return savePortWhitelist(ctx, func(current []firewall.PortWhitelist) ([]firewall.PortWhitelist, error) { - index, err := findPortWhitelistRule(current, *request.Rule) - if err != nil { - return nil, err - } - return slices.Delete(current, index, index+1), nil - }) -} - -func findPortWhitelistRule(rules []firewall.PortWhitelist, target firewall.PortWhitelist) (int, error) { - index := slices.IndexFunc(rules, func(rule firewall.PortWhitelist) bool { - return samePortWhitelistRule(rule, target) - }) - if index < 0 { - return -1, fmt.Errorf("firewall port whitelist rule has changed or no longer exists; refresh and retry") - } - return index, nil -} - -func samePortWhitelistRule(left, right firewall.PortWhitelist) bool { - if reflect.DeepEqual(left, right) { - return true - } - normalizedLeft, err := firewall.ValidatePortWhitelist([]firewall.PortWhitelist{left}) - if err != nil { - return false - } - normalizedRight, err := firewall.ValidatePortWhitelist([]firewall.PortWhitelist{right}) - if err != nil { - return false - } - slices.Sort(normalizedLeft[0].Sources) - slices.Sort(normalizedRight[0].Sources) - return reflect.DeepEqual(normalizedLeft[0], normalizedRight[0]) -} - -func savePortWhitelist(ctx context.Context, change func([]firewall.PortWhitelist) ([]firewall.PortWhitelist, error)) error { firewallWhitelistMu.Lock() defer firewallWhitelistMu.Unlock() firewallRuleMutationMu.Lock() defer firewallRuleMutationMu.Unlock() - defer filterruntime.InvalidateInventory() return global.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { current, err := loadPortWhitelistSetting(tx) if err != nil { return err } - desired, err := change(current) + index, err := findPortWhitelistRule(current, *request.Rule) if err != nil { return err } - desired, err = firewall.ValidatePortWhitelist(desired) + current = slices.Delete(current, index, index+1) + current, err = firewall.ValidatePortWhitelist(current) if err != nil { return err } - value, err := json.Marshal(desired) + value, err := json.Marshal(current) if err != nil { return err } - return tx.Where("key = ?", constant.FirewallPortWhiteList).Assign(map[string]interface{}{"value": string(value)}). - FirstOrCreate(&model.Setting{Key: constant.FirewallPortWhiteList}).Error - }) -} - -func checkFirewallRuleWhitelistProtection(provider filter.Provider, record model.FirewallRule) error { - ports, err := loadFirewallPortWhiteList() - if err != nil { - return err - } - rules, err := record.RulesForProvider(provider) - if err != nil { - return err - } - for _, rule := range rules { - if filter.RuleMatchesPortWhitelist(rule, ports) { - return filter.ErrProtectedRule - } - } - return nil -} - -func loadPortWhitelistSetting(db *gorm.DB) ([]firewall.PortWhitelist, error) { - var setting model.Setting - if err := db.Where("key = ?", constant.FirewallPortWhiteList).First(&setting).Error; errors.Is(err, gorm.ErrRecordNotFound) { - setting.Value = constant.FirewallPortWhiteListValue - } else if err != nil { - return nil, err - } - var rules []firewall.PortWhitelist - err := json.Unmarshal([]byte(setting.Value), &rules) - return rules, err -} - -func loadSSHWhitelistPortFrom(path string) (string, error) { - directives, _, err := parseSSHConfigTree(path) - if errors.Is(err, os.ErrNotExist) { - return defaultSSHPort, nil - } - if err != nil { - return "", err - } - return loadSSHPortValues(directives)[0], nil -} - -func customWhitelist(entries []firewall.PortWhitelist) []firewall.PortWhitelist { - result := make([]firewall.PortWhitelist, 0, len(entries)) - for _, entry := range entries { - if entry.Type == "" { - result = append(result, entry) - } - } - return result -} - -func InitializeFirewallWhitelistPorts(entries []firewall.PortWhitelist) ([]firewall.PortWhitelist, error) { - entries = slices.Clone(entries) - var sshPort string - for i := range entries { - entry := &entries[i] - if entry.Type == "" || entry.Port != "" { - continue - } - switch entry.Type { - case firewall.PortWhitelistTypePanel: - entry.Port = LoadPanelPort() - case firewall.PortWhitelistTypeSSH: - if sshPort == "" { - var err error - sshPort, err = loadSSHWhitelistPortFrom(sshPath) - if err != nil { - return nil, err - } - } - entry.Port = sshPort - } - } - return firewall.ValidatePortWhitelist(entries) -} - -func updateSystemAccessPortWhitelist(ctx context.Context, serviceType string, ports []string) error { - return savePortWhitelist(ctx, func(entries []firewall.PortWhitelist) ([]firewall.PortWhitelist, error) { - for i := range entries { - if entries[i].Type == serviceType { - if len(ports) == 0 { - return nil, fmt.Errorf("firewall whitelist %s requires a port", serviceType) - } - entries[i].Port = ports[0] - } + err = tx.Where("key = ?", constant.FirewallPortWhiteList).Assign(map[string]interface{}{"value": string(value)}).FirstOrCreate(&model.Setting{Key: constant.FirewallPortWhiteList}).Error + if err != nil { + return err } - return entries, nil + return nil }) } @@ -233,9 +133,10 @@ func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings for _, name := range lifecycle.InstalledProviders() { installed[name] = true } - result.System.Selected = configuredSystemFirewallBackend() + systemBackend, _ := settingRepo.GetValueByKey(constant.FirewallSystemBackendKey) + result.System.Selected = strings.TrimSpace(systemBackend) if result.System.Selected == "" { - if client, err := lifecycle.NewClient(); err == nil { + if client, err := lifecycle.NewClient(""); err == nil { result.System.Selected = client.Name() } } @@ -248,16 +149,16 @@ func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings } { option := dto.FirewallBackendOption{Name: name, Installed: installed[name], Supported: true} if option.Installed && name == result.System.Selected { - client, err := lifecycle.NewClientFor(name) + client, err := lifecycle.NewClient(name) if err != nil { option.Message = err.Error() - } else if supportsManagedFilterChains(name) { - option.Initialized, option.Bound, err = loadFirewallInitStatus(name, "base") + } else if name == constant.FirewallProviderIptables || name == constant.FirewallProviderNftables { + overview, err := loadSystemFirewallOverview(name, "base") if err != nil { option.Message = err.Error() } - option.IPv4 = loadSystemFirewallFamilyInfo(name, constant.FirewallFamilyIPv4) - option.IPv6 = loadSystemFirewallFamilyInfo(name, constant.FirewallFamilyIPv6) + option.Initialized, option.Bound = overview.IsInit, overview.IsBind + option.IPv4, option.IPv6 = overview.IPv4, overview.IPv6 } else if option.Active, err = client.Status(); err != nil { option.Message = err.Error() } @@ -270,32 +171,29 @@ func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings result.System.Options = append(result.System.Options, option) } - result.Forwarding.Selected = configuredForwardingBackend() + forwardingBackend, _ := settingRepo.GetValueByKey(constant.FirewallForwardingBackendKey) + result.Forwarding.Selected = strings.TrimSpace(forwardingBackend) + if result.Forwarding.Selected == "" { + result.Forwarding.Selected = constant.FirewallProviderIptables + } result.Forwarding.Current = result.Forwarding.Selected for _, name := range []string{constant.FirewallProviderIptables, constant.FirewallProviderNftables} { option := dto.FirewallBackendOption{Name: name, Installed: installed[name], Supported: true} if option.Installed && name == result.Forwarding.Selected { - manager, err := newForwardingManagerFor(name) + manager, err := newForwardingAdapterFor(name) if err != nil { option.Message = err.Error() - } else if status, err := manager.Status(); err != nil { - option.Message = err.Error() } else { - option.Initialized, option.Bound = status.IsInit, status.IsBind - ipv4Init, ipv4Bound, ipv4Err := manager.FamilyStatus(constant.FirewallFamilyIPv4) - ipv6Init, ipv6Bound, ipv6Err := manager.FamilyStatus(constant.FirewallFamilyIPv6) - option.IPv4 = dto.FirewallBackendFamilyStatus{ - Available: ipv4Err == nil, Initialized: ipv4Init, Bound: ipv4Bound, - } - option.IPv6 = dto.FirewallBackendFamilyStatus{ - Available: ipv6Err == nil, Initialized: ipv6Init, Bound: ipv6Bound, + status, statusErr := loadForwardingFirewallOverview(manager) + option.IPv4, option.IPv6 = status.IPv4, status.IPv6 + if statusErr != nil { + option.Message = statusErr.Error() + } else { + option.Initialized, option.Bound = status.IsInit, status.IsBind } - if name == constant.FirewallProviderIptables { - if commands, commandErr := lifecycle.ResolveIptablesCommands(); commandErr == nil { - option.IPv6.Available = option.IPv6.Available && commands.IPv6Available() - if !commands.IPv6Available() { - option.IPv6.Reason = docker_guard.ReasonCommandMissing - } + if name == constant.FirewallProviderIptables && !option.IPv6.Available { + if commands, err := lifecycle.ResolveIptablesCommands(); err == nil && !commands.IPv6Available() { + option.IPv6.Reason = dockerfirewall.ReasonCommandMissing } } } @@ -313,7 +211,11 @@ func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings if dockerInstalled { dockerVersion = loadDockerEngineVersion(ctx) } - result.Docker.Selected = configuredDockerFirewallBackend() + dockerBackend, _ := settingRepo.GetValueByKey(constant.FirewallDockerBackendKey) + dockerBackend = strings.ToLower(strings.TrimSpace(dockerBackend)) + if dockerBackend == constant.FirewallProviderIptables || dockerBackend == constant.FirewallProviderNftables { + result.Docker.Selected = dockerBackend + } result.Docker.Current = result.Docker.Selected for _, name := range []string{constant.FirewallProviderIptables, constant.FirewallProviderNftables} { option := dto.FirewallBackendOption{ @@ -326,14 +228,14 @@ func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings option.Active = false } if option.Active { - guard := docker_guard.NewRuntime(name) - ipv4, ipv6 := guard.Status(docker_guard.FamilyIPv4), guard.Status(docker_guard.FamilyIPv6) + guard := newDockerFirewallRuntime(name) + ipv4, ipv6 := guard.Status(dockerfirewall.FamilyIPv4), guard.Status(dockerfirewall.FamilyIPv6) option.Initialized = ipv4.Initialized || ipv6.Initialized option.Bound = ipv4.Bound || ipv6.Bound option.IPv4.Initialized, option.IPv4.Bound = ipv4.Initialized, ipv4.Bound option.IPv6.Initialized, option.IPv6.Bound = ipv6.Initialized, ipv6.Bound - option.IPv4.Available = ipv4.Reason != docker_guard.ReasonCommandMissing - option.IPv6.Available = ipv6.Reason != docker_guard.ReasonCommandMissing + option.IPv4.Available = ipv4.Reason != dockerfirewall.ReasonCommandMissing + option.IPv6.Available = ipv6.Reason != dockerfirewall.ReasonCommandMissing option.IPv4.Reason, option.IPv6.Reason = ipv4.Reason, ipv6.Reason } result.Docker.Options = append(result.Docker.Options, option) @@ -353,32 +255,6 @@ func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings return result, err } -func loadSystemFirewallFamilyStatus(provider, family string) (bool, bool, error) { - switch provider { - case constant.FirewallProviderIptables: - return iptables_helper.LoadFamilyInitStatus(family, "base") - case constant.FirewallProviderNftables: - return nftables_helper.LoadFamilyInitStatus(filter.Family(family), "base") - default: - return false, false, fmt.Errorf("unsupported firewall provider %q", provider) - } -} - -func loadSystemFirewallFamilyInfo(provider, family string) dto.FirewallBackendFamilyStatus { - if provider == constant.FirewallProviderIptables && family == constant.FirewallFamilyIPv6 { - commands, err := lifecycle.ResolveIptablesCommands() - if err != nil || !commands.IPv6Available() { - return dto.FirewallBackendFamilyStatus{Reason: docker_guard.ReasonCommandMissing} - } - } - initialized, bound, err := loadSystemFirewallFamilyStatus(provider, family) - return dto.FirewallBackendFamilyStatus{ - Available: err == nil, - Initialized: initialized, - Bound: bound, - } -} - func (s *FirewallSettingService) Operate(ctx context.Context, request dto.FirewallBackendOperation) error { if err := lockFirewallLifecycleIdle(); err != nil { return err @@ -387,7 +263,7 @@ func (s *FirewallSettingService) Operate(ctx context.Context, request dto.Firewa if request.Subsystem != "system" && request.Backend != constant.FirewallProviderIptables && request.Backend != constant.FirewallProviderNftables { return fmt.Errorf("%s only supports iptables or nftables", request.Subsystem) } - if request.Subsystem == "system" && !supportsManagedFilterChains(request.Backend) && request.Operation != "select" { + if request.Subsystem == "system" && (request.Backend != constant.FirewallProviderIptables && request.Backend != constant.FirewallProviderNftables) && request.Operation != "select" { return fmt.Errorf("%s does not support initialization or cleanup", request.Backend) } switch request.Subsystem { @@ -411,64 +287,14 @@ func (s *FirewallSettingService) Operate(ctx context.Context, request dto.Firewa } } -func (s *FirewallSettingService) operateDocker(ctx context.Context, request dto.FirewallBackendOperation) error { - guard := docker_guard.NewRuntime(request.Backend) - if request.Operation == "cleanup" { - if err := guard.Cleanup(); err != nil { - return err - } - return settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusDisable) - } - previous, _ := settingRepo.GetValueByKey(constant.FirewallDockerBackendKey) - if request.Operation == "select" { - current := previous - if current == "" { - current = alternateDirectBackend(request.Backend) - } - initialized, err := dockerGuardBackendInitialized(current) - if err != nil { - return err - } - if current != request.Backend && initialized { - return firewallBackendCleanupRequired(current, request.Backend) - } - } - if err := settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, request.Backend); err != nil { - return err - } - if request.Operation == "select" { - if err := (&DockerService{}).UpdateFirewallBackend(request.Backend); err != nil { - _ = settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, previous) - return err - } - } - if request.Operation == "initialize" { - if err := newDockerPortGuardService().Operate(ctx, dto.DockerPortGuardOperation{Operation: "initialize"}); err != nil { - _ = settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, previous) - return err - } - } - return nil -} - -func dockerGuardBackendInitialized(backend string) (bool, error) { - guard := docker_guard.NewRuntime(backend) - for _, family := range []string{docker_guard.FamilyIPv4, docker_guard.FamilyIPv6} { - initialized, err := guard.Initialized(family) - if err != nil { - return false, err - } - if initialized { - return true, nil - } - } - return false, nil +func NewIFirewallSettingService() IFirewallSettingService { + return &FirewallSettingService{} } func (s *FirewallSettingService) operateSystem(request dto.FirewallBackendOperation) error { firewallRuleMutationMu.Lock() defer firewallRuleMutationMu.Unlock() - if _, err := lifecycle.NewClientFor(request.Backend); err != nil { + if _, err := lifecycle.NewClient(request.Backend); err != nil { return err } if request.Operation == "cleanup" { @@ -476,7 +302,7 @@ func (s *FirewallSettingService) operateSystem(request dto.FirewallBackendOperat } previous, _ := settingRepo.GetValueByKey(constant.FirewallSystemBackendKey) if previous == "" { - if client, err := lifecycle.NewClient(); err == nil { + if client, err := lifecycle.NewClient(""); err == nil { previous = client.Name() } } @@ -486,7 +312,7 @@ func (s *FirewallSettingService) operateSystem(request dto.FirewallBackendOperat return err } if initialized { - return firewallBackendCleanupRequired(previous, request.Backend) + return buserr.WithMap("ErrFirewallBackendCleanupRequired", map[string]interface{}{"current": previous, "target": request.Backend}, nil) } } if err := settingRepo.UpdateOrCreate(constant.FirewallSystemBackendKey, request.Backend); err != nil { @@ -512,21 +338,14 @@ func (s *FirewallSettingService) operateSystem(request dto.FirewallBackendOperat } func systemFirewallBackendInitialized(backend string) (bool, error) { - return systemFirewallBackendInitializedWithClientFactory(backend, lifecycle.NewClientFor) -} - -func systemFirewallBackendInitializedWithClientFactory( - backend string, - newClient func(string) (lifecycle.Client, error), -) (bool, error) { - client, err := newClient(backend) + client, err := lifecycle.NewClient(backend) if err != nil { if errors.Is(err, lifecycle.ErrNotInstalled) { return false, nil } return false, err } - if supportsManagedFilterChains(backend) { + if backend == constant.FirewallProviderIptables || backend == constant.FirewallProviderNftables { for _, family := range []string{constant.FirewallFamilyIPv4, constant.FirewallFamilyIPv6} { initialized, _, err := loadSystemFirewallFamilyStatus(backend, family) if family == constant.FirewallFamilyIPv6 && errors.Is(err, filter.ErrFamilyUnavailable) { @@ -544,30 +363,8 @@ func systemFirewallBackendInitializedWithClientFactory( return client.Status() } -func cleanupSystemBackend(backend string) error { - switch backend { - case constant.FirewallProviderIptables: - return newIptablesHelperManager().Cleanup() - case constant.FirewallProviderNftables: - return newNftablesHelperManager().Cleanup() - default: - return fmt.Errorf("cleanup is only available for 1Panel-owned iptables and nftables resources") - } -} - -func cleanupInactiveSystemBackend(backend string) error { - switch backend { - case constant.FirewallProviderIptables: - return (&iptables_helper.Manager{}).Cleanup() - case constant.FirewallProviderNftables: - return (&nftables_helper.Manager{}).Cleanup() - default: - return fmt.Errorf("cleanup is only available for 1Panel-owned iptables and nftables resources") - } -} - func (s *FirewallSettingService) operateForwarding(request dto.FirewallBackendOperation) error { - manager, err := newForwardingManagerFor(request.Backend) + manager, err := newForwardingAdapterFor(request.Backend) if err != nil { return err } @@ -585,7 +382,7 @@ func (s *FirewallSettingService) operateForwarding(request dto.FirewallBackendOp if request.Operation == "select" { current := previous if current == "" { - detected, err := newForwardingManager() + detected, err := newForwardingAdapter() if err != nil { return err } @@ -596,7 +393,7 @@ func (s *FirewallSettingService) operateForwarding(request dto.FirewallBackendOp return err } if current != request.Backend && initialized { - return firewallBackendCleanupRequired(current, request.Backend) + return buserr.WithMap("ErrFirewallBackendCleanupRequired", map[string]interface{}{"current": current, "target": request.Backend}, nil) } } if err := settingRepo.UpdateOrCreate(constant.FirewallForwardingBackendKey, request.Backend); err != nil { @@ -610,7 +407,7 @@ func (s *FirewallSettingService) operateForwarding(request dto.FirewallBackendOp } func forwardingBackendInitialized(backend string) (bool, error) { - manager, err := newForwardingManagerFor(backend) + manager, err := newForwardingAdapterFor(backend) if err != nil { if errors.Is(err, lifecycle.ErrNotInstalled) { return false, nil @@ -629,9 +426,59 @@ func forwardingBackendInitialized(backend string) (bool, error) { return false, nil } -func alternateDirectBackend(backend string) string { - if backend == constant.FirewallProviderNftables { - return constant.FirewallProviderIptables +func (s *FirewallSettingService) operateDocker(ctx context.Context, request dto.FirewallBackendOperation) error { + guard := newDockerFirewallRuntime(request.Backend) + if request.Operation == "cleanup" { + if err := guard.Cleanup(); err != nil { + return err + } + return settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusDisable) } - return constant.FirewallProviderNftables + previous, _ := settingRepo.GetValueByKey(constant.FirewallDockerBackendKey) + if request.Operation == "select" { + current := previous + if current == "" { + current = constant.FirewallProviderNftables + if request.Backend == constant.FirewallProviderNftables { + current = constant.FirewallProviderIptables + } + } + initialized, err := dockerGuardBackendInitialized(current) + if err != nil { + return err + } + if current != request.Backend && initialized { + return buserr.WithMap("ErrFirewallBackendCleanupRequired", map[string]interface{}{"current": current, "target": request.Backend}, nil) + } + } + if err := settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, request.Backend); err != nil { + return err + } + if request.Operation == "select" { + if err := (&DockerService{}).UpdateFirewallBackend(request.Backend); err != nil { + _ = settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, previous) + return err + } + } + if request.Operation == "initialize" { + if err := newDockerPortGuardService().Operate(ctx, dto.DockerPortGuardOperation{Operation: "initialize"}); err != nil { + _ = settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, previous) + return err + } + } + return nil +} + +func dockerGuardBackendInitialized(backend string) (bool, error) { + guard := newDockerFirewallRuntime(backend) + for _, family := range []string{dockerfirewall.FamilyIPv4, dockerfirewall.FamilyIPv6} { + initialized, err := guard.Initialized(family) + if err != nil { + return false, err + } + if initialized { + return true, nil + } + } + return false, nil } diff --git a/agent/app/service/firewall_sync.go b/agent/app/service/firewall_sync.go index 193c2c317782..837ec795f33e 100644 --- a/agent/app/service/firewall_sync.go +++ b/agent/app/service/firewall_sync.go @@ -2,11 +2,12 @@ package service import ( "context" + "encoding/json" "errors" "fmt" - "sort" "strings" "sync" + "time" "github.com/1Panel-dev/1Panel/agent/app/dto" "github.com/1Panel-dev/1Panel/agent/app/model" @@ -16,13 +17,11 @@ import ( "github.com/1Panel-dev/1Panel/agent/global" "github.com/1Panel-dev/1Panel/agent/i18n" "github.com/1Panel-dev/1Panel/agent/utils/firewall" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" + dockerfirewall "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" - filterruntime "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/runtime" "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" firewallsync "github.com/1Panel-dev/1Panel/agent/utils/firewall/sync" "github.com/google/uuid" - "gorm.io/gorm" ) var ( @@ -35,657 +34,10 @@ type firewallDatabaseSyncAdapter interface { syncRules(context.Context, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error) } -type firewallSyncRule struct { - dto.FirewallRuleSyncItem - desired filter.DesiredRule - observed *filter.ObservedRule - done bool -} - -func firewallSyncSubsystem(value string) string { - if value = strings.TrimSpace(value); value == "" { - return "system" - } - return value -} - -func (s *FirewallService) PreviewRuleSync(ctx context.Context, clientIP string, request dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncPreview, error) { - switch firewallSyncSubsystem(request.Subsystem) { - case "forwarding": - return s.forwardingRuleSyncService().previewRuleSync(ctx, request) - case "docker": - return s.dockerRuleSyncService().previewRuleSync(ctx, request) - } - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - _, rules, _, err := s.loadFirewallSyncRules(ctx, request) - preview := dto.FirewallRuleSyncPreview{Subsystem: "system", TargetProvider: request.TargetProvider, Items: make([]dto.FirewallRuleSyncItem, 0, len(rules))} - for _, rule := range rules { - preview.Add(rule.FirewallRuleSyncItem) - } - return preview, err -} - -func (s *FirewallService) loadFirewallSyncRules(ctx context.Context, request dto.FirewallRuleSyncRequest) (*filterruntime.Engine, []*firewallSyncRule, []filter.Snapshot, error) { - if request.SourceProvider != "" || request.ResetSource { - return nil, nil, nil, fmt.Errorf("%w: system synchronization reads rules from the database", filter.ErrInvalidRule) - } - if err := s.checkSelectedProvider(ctx, request.TargetProvider); err != nil { - return nil, nil, nil, err - } - runtime, err := s.adapters.Resolve(request.TargetProvider) - if err != nil { - return nil, nil, nil, err - } - stored, err := s.rules.List(ctx) - if err != nil { - return nil, nil, nil, err - } - model.SortFirewallRules(stored, request.TargetProvider) - ports, err := loadFirewallPortWhiteList() - if err != nil { - return nil, nil, nil, err - } - required, err := firewall.RequiredPortWhitelist(ports) - if err != nil { - return nil, nil, nil, err - } - whitelistDesired := whitelistRules(request.TargetProvider, firewall.ExpandPortWhitelist(customWhitelist(ports)), firewall.ExpandPortWhitelist(required)) - rules := make([]*firewallSyncRule, 0, len(stored)) - preservedMarkers := make(map[string]bool) - compileFailed := false - for _, record := range stored { - desired, preserved, err := s.compileRestorableFirewallRules(ctx, record, request.TargetProvider) - if err != nil { - compileFailed = true - rules = append(rules, &firewallSyncRule{FirewallRuleSyncItem: dto.FirewallRuleSyncItem{SourceUUID: record.UUID, Status: firewallsync.StatusBlocked, Reason: err.Error()}}) - continue - } - for _, rule := range preserved { - preservedMarkers[rule.Marker] = true - } - for _, native := range desired { - rule := native.Rule - native.Protected = filter.RuleMatchesPortWhitelist(rule, ports) - rules = append(rules, &firewallSyncRule{desired: native, FirewallRuleSyncItem: dto.FirewallRuleSyncItem{SourceUUID: record.UUID, Rule: &rule}}) - } - } - for _, candidate := range whitelistDesired { - prepared, err := s.prepareCreate(ctx, request.TargetProvider, dto.FirewallRuleCreateItem{Rule: candidate}) - if err != nil { - return nil, nil, nil, err - } - rule := prepared.request.Rule - duplicate := false - for _, existing := range rules { - if existing.Rule != nil { - if same, err := filter.SameRuleContent(*existing.Rule, rule); err == nil && same { - duplicate = true - break - } - } - } - if duplicate { - continue - } - key, err := filter.RuleKey(rule) - if err != nil { - return nil, nil, nil, err - } - rule.UUID = uuid.NewSHA1(uuid.NameSpaceOID, []byte(key)).String() - rules = append(rules, &firewallSyncRule{ - desired: filter.DesiredRule{UUID: rule.UUID, Rule: rule, RuleKey: key, Origin: filter.RuleOriginCreated, Protected: true}, - FirewallRuleSyncItem: dto.FirewallRuleSyncItem{SourceUUID: rule.UUID, Rule: &rule}, - }) - } - snapshots := make([]filter.Snapshot, 0) - for _, scope := range filter.ManagedInputScopes(request.TargetProvider) { - desired := make([]filter.DesiredRule, 0) - scoped := make([]*firewallSyncRule, 0) - byUUID := make(map[string]*firewallSyncRule) - for _, rule := range rules { - if rule.Rule != nil && rule.Rule.Scope.Key() == scope.Key() { - scoped = append(scoped, rule) - desired = append(desired, rule.desired) - byUUID[rule.desired.Rule.UUID] = rule - } - } - if compileFailed && len(desired) == 0 { - continue - } - snapshot, err := runtime.ObserveMutation(ctx, scope) - if errors.Is(err, filter.ErrFamilyUnavailable) { - for _, rule := range scoped { - rule.Status, rule.Reason = firewallsync.StatusBlocked, err.Error() - } - continue - } - if err != nil { - return nil, nil, nil, err - } - snapshots = append(snapshots, snapshot) - inventory, err := filter.MergeInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desired}) - if err != nil { - return nil, nil, nil, err - } - matched := make(map[string]filter.InventoryItem) - for _, item := range inventory { - if item.Desired == nil { - if compileFailed || item.Observed == nil || !strings.HasPrefix(item.Observed.Marker, "1panel-rule:") || preservedMarkers[item.Observed.Marker] { - continue - } - observed := item.Observed - if observed.Rule.Scope.Chain == filter.BasicBeforeChain { - continue - } - rule := &firewallSyncRule{observed: observed, FirewallRuleSyncItem: dto.FirewallRuleSyncItem{ - SourceUUID: strings.TrimPrefix(observed.Marker, "1panel-rule:"), Rule: &observed.Rule, Status: firewallsync.StatusRemove, - ReasonCode: firewallsync.ReasonManagedOnlyInTarget, Reason: firewallsync.ReasonMessage(firewallsync.ReasonManagedOnlyInTarget), - }} - if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque { - rule.Status, rule.ReasonCode = firewallsync.StatusBlocked, firewallsync.ReasonUnsafeRemoval - rule.Reason = firewallsync.ReasonMessage(rule.ReasonCode) - } - rules = append(rules, rule) - continue - } - rule := byUUID[item.Desired.Rule.UUID] - matched[item.Desired.Rule.UUID] = item - rule.observed = item.Observed - divergent := item.Observed != nil && item.Observed.Persistence != "" && item.Observed.Persistence != filter.PersistenceStatusConverged - switch { - case item.Match == filter.InventoryMatchExact && !divergent: - rule.Status, rule.Reason = firewallsync.StatusExisting, "rule already matches database policy" - case item.Match == filter.InventoryMatchMissing || item.Match == filter.InventoryMatchChanged || item.Match == filter.InventoryMatchExact: - rule.Status, rule.Reason = firewallsync.StatusReady, "target rule differs from database policy" - if item.Observed != nil && item.Observed.Protected { - rule.Status, rule.Reason = firewallsync.StatusBlocked, filter.ErrProtectedRule.Error() - } - default: - rule.Status, rule.Reason = firewallsync.StatusBlocked, fmt.Sprintf("target rule cannot be synchronized: %s", item.Match) - } - } - ordered := make([]filter.InventoryItem, 0, len(scoped)) - for _, rule := range scoped { - ordered = append(ordered, matched[rule.desired.Rule.UUID]) - } - drifted := firewallsync.RuleOrder(snapshot, ordered) - for _, rule := range scoped { - if rule.desired.Protected || !drifted[rule.desired.Marker] { - continue - } - if rule.Status == firewallsync.StatusExisting { - rule.Status, rule.Reason = firewallsync.StatusReady, "managed rule order differs from database sequence" - } - } - } - return runtime, rules, snapshots, nil -} - -func (s *FirewallService) syncRules(ctx context.Context, _ string, request dto.FirewallRuleSyncRequest, t *task.Task) (result dto.FirewallRuleSyncResult, err error) { - if request.SourceProvider != "" || request.ResetSource { - return result, fmt.Errorf("%w: system synchronization reads rules from the database", filter.ErrInvalidRule) - } - if err := s.checkSelectedProvider(ctx, request.TargetProvider); err != nil { - return result, err - } - whitelistErr := s.SyncPortWhitelist(ctx) - if t != nil { - t.LogWithStatus(i18n.GetMsgByKey("FirewallSyncWhitelistStep"), whitelistErr) - } - if whitelistErr != nil { - return result, whitelistErr - } - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - result = dto.FirewallRuleSyncResult{Subsystem: "system", TargetProvider: request.TargetProvider} - if err := ctx.Err(); err != nil { - return result, err - } - runtime, rules, snapshots, err := s.loadFirewallSyncRules(ctx, request) - if err != nil { - return result, err - } - - created, removed, unexecuted := 0, 0, 0 - stopped := make(map[string]error) - failedRemovals := make(map[string]error) - var firewalldFinalSnapshot *filter.Snapshot - record := func(operation string, rule *firewallSyncRule, cause error, skipped bool) { - item := rule.FirewallRuleSyncItem - if operation == "TaskDelete" && rule.observed != nil { - item.Rule = &rule.observed.Rule - } - switch { - case skipped: - result.Skipped++ - unexecuted++ - rule.done = true - case cause != nil: - appendDatabaseSyncFailure(&result, item, cause) - rule.done = true - if operation == "TaskDelete" { - failedRemovals[rule.SourceUUID] = cause - } - case operation == "TaskDelete": - removed++ - if rule.Status == firewallsync.StatusRemove { - result.Removed++ - rule.done = true - } - case operation == "TaskCreate": - created++ - result.Succeeded++ - rule.done = true - default: - result.Skipped++ - rule.done = true - } - if t == nil { - return - } - label := fmt.Sprintf("%s %s", i18n.GetMsgByKey(operation), item.SourceUUID) - if r := item.Rule; r != nil { - label += fmt.Sprintf(" [%s] %s %s:%s -> %s:%s %s", r.Scope.Key(), r.Protocol, r.SourceAddress, r.SourcePort, r.DestinationAddress, r.DestinationPort, r.Action) - } - if skipped { - t.Logf("%s %s: %v", label, i18n.GetMsgByKey("FirewallCreateRuleSkipped"), cause) - } else if operation == task.TaskSync && cause == nil { - t.Logf("%s %s", label, i18n.GetMsgByKey("FirewallSyncRuleUnchanged")) - } else { - t.LogWithStatus(label, cause) - } - } - defer func() { - if t != nil { - t.Log(i18n.GetMsgWithMap("FirewallSyncOperationsResult", map[string]interface{}{"created": created, "removed": removed, "failed": result.Failed, "skipped": unexecuted, "unchanged": result.Skipped - unexecuted})) - } - }() - blocked := false - for _, rule := range rules { - if rule.Status != firewallsync.StatusRemove { - result.Total++ - } - if rule.Status == firewallsync.StatusBlocked { - blocked = true - record("TaskSync", rule, errors.New(rule.Reason), false) - } - } - if blocked { - for _, rule := range rules { - if !rule.done { - record("TaskSync", rule, filter.ErrRuleOperation, true) - } - } - return result, nil - } - for _, rule := range rules { - if rule.Status == firewallsync.StatusExisting { - record("TaskSync", rule, nil, false) - } - } - for _, operation := range []filter.ChangeOperation{filter.ChangeDelete, filter.ChangeCreate} { - name := "TaskCreate" - if operation == filter.ChangeDelete { - name = "TaskDelete" - } - for _, initial := range snapshots { - scope := initial.Scope - queue := make([]*firewallSyncRule, 0) - for _, rule := range rules { - if rule.done || rule.Rule.Scope.Key() != scope.Key() { - continue - } - if operation == filter.ChangeDelete && rule.observed != nil || operation == filter.ChangeCreate && rule.Status == firewallsync.StatusReady { - queue = append(queue, rule) - } - } - if operation == filter.ChangeDelete { - sort.SliceStable(queue, func(i, j int) bool { - return syncObservedPosition(queue[i].observed) > syncObservedPosition(queue[j].observed) - }) - } - if scope.Provider == filter.ProviderFirewalld && operation == filter.ChangeCreate && len(failedRemovals) == 0 && stopped[scope.Key()] == nil { - firewalldFinalSnapshot = syncFirewalldCreates(ctx, runtime, initial, removed > 0, queue, t, record) - continue - } - markers := make([]string, 0) - for _, candidate := range rules { - if candidate.Rule != nil && candidate.Rule.Scope.Key() == scope.Key() && candidate.desired.Marker != "" { - markers = append(markers, candidate.desired.Marker) - } - } - for start := 0; start < len(queue); { - rule := queue[start] - cause := stopped[scope.Key()] - if operation == filter.ChangeCreate && cause == nil { - for id, err := range failedRemovals { - if id == rule.SourceUUID || strings.HasPrefix(id, rule.SourceUUID+"-") { - cause = err - break - } - } - } - if cause != nil { - record(name, rule, cause, true) - start++ - continue - } - snapshot := initial - err := ctx.Err() - if err == nil { - snapshot, err = runtime.ObserveMutation(ctx, scope) - } - unreadable := err != nil - end := start + 1 - if supportsNativeRuleBatch(scope.Provider) && (operation == filter.ChangeDelete || len(snapshot.Rules) == 0) && len(failedRemovals) == 0 { - end = min(start+filter.MaxAtomicExpansion, len(queue)) - } - batch := queue[start:end] - changes := make([]filter.DesiredChange, 0, len(batch)) - for _, entry := range batch { - if err != nil { - break - } - if operation == filter.ChangeDelete { - var change filter.DesiredChange - change, err = firewallsync.DeleteChange(snapshot, *entry.observed, entry.desired) - changes = append(changes, change) - } else { - after := *entry.Rule - after.OrderIndex = nil - if len(batch) == 1 && scope.Provider != filter.ProviderFirewalld { - after.OrderIndex = firewallsync.InsertionPosition(snapshot, markers, entry.desired.Marker) - } - changes = append(changes, filter.DesiredChange{Operation: operation, After: &after, Append: scope.Provider == filter.ProviderUFW && after.OrderIndex == nil}) - } - } - if err == nil { - _, err = runtime.ExecuteSync(ctx, snapshot, changes) - } - for _, entry := range batch { - record(name, entry, err, false) - } - if err != nil && (operation == filter.ChangeDelete || scope.Provider != filter.ProviderFirewalld || unreadable || firewallCreateUnavailable(err)) { - stopped[scope.Key()] = err - } - start = end - } - } - } - if firewalldFinalSnapshot != nil && result.Failed == 0 && created > 0 { - verifyErr := verifyFirewalldSyncSnapshot(*firewalldFinalSnapshot, rules) - if t != nil { - t.LogWithStatus(i18n.GetWithName("FirewallSyncStep", string(request.TargetProvider)), verifyErr) - } - if verifyErr != nil { - return result, verifyErr - } - } - return result, nil -} - -func syncFirewalldCreates( - ctx context.Context, - runtime *filterruntime.Engine, - initial filter.Snapshot, - refresh bool, - queue []*firewallSyncRule, - t *task.Task, - record func(string, *firewallSyncRule, error, bool), -) *filter.Snapshot { - if len(queue) == 0 { - return nil - } - planner, readErr := runtime.NewCreatePlanner(initial) - if readErr != nil { - for _, entry := range queue { - record(task.TaskCreate, entry, readErr, false) - } - return nil - } - pending := make([]*firewallSyncRule, 0, len(queue)) - for index, entry := range queue { - if refresh { - var snapshot filter.Snapshot - snapshot, readErr = runtime.ObserveMutation(ctx, initial.Scope) - if readErr == nil { - planner, readErr = runtime.NewCreatePlanner(snapshot) - } - if readErr != nil { - record(task.TaskCreate, entry, readErr, false) - for _, remaining := range queue[index+1:] { - record(task.TaskCreate, remaining, readErr, true) - } - break - } - refresh = false - } - if t != nil { - t.Logf("[%d/%d] %s %s", index+1, len(queue), i18n.GetMsgByKey(task.TaskCreate), entry.SourceUUID) - } - after := *entry.Rule - after.OrderIndex = nil - _, err := runtime.ExecutePlannedCreate(ctx, planner, filter.DesiredChange{ - Operation: filter.ChangeCreate, After: &after, CommandOnly: true, - }) - if err != nil { - record(task.TaskCreate, entry, err, false) - if firewallCreateUnavailable(err) { - for _, remaining := range queue[index+1:] { - record(task.TaskCreate, remaining, err, true) - } - break - } - refresh = true - continue - } - pending = append(pending, entry) - } - var actual filter.Snapshot - if readErr == nil { - actual, readErr = runtime.ObserveMutation(ctx, initial.Scope) - } - if readErr != nil { - for _, entry := range pending { - record(task.TaskCreate, entry, readErr, false) - } - return nil - } - states := firewalldRuleStates(actual) - for _, entry := range pending { - key, err := filter.RuleKey(*entry.Rule) - if err == nil && states[key] != 1 { - err = filter.ErrVerificationFailed - } - record(task.TaskCreate, entry, err, false) - } - return &actual -} - -func firewalldRuleStates(snapshot filter.Snapshot) map[string]int { - states := make(map[string]int, len(snapshot.Rules)) - for _, observed := range snapshot.Rules { - if observed.ParseStatus != filter.ParseStatusSupported || observed.Persistence != filter.PersistenceStatusConverged { - continue - } - if key, err := filter.RuleKey(observed.Rule); err == nil { - states[key]++ - } - } - return states -} - -func verifyFirewalldSyncSnapshot(snapshot filter.Snapshot, rules []*firewallSyncRule) error { - states := firewalldRuleStates(snapshot) - for _, entry := range rules { - if entry.Status == firewallsync.StatusRemove { - continue - } - if entry.Rule == nil { - return filter.ErrVerificationFailed - } - key, err := filter.RuleKey(*entry.Rule) - if err != nil { - return err - } - if states[key] != 1 { - return fmt.Errorf("%w: %s", filter.ErrVerificationFailed, entry.SourceUUID) - } - } - return nil -} - -func syncObservedPosition(rule *filter.ObservedRule) int { - if rule.Locator.Position != nil { - return *rule.Locator.Position - } - return 0 -} - -func (s *FirewallService) restoreStoredFirewallRules(ctx context.Context, provider filter.Provider, t *task.Task) error { - result, err := s.syncRules(ctx, "", dto.FirewallRuleSyncRequest{TargetProvider: provider}, t) - if err != nil { - return fmt.Errorf("restore database firewall rules: %w", err) - } - failures := make([]error, 0, len(result.Errors)) - for _, failure := range result.Errors { - failures = append(failures, fmt.Errorf("rule %s: %s", failure.SourceUUID, failure.Error)) - } - return errors.Join(failures...) -} - -func (s *FirewallService) SyncRules( - ctx context.Context, - clientIP string, - request dto.FirewallRuleSyncRequest, -) (dto.FirewallRuleSyncResult, error) { - switch firewallSyncSubsystem(request.Subsystem) { - case "forwarding": - return s.forwardingRuleSyncService().syncRules(ctx, request) - case "docker": - return s.dockerRuleSyncService().syncRules(ctx, request) - default: - return s.syncSystemRules(ctx, clientIP, request) - } -} - -func (s *FirewallService) syncSystemRules( - ctx context.Context, - clientIP string, - request dto.FirewallRuleSyncRequest, -) (dto.FirewallRuleSyncResult, error) { - if err := lockFirewallLifecycleIdle(); err != nil { - return dto.FirewallRuleSyncResult{}, err - } - defer firewallLifecycleTaskMu.Unlock() - firewallRuleSyncTaskMu.Lock() - defer firewallRuleSyncTaskMu.Unlock() - - running, err := currentFirewallRuleSyncTaskLocked() - if err != nil { - return dto.FirewallRuleSyncResult{}, err - } - if running.Executing { - return dto.FirewallRuleSyncResult{ - Subsystem: firewallSyncSubsystem(request.Subsystem), - TargetProvider: request.TargetProvider, - TaskID: running.TaskID, - Queued: true, - }, nil - } - if firewallSyncSubsystem(request.Subsystem) != "system" { - return dto.FirewallRuleSyncResult{}, fmt.Errorf("%w: firewall synchronization tasks are only available for the system firewall", filter.ErrInvalidRule) - } - taskItem, err := task.NewTask(firewallTaskName(task.TaskSync, firewallTaskHost, string(request.TargetProvider)), task.TaskSync, task.TaskScopeFirewall, "", 0) - if err != nil { - return dto.FirewallRuleSyncResult{}, fmt.Errorf("create firewall sync task: %w", err) - } - taskItem.AddSubTaskWithOps(i18n.GetWithName("FirewallSyncStep", string(request.TargetProvider)), func(t *task.Task) error { - result, err := s.syncRules(t.TaskCtx, clientIP, request, t) - if err != nil { - return err - } - if result.Failed > 0 { - return errors.New(i18n.GetMsgWithMap("FirewallSyncFailed", map[string]interface{}{"failed": result.Failed})) - } - return nil - }, nil, 0, 0) - - if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { - taskItem.LogFailedWithErr(taskItem.Name, err) - closeUnstartedFirewallTask(taskItem) - return dto.FirewallRuleSyncResult{}, fmt.Errorf("save firewall sync task: %w", err) - } - firewallRuleSyncTaskID = taskItem.TaskID - go func() { - defer func() { - firewallRuleSyncTaskMu.Lock() - if firewallRuleSyncTaskID == taskItem.TaskID { - firewallRuleSyncTaskID = "" - } - firewallRuleSyncTaskMu.Unlock() - }() - if err := taskItem.Execute(); err != nil && global.LOG != nil { - global.LOG.Errorf("firewall sync task %s failed: %v", taskItem.TaskID, err) - } - }() - return dto.FirewallRuleSyncResult{ - Subsystem: "system", - TargetProvider: request.TargetProvider, - TaskID: taskItem.TaskID, - Queued: true, - }, nil -} - -func (s *FirewallService) CurrentRuleSyncTask() (dto.FirewallRuleSyncTask, error) { - firewallRuleSyncTaskMu.Lock() - defer firewallRuleSyncTaskMu.Unlock() - return currentFirewallRuleSyncTaskLocked() -} - -func currentFirewallRuleSyncTaskLocked() (dto.FirewallRuleSyncTask, error) { - if firewallRuleSyncTaskID != "" { - return dto.FirewallRuleSyncTask{TaskID: firewallRuleSyncTaskID, Executing: true}, nil - } - if global.TaskDB == nil { - return dto.FirewallRuleSyncTask{}, nil - } - taskRepo := repo.NewITaskRepo() - record, err := taskRepo.GetFirst( - repo.WithByStatus(constant.StatusExecuting), - repo.WithByType(task.TaskScopeFirewall), - taskRepo.WithOperate(task.TaskSync), - ) - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return dto.FirewallRuleSyncTask{}, nil - } - return dto.FirewallRuleSyncTask{}, err - } - return dto.FirewallRuleSyncTask{TaskID: record.ID, Executing: true}, nil -} - -func whitelistRules(provider filter.Provider, ports, required []firewall.SystemPort) []filter.FirewallRule { - rules := make([]filter.FirewallRule, 0, len(ports)+len(required)) - for _, port := range required { - rule := systemPortRule(provider, port) - if isDirectFirewallProvider(provider) { - rule.Scope.Chain = filter.BasicBeforeChain - } - rules = append(rules, rule) - } - for _, port := range ports { - rules = append(rules, systemPortRule(provider, port)) - } - return rules -} - func (s *FirewallService) SyncPortWhitelist(ctx context.Context) error { firewallWhitelistMu.Lock() defer firewallWhitelistMu.Unlock() - filterruntime.InvalidateInventory() - defer filterruntime.InvalidateInventory() + ports, err := loadFirewallPortWhiteList() if err != nil { return err @@ -699,20 +51,7 @@ func (s *FirewallService) SyncPortWhitelist(ctx context.Context) error { return err } rules := whitelistRules(provider, firewall.ExpandPortWhitelist(customWhitelist(ports)), firewall.ExpandPortWhitelist(required)) - activeFamilies := make(map[filter.Family]bool) - if isDirectFirewallProvider(provider) { - for _, rule := range rules { - family := rule.Scope.Family - if _, checked := activeFamilies[family]; checked { - continue - } - initialized, bound, err := loadSystemFirewallFamilyStatus(string(provider), string(family)) - if err != nil { - return err - } - activeFamilies[family] = initialized && bound - } - } else if len(rules) > 0 { + if provider != filter.ProviderIptables && provider != filter.ProviderNftables && len(rules) > 0 { client, err := s.baseClient() if err != nil { return err @@ -722,18 +61,34 @@ func (s *FirewallService) SyncPortWhitelist(ctx context.Context) error { return err } } - prepared := make([]preparedFirewallRuleCreate, 0, len(rules)) + client, err := s.firewallAdapter(provider) + if err != nil { + return err + } + prepared := make([]dto.FirewallRuleCreateItem, 0, len(rules)) var failures []error + activeFamilies := make(map[filter.Family]bool) for _, rule := range rules { if err := ctx.Err(); err != nil { return err } - if isDirectFirewallProvider(provider) && !activeFamilies[rule.Scope.Family] { - continue + if provider == filter.ProviderIptables || provider == filter.ProviderNftables { + active, inspected := activeFamilies[rule.Scope.Family] + if !inspected { + initialized, bound, err := loadSystemFirewallFamilyStatus(string(provider), string(rule.Scope.Family)) + if err != nil { + return err + } + active = initialized && bound + activeFamilies[rule.Scope.Family] = active + } + if !active { + continue + } } port := firewall.SystemPort{Family: string(rule.Scope.Family), Port: rule.DestinationPort, Protocol: rule.Protocol, SourceAddress: rule.SourceAddress} - item, err := s.prepareCreate(ctx, provider, dto.FirewallRuleCreateItem{ - Rule: rule, SourceKind: constant.FirewallRuleSourceSecurity, SourceID: systemPortSourceID(port), + item, err := prepareFirewallCreateRule(ctx, client, dto.FirewallRuleCreateItem{ + Rule: rule, SourceKind: constant.FirewallRuleSourceSecurity, SourceID: constant.FirewallSystemAcceptedPortSourcePrefix + firewall.SystemPortKey(firewall.SystemPort(port)), }) if err != nil { failures = append(failures, fmt.Errorf("prepare whitelist rule %s: %w", firewall.SystemPortKey(port), err)) @@ -744,251 +99,271 @@ func (s *FirewallService) SyncPortWhitelist(ctx context.Context) error { if len(failures) > 0 { return errors.Join(failures...) } - if provider == filter.ProviderUFW && len(prepared) > 0 { - return s.syncUFWPortWhitelist(ctx, prepared) + if len(prepared) == 0 { + return nil } + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + scopes := make([]filter.Scope, 0, len(prepared)) for _, item := range prepared { - if err := ctx.Err(); err != nil { - return err - } - if err := s.addWhitelistRule(ctx, item); err != nil { - rule := item.request.Rule - failures = append(failures, fmt.Errorf("add whitelist rule %s %s/%s [%s]: %w", rule.Scope.Family, rule.DestinationPort, rule.Protocol, rule.SourceAddress, err)) - } + scopes = append(scopes, item.Rule.Scope) } - return errors.Join(failures...) -} - -func (s *FirewallService) syncUFWPortWhitelist(ctx context.Context, prepared []preparedFirewallRuleCreate) error { - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - runtime := prepared[0].runtime - snapshots, err := runtime.ObserveScopes(ctx, filter.ManagedInputScopes(filter.ProviderUFW)) + snapshots, err := readMutableFirewallRuleScopes(client, ctx, scopes) if err != nil { return err } - for itemIndex, item := range prepared { + byScope := make(map[string]filter.RuleSet, len(snapshots)) + byScopeIdentity := make(map[string]firewallRuleCollisionIndex, len(snapshots)) + for _, snapshot := range snapshots { + identities, err := observedFirewallRuleCollisionIndex(snapshot) + if err != nil { + return err + } + byScope[snapshot.Scope.Key()] = snapshot + byScopeIdentity[snapshot.Scope.Key()] = identities + } + var byMatch map[string][]filter.DesiredRule + type whitelistBatch struct { + snapshot filter.RuleSet + changes []filter.RuleChange + records []*model.FirewallRule + } + batches := make([]whitelistBatch, 0) + batchByScope := make(map[string]int) + for _, item := range prepared { if err := ctx.Err(); err != nil { return err } - rule := item.request.Rule - var snapshot *filter.Snapshot - for index := range snapshots { - if snapshots[index].Scope.Key() == rule.Scope.Key() { - snapshot = &snapshots[index] - break - } - } - if snapshot == nil { - return fmt.Errorf("%w: missing UFW whitelist scope %s", filter.ErrInventoryUnavailable, rule.Scope.Key()) + rule := item.Rule + scope := rule.Scope + identities := byScopeIdentity[scope.Key()] + if err := identities.CheckDuplicate(rule); errors.Is(err, filter.ErrRuleOperation) { + continue + } else if err != nil { + return err } - for _, notice := range snapshot.Notices { - if notice.Code == filter.ScopeNoticeManagedScopeInactive || notice.Code == filter.ScopeNoticeManagedScopeMissing { - return filter.ErrProviderUnavailable - } + if err := identities.Check(rule); err != nil { + return err } - added, err := s.addWhitelistRuleFromSnapshot(ctx, item, *snapshot) + key, err := filter.RuleMatchKey(rule) if err != nil { - return fmt.Errorf("add whitelist rule %s %s/%s [%s]: %w", rule.Scope.Family, rule.DestinationPort, rule.Protocol, rule.SourceAddress, err) - } - if !added { - continue + return err } - - if itemIndex+1 < len(prepared) { - if err := ctx.Err(); err != nil { - return err - } - snapshots, err = runtime.ObserveScopes(ctx, filter.ManagedInputScopes(filter.ProviderUFW)) + if byMatch == nil { + stored, err := s.rules.List(ctx) if err != nil { return err } + byMatch = make(map[string][]filter.DesiredRule) + for _, record := range stored { + rules, err := compileStoredFirewallRules(ctx, record, client) + if isFirewallPolicyIncompatible(err) { + continue + } + if err != nil { + return err + } + for _, candidate := range rules { + key, err := filter.RuleMatchKey(candidate.Rule) + if err != nil { + return err + } + byMatch[key] = append(byMatch[key], candidate) + } + } } - } - return nil -} - -func (s *FirewallService) addWhitelistRule(ctx context.Context, prepared preparedFirewallRuleCreate) error { - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - if err := ctx.Err(); err != nil { - return err - } - rule, runtime := prepared.request.Rule, prepared.runtime - snapshot, err := runtime.ObserveMutation(ctx, rule.Scope) - if err != nil { - return err - } - _, err = s.addWhitelistRuleFromSnapshot(ctx, prepared, snapshot) - return err -} - -func (s *FirewallService) addWhitelistRuleFromSnapshot(ctx context.Context, prepared preparedFirewallRuleCreate, snapshot filter.Snapshot) (bool, error) { - rule, runtime := prepared.request.Rule, prepared.runtime - for _, observed := range snapshot.Rules { - if observed.ParseStatus == filter.ParseStatusSupported { - if same, err := filter.SameRuleContent(rule, observed.Rule); err == nil && same { - return false, nil + var existing *filter.DesiredRule + for _, candidate := range byMatch[key] { + if collision := checkCollisionActions(rule.Action, candidate.Rule.Action); errors.Is(collision, filter.ErrRuleOperation) { + existing = &candidate + } else if collision != nil { + return collision } } + var record *model.FirewallRule + if existing != nil && scope.Chain != filter.BasicBeforeChain { + rule = existing.Rule + } else { + rule.UUID = uuid.NewString() + if scope.Chain != filter.BasicBeforeChain { + created, err := firewallRuleModelForCreate(rule, item, constant.FirewallRuleOriginCreated) + if err != nil { + return err + } + created.UUID = rule.UUID + record = &created + } + } + if provider != filter.ProviderFirewalld { + position := int64(1) + rule.OrderIndex = &position + } + index, exists := batchByScope[scope.Key()] + if !exists || provider != filter.ProviderIptables && provider != filter.ProviderNftables { + index = len(batches) + batchByScope[scope.Key()] = index + batches = append(batches, whitelistBatch{snapshot: byScope[scope.Key()]}) + } + batches[index].changes = append(batches[index].changes, filter.RuleChange{Operation: filter.ChangeCreate, After: &rule, CommandOnly: true}) + batches[index].records = append(batches[index].records, record) + if err := identities.Add(rule); err != nil { + return err + } } - if err := filter.CheckObservedRuleCollisions(snapshot, rule, nil); err != nil { - return false, err - } - if rule.Scope.Chain == filter.BasicBeforeChain { - position := int64(1) - rule.OrderIndex = &position - rule.UUID = uuid.NewString() - return true, runtime.ExecuteCreate(ctx, snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}) - } - stored, err := s.rules.List(ctx) - if err != nil { - return false, err - } - model.SortFirewallRules(stored, rule.Scope.Provider) - markers := make([]string, 0, len(stored)) - var existing *filter.DesiredRule - var firstSequence *int64 - for _, record := range stored { - compiled, err := s.compileStoredFirewallRules(ctx, record, rule.Scope.Provider) - if err != nil { - if isFirewallPolicyIncompatible(err) { + saveRecords := func(batch whitelistBatch) { + saveCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + for _, record := range batch.records { + if record == nil { continue } - return false, err - } - for _, candidate := range compiled { - if candidate.Rule.Scope.Key() == rule.Scope.Key() && candidate.Marker != "" { - markers = append(markers, candidate.Marker) - if firstSequence == nil && record.Sequence != nil { - firstSequence = record.Sequence + if err := s.saveFirewallRule(saveCtx, record); err != nil { + failures = append(failures, err) + if global.LOG != nil { + global.LOG.Errorf("save firewall whitelist rule %s failed: %v", record.UUID, err) } } - if err := filter.CheckRuleCollision(rule, candidate.Rule); errors.Is(err, filter.ErrRuleOperation) { - copy := candidate - existing = © - } else if err != nil { - return false, err - } } + cancel() } - if existing != nil { - existing.Rule.OrderIndex = firewallsync.InsertionPosition(snapshot, markers, existing.Marker) - return true, runtime.ExecuteCreate(ctx, snapshot, []filter.DesiredChange{{ - Operation: filter.ChangeCreate, After: &existing.Rule, - Append: rule.Scope.Provider == filter.ProviderUFW && existing.Rule.OrderIndex == nil, - }}) - } - if rule.Scope.Provider != filter.ProviderFirewalld { - position := int64(1) - rule.OrderIndex = &position - record, err := firewallRuleModelForCreate(rule, prepared.request, constant.FirewallRuleOriginCreated) + _, savesRules := client.(filter.RuleSaver) + plans := make([]filter.CommandBatch, 0, len(batches)) + completed := make([]int, 0, len(batches)) + for batchIndex, batch := range batches { + if err := ctx.Err(); err != nil { + failures = append(failures, err) + break + } + commands, err := client.BuildCommands(batch.snapshot, batch.changes) + commands.CommandOnly = true + if err == nil { + err = client.RunCommands(ctx, commands) + } if err != nil { - return false, err + failures = append(failures, err) + if global.LOG != nil { + global.LOG.Errorf("create firewall whitelist rules for %s failed: %v", batch.snapshot.Scope.Key(), err) + } + continue } - sequence := model.FirewallRuleSequenceStep - if firstSequence != nil { - sequence = *firstSequence - model.FirewallRuleSequenceStep + verification := filter.CommandBatch{Provider: commands.Provider, Scope: commands.Scope} + for _, command := range commands.Rules { + expected := command.Expected + expected.Locator.Position = nil + verification.Rules = append(verification.Rules, filter.RuleCommands{Operation: filter.ChangeCreate, Expected: expected}) } - record.UUID, record.Sequence = uuid.NewString(), &sequence - rule.UUID = record.UUID - if err := runtime.ExecuteCreate(ctx, snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}); err != nil { - return false, firewallCreateExecutionError(err) + plans = append(plans, verification) + completed = append(completed, batchIndex) + if !savesRules { + saveRecords(batch) } - return true, s.saveFirewallRule(ctx, &record) - } - return true, s.applyCreateRules(ctx, runtime, snapshot, stored, []preparedFirewallRuleCreate{prepared})[0] -} - -type forwardingRuleSyncCandidate struct { - rule forwarding.Rule - err error -} - -func (s *ForwardingService) loadRuleSyncCandidates( - ctx context.Context, - targetProvider filter.Provider, -) (*forwarding.Manager, []forwardingRuleSyncCandidate, []forwarding.Rule, bool, error) { - target, err := s.managerFactory() - if err != nil { - return nil, nil, nil, false, err } - if target.Name() != string(targetProvider) { - return nil, nil, nil, false, fmt.Errorf( - "%w: selected forwarding backend is %s, requested target is %s", - filter.ErrProviderUnavailable, target.Name(), targetProvider, - ) + persistenceErrors := persistFirewallRuleBatches(ctx, client, plans) + for index, batchIndex := range completed { + if err := persistenceErrors[index]; err != nil { + failures = append(failures, err) + continue + } + if savesRules { + saveRecords(batches[batchIndex]) + } } - stored, err := s.rules.List(ctx) + + verification, err := verifyFirewallCommands(ctx, client, plans...) if err != nil { - return nil, nil, nil, false, err + failures = append(failures, err) + } else if !verification.Matched { + failures = append(failures, filter.ErrVerificationFailed) } - candidates := make([]forwardingRuleSyncCandidate, 0, len(stored)) - for _, record := range stored { - rule := forwarding.Rule{ - Family: record.Family, Protocol: record.Protocol, Port: record.Port, TargetIP: record.TargetIP, - TargetPort: record.TargetPort, Interface: record.Interface, + return errors.Join(failures...) +} + +func (s *FirewallService) PreviewRuleSync(ctx context.Context, clientIP string, request dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncPreview, error) { + switch strings.TrimSpace(request.Subsystem) { + case "forwarding": + service := s.forwardingSync + if service == nil { + service = newForwardingService() + } + return service.previewRuleSync(ctx, request) + case "docker": + service := s.dockerSync + if service == nil { + service = newDockerPortGuardService() } - normalized, normalizeErr := forwarding.NormalizeRule(rule) - candidates = append(candidates, forwardingRuleSyncCandidate{rule: normalized, err: normalizeErr}) + return service.previewRuleSync(ctx, request) } - targetStatus, err := target.Status() - if err != nil { - return nil, nil, nil, false, err + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + _, rules, _, err := s.loadFirewallSyncRules(ctx, request) + preview := dto.FirewallRuleSyncPreview{Subsystem: "system", TargetProvider: request.TargetProvider, Items: make([]dto.FirewallRuleSyncItem, 0, len(rules))} + for _, rule := range rules { + preview.Add(rule.FirewallRuleSyncItem) } - targetRules := make([]forwarding.Rule, 0) - if targetStatus.IsInit { - targetRules, err = target.List("", "") - if err != nil { - return nil, nil, nil, false, err + return preview, err +} + +func (s *FirewallService) SyncRules(ctx context.Context, clientIP string, request dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error) { + switch strings.TrimSpace(request.Subsystem) { + case "forwarding": + service := s.forwardingSync + if service == nil { + service = newForwardingService() } - targetRules, err = normalizeForwardingRuntimeRules(targetRules) - if err != nil { - return nil, nil, nil, false, err + return service.syncRules(ctx, request) + case "docker": + service := s.dockerSync + if service == nil { + service = newDockerPortGuardService() } + return service.syncRules(ctx, request) + default: + return s.syncSystemRules(ctx, clientIP, request) } - return target, candidates, targetRules, targetStatus.IsInit, nil } -func verifyForwardingRuleSync(target *forwarding.Manager, desired []forwarding.Rule) error { - actual, err := target.List("", "") - if err != nil { - return fmt.Errorf("verify synchronized forwarding rules: %w", err) - } - actual, err = normalizeForwardingRuntimeRules(actual) - if err != nil { - return fmt.Errorf("verify synchronized forwarding rules: %w", err) - } - if !firewallsync.StatesEqual(actual, desired, func(rule forwarding.Rule) string { return rule.Identity() }) { - return fmt.Errorf("verify synchronized forwarding rules: target rules do not match the database") - } - return nil +func (s *FirewallService) CurrentRuleSyncTask() (dto.FirewallRuleSyncTask, error) { + firewallRuleSyncTaskMu.Lock() + defer firewallRuleSyncTaskMu.Unlock() + return currentFirewallRuleSyncTaskLocked() } -func normalizeForwardingRuntimeRules(rules []forwarding.Rule) ([]forwarding.Rule, error) { - normalized := make([]forwarding.Rule, 0, len(rules)) - for _, rule := range rules { - item, err := forwarding.NormalizeRule(rule) +func (s *DockerPortGuardService) Reconcile(ctx context.Context) error { + dockerPortGuardServiceMu.Lock() + defer dockerPortGuardServiceMu.Unlock() + return s.reconcileLocked(ctx) +} + +func (s *ForwardingService) Restore(ctx context.Context) error { + forwardingMutationMu.Lock() + defer forwardingMutationMu.Unlock() + enabled, err := s.forwardingEnabled() + if err != nil || !enabled { if err != nil { - return nil, fmt.Errorf("normalize target forwarding rule %s: %w", rule.Identity(), err) + recordForwardingSyncError(err) } - normalized = append(normalized, item) + return err } - return normalized, nil -} - -func forwardingRuleSyncDTO(rule forwarding.Rule) *dto.ForwardRule { - return &dto.ForwardRule{ - Family: rule.Family, Protocol: rule.Protocol, Port: rule.Port, TargetIP: rule.TargetIP, - TargetPort: rule.TargetPort, Interface: rule.Interface, + manager, err := s.clientFactory() + if err != nil { + recordForwardingSyncError(err) + return err + } + stored, err := s.rules.List(ctx) + if err != nil { + recordForwardingSyncError(err) + return err + } + if err := s.initializeForwarding(manager); err != nil { + recordForwardingSyncError(err) + return err } + err = manager.ReplaceRules(forwardingRulesFromModels(stored)) + recordForwardingSyncError(err) + return err } -func (s *ForwardingService) previewRuleSync( - ctx context.Context, - request dto.FirewallRuleSyncRequest, -) (dto.FirewallRuleSyncPreview, error) { +func (s *ForwardingService) previewRuleSync(ctx context.Context, request dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncPreview, error) { targetProvider, err := databaseRuleSyncTarget(request, "forwarding") if err != nil { return dto.FirewallRuleSyncPreview{}, err @@ -1000,10 +375,7 @@ func (s *ForwardingService) previewRuleSync( return forwardingSyncPreview(filter.Provider(target.Name()), candidates, targetRules), nil } -func (s *ForwardingService) syncRules( - ctx context.Context, - request dto.FirewallRuleSyncRequest, -) (dto.FirewallRuleSyncResult, error) { +func (s *ForwardingService) syncRules(ctx context.Context, request dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error) { forwardingMutationMu.Lock() defer forwardingMutationMu.Unlock() @@ -1033,11 +405,11 @@ func (s *ForwardingService) syncRules( if err := s.persistForwardingEnabled(); err != nil { return err } - if err := s.activateManager(target); err != nil { + if err := s.initializeForwarding(target); err != nil { return err } } - if err := target.Reconcile(desired); err != nil { + if err := target.ReplaceRules(desired); err != nil { return err } return verifyForwardingRuleSync(target, desired) @@ -1053,125 +425,7 @@ func (s *ForwardingService) syncRules( return result, nil } -func forwardingSyncPreview( - target filter.Provider, - candidates []forwardingRuleSyncCandidate, - actual []forwarding.Rule, -) dto.FirewallRuleSyncPreview { - desired := make([]firewallsync.Desired[forwarding.Rule, dto.FirewallRuleSyncItem], 0, len(candidates)) - for _, candidate := range candidates { - desired = append(desired, firewallsync.Desired[forwarding.Rule, dto.FirewallRuleSyncItem]{ - Value: candidate.rule, - Payload: dto.FirewallRuleSyncItem{ - SourceUUID: candidate.rule.Identity(), ForwardRule: forwardingRuleSyncDTO(candidate.rule), - }, - Err: candidate.err, - }) - } - return firewallDiffPreview( - "forwarding", target, desired, actual, - func(rule forwarding.Rule) string { return rule.Identity() }, - func(rule forwarding.Rule) dto.FirewallRuleSyncItem { - return dto.FirewallRuleSyncItem{SourceUUID: rule.Identity(), ForwardRule: forwardingRuleSyncDTO(rule)} - }, - ) -} - -func (s *FirewallService) forwardingRuleSyncService() firewallDatabaseSyncAdapter { - if s.forwardingSync == nil { - return newForwardingService() - } - return s.forwardingSync -} - -func (s *DockerPortGuardService) loadRuleSyncCandidates( - ctx context.Context, - request dto.FirewallRuleSyncRequest, -) (string, []model.DockerPortGuardPolicy, dockerGuardRuntime, error) { - targetProvider, err := databaseRuleSyncTarget(request, "Docker") - if err != nil { - return "", nil, nil, err - } - target := string(targetProvider) - selected, err := s.selectedRuleSyncBackend(ctx) - if err != nil { - return "", nil, nil, err - } - if target != selected { - return "", nil, nil, fmt.Errorf( - "%w: selected Docker firewall backend is %s, requested target is %s", - filter.ErrProviderUnavailable, selected, target, - ) - } - policies, err := s.policies.ListManaged(ctx) - if err != nil { - return "", nil, nil, err - } - return target, policies, s.guardRuntime(target), nil -} - -func (s *DockerPortGuardService) selectedRuleSyncBackend(ctx context.Context) (string, error) { - if global.DB != nil { - selected, _ := settingRepo.GetValueByKey(constant.FirewallDockerBackendKey) - selected = strings.ToLower(strings.TrimSpace(selected)) - if selected == constant.FirewallProviderIptables || selected == constant.FirewallProviderNftables { - return selected, nil - } - } - if s.client == nil { - return "", fmt.Errorf("%w: Docker firewall backend is unavailable", ErrDockerUnavailable) - } - cli, err := s.client() - if err != nil { - return "", fmt.Errorf("%w: %v", ErrDockerUnavailable, err) - } - defer cli.Close() - info, err := cli.Info(ctx) - if err != nil { - return "", fmt.Errorf("%w: %v", ErrDockerUnavailable, err) - } - return selectedDockerFirewallBackend(dockerFirewallBackend(info)), nil -} - -func dockerGuardPoliciesFromModels(policies []model.DockerPortGuardPolicy) []docker_guard.Policy { - result := make([]docker_guard.Policy, 0, len(policies)) - for _, policy := range policies { - result = append(result, dockerGuardPolicyFromModel(policy)) - } - return result -} - -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: docker_guard.DecodeSources(policy.Sources), Description: policy.Description, - TrafficPath: dockerTrafficPathUnknown, ManagementTarget: dockerManagementNeedsDiagnosis, - ManagementReason: dockerReasonNoMatchingPath, - } -} - -func dockerGuardRuntimeRuleSyncDTO(policy docker_guard.Policy) *dto.DockerPortGuardEndpoint { - return &dto.DockerPortGuardEndpoint{ - Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, - PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: append([]string(nil), policy.Sources...), - TrafficPath: dockerTrafficPathUnknown, ManagementTarget: dockerManagementNeedsDiagnosis, - ManagementReason: dockerReasonNoMatchingPath, - } -} - -func dockerGuardReadOnlyRuleSyncDTO(policy docker_guard.ReadOnlyPolicy) *dto.DockerPortGuardEndpoint { - return &dto.DockerPortGuardEndpoint{ - Family: policy.Policy.Family, HostIP: policy.Policy.HostIP, HostPort: policy.Policy.HostPort, - Protocol: policy.Policy.Protocol, PolicyUUID: dockerGuardReadOnlyPolicyUUID(policy), Sources: append([]string(nil), policy.Policy.Sources...), - NativeAction: policy.Action, ReadOnly: true, TrafficPath: dockerTrafficPathUnknown, - ManagementTarget: dockerManagementNeedsDiagnosis, ManagementReason: dockerReasonNoMatchingPath, - } -} - -func (s *DockerPortGuardService) previewRuleSync( - ctx context.Context, - request dto.FirewallRuleSyncRequest, -) (dto.FirewallRuleSyncPreview, error) { +func (s *DockerPortGuardService) previewRuleSync(ctx context.Context, request dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncPreview, error) { target, policies, runtime, err := s.loadRuleSyncCandidates(ctx, request) if err != nil { return dto.FirewallRuleSyncPreview{}, err @@ -1183,10 +437,7 @@ func (s *DockerPortGuardService) previewRuleSync( return dockerSyncPreview(filter.Provider(target), policies, targetInventory), nil } -func (s *DockerPortGuardService) syncRules( - ctx context.Context, - request dto.FirewallRuleSyncRequest, -) (dto.FirewallRuleSyncResult, error) { +func (s *DockerPortGuardService) syncRules(ctx context.Context, request dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error) { dockerPortGuardServiceMu.Lock() defer dockerPortGuardServiceMu.Unlock() @@ -1194,7 +445,15 @@ func (s *DockerPortGuardService) syncRules( if err != nil { return dto.FirewallRuleSyncResult{}, err } - runtimePolicies := dockerGuardPoliciesFromModels(policies) + runtimePolicies := make([]dockerfirewall.Policy, 0, len(policies)) + for _, policy := range policies { + sources := []string{} + _ = json.Unmarshal([]byte(policy.Sources), &sources) + runtimePolicies = append(runtimePolicies, dockerfirewall.Policy{ + UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, + Protocol: policy.Protocol, Mode: policy.Mode, Sources: sources, + }) + } targetInventory, err := targetRuntime.ListPolicies() if err != nil { return dto.FirewallRuleSyncResult{}, err @@ -1209,13 +468,10 @@ func (s *DockerPortGuardService) syncRules( return firewallSyncResult(preview, nil, false), nil } reconcileErr := func() error { - if err := s.replaceRuntimeReadOnlyPolicies(ctx, targetInventory.ReadOnly); err != nil { - return err - } - if err := docker_guard.ReconcileTarget(target, runtimePolicies, targetRuntime); err != nil { + if err := reconcileDockerFirewall(target, runtimePolicies, targetRuntime, targetInventory); err != nil { return err } - if err := docker_guard.Verify(targetRuntime, runtimePolicies, targetInventory.ReadOnly); err != nil { + if err := verifyDockerFirewall(targetRuntime, runtimePolicies, targetInventory.ReadOnly); err != nil { return err } if len(policies) == 0 { @@ -1227,109 +483,144 @@ func (s *DockerPortGuardService) syncRules( return settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusEnable) }() result := firewallSyncResult(preview, reconcileErr, true) - recordDockerPortGuardReconcileError(reconcileErr) if reconcileErr != nil { return result, reconcileErr } return result, nil } -func dockerSyncPreview( - target filter.Provider, - policies []model.DockerPortGuardPolicy, - inventory docker_guard.PolicyInventory, -) dto.FirewallRuleSyncPreview { - desired := make([]firewallsync.Desired[docker_guard.Policy, dto.FirewallRuleSyncItem], 0, len(policies)) - for _, policy := range policies { - desired = append(desired, firewallsync.Desired[docker_guard.Policy, dto.FirewallRuleSyncItem]{ - Value: dockerGuardPolicyFromModel(policy), - Payload: dto.FirewallRuleSyncItem{SourceUUID: policy.UUID, DockerRule: dockerGuardRuleSyncDTO(policy)}, - }) +func (s *FirewallService) syncSystemRules(ctx context.Context, clientIP string, request dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error) { + subsystem := strings.TrimSpace(request.Subsystem) + if subsystem == "" { + subsystem = "system" } - preview := firewallDiffPreview( - "docker", target, desired, inventory.Policies, docker_guard.PolicySyncKey, - func(policy docker_guard.Policy) dto.FirewallRuleSyncItem { - return dto.FirewallRuleSyncItem{SourceUUID: policy.UUID, DockerRule: dockerGuardRuntimeRuleSyncDTO(policy)} - }, - ) - for _, policy := range inventory.ReadOnly { - preview.Add(dto.FirewallRuleSyncItem{ - SourceUUID: dockerGuardReadOnlyPolicyUUID(policy), - DockerRule: dockerGuardReadOnlyRuleSyncDTO(policy), - Status: firewallsync.StatusBlocked, - ReasonCode: firewallsync.ReasonReadOnlyRule, - Reason: firewallsync.ReasonMessage(firewallsync.ReasonReadOnlyRule), - }) + if err := lockFirewallLifecycleIdle(); err != nil { + return dto.FirewallRuleSyncResult{}, err } - return preview -} + defer firewallLifecycleTaskMu.Unlock() + firewallRuleSyncTaskMu.Lock() + defer firewallRuleSyncTaskMu.Unlock() -func (s *FirewallService) dockerRuleSyncService() firewallDatabaseSyncAdapter { - if s.dockerSync == nil { - return newDockerPortGuardService() + running, err := currentFirewallRuleSyncTaskLocked() + if err != nil { + return dto.FirewallRuleSyncResult{}, err } - return s.dockerSync -} - -func databaseRuleSyncTarget(request dto.FirewallRuleSyncRequest, subsystem string) (filter.Provider, error) { - if request.SourceProvider != "" { - return "", fmt.Errorf("%w: %s synchronization reads rules from the database and does not accept a source provider", filter.ErrInvalidRule, subsystem) + if running.Executing { + return dto.FirewallRuleSyncResult{ + Subsystem: subsystem, + TargetProvider: request.TargetProvider, + TaskID: running.TaskID, + Queued: true, + }, nil } - if request.ResetSource { - return "", fmt.Errorf("%w: %s synchronization does not have a source firewall to reset", filter.ErrInvalidRule, subsystem) + if subsystem != "system" { + return dto.FirewallRuleSyncResult{}, fmt.Errorf("%w: firewall synchronization tasks are only available for the system firewall", filter.ErrInvalidRule) } - if request.TargetProvider != filter.ProviderIptables && request.TargetProvider != filter.ProviderNftables { - return "", fmt.Errorf("%w: %s synchronization only supports iptables and nftables targets", filter.ErrInvalidRule, subsystem) + taskItem, err := task.NewTask(firewallTaskName(task.TaskSync, firewallTaskHost, string(request.TargetProvider)), task.TaskSync, task.TaskScopeFirewall, "", 0) + if err != nil { + return dto.FirewallRuleSyncResult{}, fmt.Errorf("create firewall sync task: %w", err) } - return request.TargetProvider, nil -} + taskItem.AddSubTaskWithOps(i18n.GetWithName("FirewallSyncStep", string(request.TargetProvider)), func(t *task.Task) error { + result, err := s.syncRules(t.TaskCtx, clientIP, request, t) + if err != nil { + return err + } + if result.Failed > 0 { + return errors.New(i18n.GetMsgWithMap("FirewallSyncFailed", map[string]interface{}{"failed": result.Failed})) + } + return nil + }, nil, 0, 0) -func appendDatabaseSyncFailure(result *dto.FirewallRuleSyncResult, item dto.FirewallRuleSyncItem, err error) { - if err == nil { - err = errors.New("database synchronization failed") + if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { + taskItem.LogFailedWithErr(taskItem.Name, err) + closeUnstartedFirewallTask(taskItem) + return dto.FirewallRuleSyncResult{}, fmt.Errorf("save firewall sync task: %w", err) } - result.Failed++ - result.Errors = append(result.Errors, dto.FirewallRuleSyncFailure{ - SourceUUID: item.SourceUUID, - Rule: item.Rule, ForwardRule: item.ForwardRule, DockerRule: item.DockerRule, - Error: err.Error(), - }) + firewallRuleSyncTaskID = taskItem.TaskID + go func() { + defer func() { + firewallRuleSyncTaskMu.Lock() + if firewallRuleSyncTaskID == taskItem.TaskID { + firewallRuleSyncTaskID = "" + } + firewallRuleSyncTaskMu.Unlock() + }() + if err := taskItem.Execute(); err != nil && global.LOG != nil { + global.LOG.Errorf("firewall sync task %s failed: %v", taskItem.TaskID, err) + } + }() + return dto.FirewallRuleSyncResult{ + Subsystem: "system", + TargetProvider: request.TargetProvider, + TaskID: taskItem.TaskID, + Queued: true, + }, nil } -func firewallDiffPreview[T any](subsystem string, target filter.Provider, desired []firewallsync.Desired[T, dto.FirewallRuleSyncItem], actual []T, key func(T) string, actualItem func(T) dto.FirewallRuleSyncItem) dto.FirewallRuleSyncPreview { - preview := dto.FirewallRuleSyncPreview{Subsystem: subsystem, TargetProvider: target, Items: make([]dto.FirewallRuleSyncItem, 0)} - for _, item := range firewallsync.Diff(desired, actual, key, actualItem) { - row := item.Payload - row.Status, row.ReasonCode, row.Reason = item.Status, item.ReasonCode, item.Reason - preview.Add(row) +func verifyForwardingRuleSync(target forwarding.Adapter, desired []forwarding.Rule) error { + actual, err := target.List() + if err != nil { + return fmt.Errorf("verify synchronized forwarding rules: %w", err) + } + actual, err = normalizeForwardingRuntimeRules(actual) + if err != nil { + return fmt.Errorf("verify synchronized forwarding rules: %w", err) } - return preview + if !firewallRuleStatesEqual(actual, desired, func(rule forwarding.Rule) string { return rule.Identity() }) { + return fmt.Errorf("verify synchronized forwarding rules: target rules do not match the database") + } + return nil } -func firewallSyncResult(preview dto.FirewallRuleSyncPreview, cause error, executed bool) dto.FirewallRuleSyncResult { - result := dto.FirewallRuleSyncResult{Subsystem: preview.Subsystem, TargetProvider: preview.TargetProvider} - for _, item := range preview.Items { - if item.Status != firewallsync.StatusRemove { - result.Total++ - } - switch item.Status { - case firewallsync.StatusExisting: - result.Skipped++ - case firewallsync.StatusBlocked: - appendDatabaseSyncFailure(&result, item, errors.New(item.Reason)) - case firewallsync.StatusReady: - if executed { - if cause != nil { - appendDatabaseSyncFailure(&result, item, cause) - } else { - result.Succeeded++ - } - } - case firewallsync.StatusRemove: - if executed && cause == nil { - result.Removed++ +func reconcileDockerFirewall(backend string, policies []dockerfirewall.Policy, runtime dockerfirewall.Runtime, inventory dockerfirewall.PolicyInventory) 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{dockerfirewall.FamilyIPv4, dockerfirewall.FamilyIPv6} { + status := runtime.Status(family) + if status.Reason == dockerfirewall.ReasonInspectFailed { + return fmt.Errorf("inspect Docker firewall target %s for %s failed", backend, family) } + initialized = initialized || status.Initialized + } + if initialized { + return runtime.ReplacePolicies(nil, inventory) + } + 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, inventory) + } else { + if needsBind { + err = runtime.Bind() + } + if err == nil { + err = runtime.ReplacePolicies(policies, inventory) + } + } + 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 result + return nil +} + +func ReconcileDockerPortGuardBestEffort(ctx context.Context) { + if err := ReconcileDockerPortGuard(ctx); err != nil { + global.LOG.Warnf("reconcile Docker port guard failed, err: %v", err) + } } diff --git a/agent/app/service/firewall_utils.go b/agent/app/service/firewall_utils.go new file mode 100644 index 000000000000..33058c42b7be --- /dev/null +++ b/agent/app/service/firewall_utils.go @@ -0,0 +1,3948 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/netip" + "os" + "reflect" + "slices" + "sort" + "strconv" + "strings" + "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/app/task" + "github.com/1Panel-dev/1Panel/agent/buserr" + "github.com/1Panel-dev/1Panel/agent/constant" + "github.com/1Panel-dev/1Panel/agent/global" + "github.com/1Panel-dev/1Panel/agent/i18n" + "github.com/1Panel-dev/1Panel/agent/utils/cmd" + "github.com/1Panel-dev/1Panel/agent/utils/controller" + "github.com/1Panel-dev/1Panel/agent/utils/docker" + "github.com/1Panel-dev/1Panel/agent/utils/firewall" + dockerfirewall "github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard" + "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" + "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" + "github.com/1Panel-dev/1Panel/agent/utils/firewall/iptables_helper" + "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" + lifecycleproviders "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle/providers" + "github.com/1Panel-dev/1Panel/agent/utils/firewall/nftables_helper" + firewallsync "github.com/1Panel-dev/1Panel/agent/utils/firewall/sync" + containertypes "github.com/docker/docker/api/types/container" + "github.com/docker/docker/api/types/system" + "github.com/docker/docker/client" + "github.com/google/uuid" + "gorm.io/gorm" +) + +const ( + firewallTaskHost = "FirewallTaskHost" + firewallTaskForwarding = "FirewallTaskForwarding" + firewallTaskDocker = "FirewallTaskDocker" +) + +const fail2BanRestoreWithFirewallMarker = "/run/1panel_fail2ban_restore_with_firewall" + +type firewallLifecycleClient struct{ lifecycle.Client } + +type managedMutationRequest struct { + Stored model.FirewallRule + Before filter.FirewallRule + After filter.FirewallRule + RuleSet filter.RuleSet + Locator filter.Locator + AdapterOperation filter.ChangeOperation + Runtime filter.Adapter +} + +type firewallDockerRestartError struct { + Err error +} + +type firewallCompletedOperationError struct { + Operation string + Err error +} + +type firewallVerification struct { + RuleSet filter.RuleSet + Matched bool +} + +type firewallRuleBatchItem struct { + snapshot filter.RuleSet + change filter.RuleChange +} + +type firewallSyncRule struct { + dto.FirewallRuleSyncItem + desired filter.DesiredRule + observed *filter.ObservedRule + done bool +} + +type forwardingRuleSyncCandidate struct { + rule forwarding.Rule + err error +} + +type forwardingInventoryItem struct { + ID uint + Rule forwarding.Rule + IsDesired bool + IsRuntime bool +} + +type observedInventoryCandidate struct { + rule filter.ObservedRule + ruleKey string + instanceKey string + claimed bool +} + +type firewallRuleCollisionIndex map[string][]filter.Action + +type firewallSyncDesired[T any] struct { + Value T + Payload dto.FirewallRuleSyncItem + Err error +} + +func LoadPanelPort() string { + if !global.IsMaster { + return global.CONF.Base.Port + } + var portSetting model.Setting + _ = global.CoreDB.Where("key = ?", "ServerPort").First(&portSetting).Error + return portSetting.Value +} + +func updateSystemAccessPortWhitelist(ctx context.Context, serviceType string, ports []string) error { + firewallWhitelistMu.Lock() + defer firewallWhitelistMu.Unlock() + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + return global.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + entries, err := loadPortWhitelistSetting(tx) + if err != nil { + return err + } + for index := range entries { + if entries[index].Type != serviceType { + continue + } + if len(ports) == 0 { + return fmt.Errorf("firewall whitelist %s requires a port", serviceType) + } + entries[index].Port = ports[0] + } + entries, err = firewall.ValidatePortWhitelist(entries) + if err != nil { + return err + } + value, err := json.Marshal(entries) + if err != nil { + return err + } + err = tx.Where("key = ?", constant.FirewallPortWhiteList).Assign(map[string]interface{}{"value": string(value)}).FirstOrCreate(&model.Setting{Key: constant.FirewallPortWhiteList}).Error + if err != nil { + return err + } + return nil + }) +} + +func loadPortWhitelistSetting(db *gorm.DB) ([]firewall.PortWhitelist, error) { + var setting model.Setting + if err := db.Where("key = ?", constant.FirewallPortWhiteList).First(&setting).Error; errors.Is(err, gorm.ErrRecordNotFound) { + setting.Value = constant.FirewallPortWhiteListValue + } else if err != nil { + return nil, err + } + var rules []firewall.PortWhitelist + err := json.Unmarshal([]byte(setting.Value), &rules) + return rules, err +} + +func NewSelectedSystemFirewallClient() (lifecycle.Client, error) { + provider, _ := settingRepo.GetValueByKey(constant.FirewallSystemBackendKey) + provider = strings.TrimSpace(provider) + client, err := lifecycle.NewClient(provider) + if err != nil { + return nil, err + } + if provider == "" { + _ = settingRepo.UpdateOrCreate(constant.FirewallSystemBackendKey, client.Name()) + } + return client, nil +} + +func loadFirewallInitStatus(provider, tab string) (bool, bool, error) { + switch provider { + case constant.FirewallProviderNftables: + return nftables_helper.LoadInitStatus(tab) + case constant.FirewallProviderIptables: + return iptables_helper.LoadInitStatus(tab) + default: + return false, false, fmt.Errorf("unsupported firewall provider: %s", provider) + } +} + +func loadSystemFirewallOverview(provider, chainGroup string) (dto.FirewallSubsystemStatus, error) { + var status dto.FirewallSubsystemStatus + var ipv4Err, ipv6Err error + status.IPv4, ipv4Err = loadSystemFirewallFamilyInfo(provider, constant.FirewallFamilyIPv4) + status.IPv6, ipv6Err = loadSystemFirewallFamilyInfo(provider, constant.FirewallFamilyIPv6) + if chainGroup != "base" { + return status, nil + } + if ipv4Err != nil { + return status, ipv4Err + } + if provider == constant.FirewallProviderIptables { + status.IsInit, status.IsBind = status.IPv4.Initialized, status.IPv4.Bound + return status, nil + } + if !status.IPv4.Initialized { + return status, nil + } + if !status.IPv4.Bound { + status.IsInit = true + return status, nil + } + if ipv6Err != nil { + return status, ipv6Err + } + if status.IPv6.Initialized { + status.IsInit, status.IsBind = true, status.IPv6.Bound + } + return status, nil +} + +func loadSystemFirewallFamilyInfo(provider, family string) (dto.FirewallBackendFamilyStatus, error) { + if provider == constant.FirewallProviderIptables && family == constant.FirewallFamilyIPv6 { + commands, err := lifecycle.ResolveIptablesCommands() + if err != nil || !commands.IPv6Available() { + return dto.FirewallBackendFamilyStatus{Reason: dockerfirewall.ReasonCommandMissing}, nil + } + } + initialized, bound, err := loadSystemFirewallFamilyStatus(provider, family) + return dto.FirewallBackendFamilyStatus{ + Available: err == nil, + Initialized: initialized, + Bound: bound, + }, err +} + +func loadSystemFirewallFamilyStatus(provider, family string) (bool, bool, error) { + switch provider { + case constant.FirewallProviderIptables: + return iptables_helper.LoadFamilyInitStatus(family, "base") + case constant.FirewallProviderNftables: + return nftables_helper.LoadFamilyInitStatus(filter.Family(family), "base") + default: + return false, false, fmt.Errorf("unsupported firewall provider %q", provider) + } +} + +func (s *FirewallService) updateRuleOrder(ctx context.Context, ruleUUID string, targetPosition *int64, priority *int, description *string) error { + if (targetPosition == nil) == (priority == nil) { + return fmt.Errorf("%w: provide either position or priority", filter.ErrInvalidRule) + } + if ruleUUID == "" { + return fmt.Errorf("%w: rule UUID is required", repo.ErrFirewallPersistenceInvalid) + } + stored, before, runtime, err := s.loadManagedRule(ctx, ruleUUID) + if err != nil { + return err + } + var snapshot filter.RuleSet + var observed filter.ObservedRule + if runtime.Provider() == filter.ProviderNftables { + snapshot, err = readMutableFirewallRules(runtime, ctx, before.Rule.Scope) + if err != nil { + return err + } + observed, err = managedFirewallObserved(snapshot, before) + if err != nil { + return err + } + } + capabilities, err := runtime.Capabilities(ctx) + if err != nil { + return err + } + after := before.Rule + adapterOperation := filter.ChangeReorder + switch { + case capabilities.ExplicitPosition || capabilities.OwnedChains: + if targetPosition == nil || *targetPosition < 1 { + return fmt.Errorf("%w: target position is required", filter.ErrInvalidRule) + } + if runtime.Provider() == filter.ProviderNftables { + if err := validateFirewallRulePosition(snapshot, before.Rule, *targetPosition); err != nil { + return err + } + } + after.OrderIndex = targetPosition + case capabilities.ExplicitPriority: + if before.Rule.NativeKind != filter.NativeKindRichRule { + return fmt.Errorf("%w: only rich rules support explicit priority", filter.ErrUnsupportedScope) + } + if priority == nil { + return fmt.Errorf("%w: priority is required", filter.ErrInvalidRule) + } + after.Priority = priority + adapterOperation = filter.ChangeUpdate + default: + return fmt.Errorf("%w: provider does not support rule reordering", filter.ErrUnsupportedScope) + } + if description != nil { + after.Description = strings.TrimSpace(*description) + } + after, err = prepareFirewallBackendRule(ctx, runtime, after) + if err != nil { + return err + } + metadataOnly, err := isFirewallMetadataOnlyUpdate(before.Rule, after, observed.Locator) + if err != nil { + return err + } + if metadataOnly { + return s.updateRuleDescription(ctx, stored.UUID, after.Description) + } + if runtime.Provider() != filter.ProviderNftables { + return s.replaceManagedRule(ctx, preparedManagedUpdate{Stored: stored, Before: before, After: after, Runtime: runtime}) + } + if err := filter.GuardMutation(observed); err != nil { + return err + } + return s.executeManagedMutation(ctx, managedMutationRequest{ + Stored: stored, Before: before.Rule, After: after, RuleSet: snapshot, Locator: observed.Locator, + AdapterOperation: adapterOperation, Runtime: runtime, + }) +} + +func (s *FirewallService) loadManagedRule(ctx context.Context, ruleUUID string) (model.FirewallRule, filter.DesiredRule, filter.Adapter, error) { + stored, err := s.rules.GetByUUID(ctx, ruleUUID) + if err != nil { + return model.FirewallRule{}, filter.DesiredRule{}, nil, err + } + if stored.Origin != constant.FirewallRuleOriginCreated && stored.Origin != constant.FirewallRuleOriginAdopted { + return model.FirewallRule{}, filter.DesiredRule{}, nil, + fmt.Errorf("%w: only created or adopted rules can be changed", filter.ErrInvalidRule) + } + selected, err := s.selectedProviderForStoredRule(ctx, stored) + if err != nil { + return model.FirewallRule{}, filter.DesiredRule{}, nil, err + } + if err := checkFirewallRuleWhitelistProtection(selected, stored); err != nil { + return model.FirewallRule{}, filter.DesiredRule{}, nil, err + } + runtime, err := s.resolveRuntime(ctx, selected) + if err != nil { + return model.FirewallRule{}, filter.DesiredRule{}, nil, err + } + desiredRules, err := compileStoredFirewallRules(ctx, stored, runtime) + if err != nil { + return model.FirewallRule{}, filter.DesiredRule{}, nil, err + } + if len(desiredRules) != 1 { + return model.FirewallRule{}, filter.DesiredRule{}, nil, + fmt.Errorf("%w: policy %q expands to %d target rules and cannot be edited atomically", filter.ErrUnsupportedScope, ruleUUID, len(desiredRules)) + } + return stored, desiredRules[0], runtime, nil +} + +func (s *FirewallService) selectedProviderForStoredRule(ctx context.Context, _ model.FirewallRule) (filter.Provider, error) { + if s.selectedProvider != nil { + return s.selectedProvider(ctx) + } + if s.adapters != nil { + providers := make([]filter.Provider, 0, len(s.adapters)) + for provider := range s.adapters { + providers = append(providers, provider) + } + if len(providers) == 1 { + return providers[0], nil + } + } + return "", fmt.Errorf("%w: selected provider is unavailable", filter.ErrProviderUnavailable) +} + +func (s *FirewallService) resolveRuntime(ctx context.Context, provider filter.Provider) (filter.Adapter, error) { + if s.selectedProvider != nil { + selected, err := s.selectedProvider(ctx) + if err != nil { + return nil, err + } + if selected != provider { + return nil, fmt.Errorf("%w: selected provider is %s, requested %s", filter.ErrProviderUnavailable, selected, provider) + } + } + return s.firewallAdapter(provider) +} + +func (s *FirewallService) firewallAdapter(provider filter.Provider) (filter.Adapter, error) { + if s.adapters != nil { + if client := s.adapters[provider]; client != nil { + return client, nil + } + return nil, fmt.Errorf("%w: %s", filter.ErrAdapterUnavailable, provider) + } + switch provider { + case filter.ProviderUFW: + return filterufw.NewAdapter(), nil + case filter.ProviderFirewalld: + return filterfirewalld.NewAdapter(), nil + case filter.ProviderIptables: + return filteriptables.NewAdapter(), nil + case filter.ProviderNftables: + return filternftables.NewAdapter(), nil + default: + return nil, fmt.Errorf("%w: %s", filter.ErrAdapterUnavailable, provider) + } +} + +func checkFirewallRuleWhitelistProtection(provider filter.Provider, record model.FirewallRule) error { + ports, err := loadFirewallPortWhiteList() + if err != nil { + return err + } + rules, err := expandStoredFirewallRule(record, provider) + if err != nil { + return err + } + whitelist := filter.NewPortWhitelistIndex(ports) + for _, rule := range rules { + if whitelist.Matches(rule) { + return filter.ErrProtectedRule + } + } + return nil +} + +func loadFirewallPortWhiteList() ([]firewall.PortWhitelist, error) { + ports, err := loadPortWhitelistSetting(global.DB) + if err != nil { + return nil, err + } + return firewall.ValidatePortWhitelist(ports) +} + +func expandStoredFirewallRule(rule model.FirewallRule, 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.ProviderUFW && strings.EqualFold(strings.TrimSpace(base.Protocol), "all") && + strings.TrimSpace(base.SourcePort) == "" && strings.TrimSpace(base.DestinationPort) != "" { + base.Protocol = "tcp/udp" + } + 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 := false, false + for _, address := range []string{base.SourceAddress, base.DestinationAddress} { + address = strings.TrimSpace(address) + if address == "" { + continue + } + if strings.Contains(address, ":") { + hasIPv6 = true + } else { + hasIPv4 = true + } + } + 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 compileStoredFirewallRules(ctx context.Context, stored model.FirewallRule, client filter.Adapter) ([]filter.DesiredRule, error) { + rules, err := expandStoredFirewallRule(stored, client.Provider()) + if err != nil { + return nil, err + } + capabilities, err := client.Capabilities(ctx) + if err != nil { + return nil, err + } + result := make([]filter.DesiredRule, 0, len(rules)) + for ordinal, rule := range rules { + prepared, err := prepareFirewallBackendRule(ctx, client, rule) + if err != nil { + return nil, err + } + ruleKey, err := filter.RuleKey(prepared) + if err != nil { + return nil, err + } + prepared.UUID = stored.UUID + if ordinal > 0 { + suffix := ruleKey + if len(suffix) > 12 { + suffix = suffix[:12] + } + prepared.UUID = fmt.Sprintf("%s-%d-%s", stored.UUID, ordinal+1, suffix) + } + desired := filter.DesiredRule{UUID: stored.UUID, Rule: prepared, RuleKey: ruleKey, Origin: filter.RuleOrigin(stored.Origin)} + if capabilities.Marker { + desired.Marker = "1panel-rule:" + prepared.UUID + } + result = append(result, desired) + } + return result, nil +} + +func prepareFirewallBackendRule(ctx context.Context, client filter.Adapter, rule filter.FirewallRule) (filter.FirewallRule, error) { + if preparer, ok := client.(filter.RulePreparer); ok { + var err error + rule, err = preparer.PrepareRule(rule) + if err != nil { + return filter.FirewallRule{}, err + } + } + if checker, ok := client.(filter.RuleChecker); ok { + if err := checker.CheckRule(ctx, rule); err != nil { + return filter.FirewallRule{}, err + } + } + return rule, nil +} + +func readMutableFirewallRules(client filter.Adapter, ctx context.Context, scope filter.Scope) (filter.RuleSet, error) { + snapshots, err := readMutableFirewallRuleScopes(client, ctx, []filter.Scope{scope}) + if err != nil { + return filter.RuleSet{}, err + } + return snapshots[0], nil +} + +func firewallScopeReadGroups(scopes []filter.Scope) [][]filter.Scope { + groups := make([][]filter.Scope, 0) + indexes := make(map[string]int) + for _, scope := range scopes { + scope = scope.Normalize() + key := scope.Key() + if scope.Provider == filter.ProviderUFW { + key = string(scope.Provider) + } + if scope.Provider == filter.ProviderIptables || scope.Provider == filter.ProviderNftables { + key = string(scope.Provider) + ":" + string(scope.Family) + ":" + scope.Table + } + if index, ok := indexes[key]; ok { + groups[index] = append(groups[index], scope) + } else { + indexes[key] = len(groups) + groups = append(groups, []filter.Scope{scope}) + } + } + return groups +} + +func listFirewallRuleScopes(client filter.Adapter, ctx context.Context, scopes []filter.Scope) ([]filter.RuleSet, error) { + unique := make([]filter.Scope, 0, len(scopes)) + seen := make(map[string]bool) + for _, scope := range scopes { + scope = scope.Normalize() + if err := scope.ValidateMVP(); err != nil { + return nil, err + } + if !seen[scope.Key()] { + seen[scope.Key()] = true + unique = append(unique, scope) + } + } + if len(unique) == 0 { + return nil, nil + } + if reader, ok := client.(filter.MultiScopeReader); ok { + snapshots, err := reader.ListRuleScopes(ctx, unique) + if err != nil { + return nil, err + } + if len(snapshots) != len(unique) { + return nil, filter.ErrInventoryUnavailable + } + return snapshots, nil + } + snapshots := make([]filter.RuleSet, 0, len(unique)) + for _, scope := range unique { + snapshot, err := client.ListRules(ctx, scope) + if err != nil { + return nil, err + } + snapshots = append(snapshots, snapshot) + } + return snapshots, nil +} + +func readMutableFirewallRuleScopes(client filter.Adapter, ctx context.Context, scopes []filter.Scope) ([]filter.RuleSet, error) { + snapshots, err := readFirewallRuleScopes(client, ctx, scopes) + if err != nil { + return nil, err + } + for _, snapshot := range snapshots { + for _, notice := range snapshot.Notices { + if notice.Code == filter.ScopeNoticeManagedScopeInactive || notice.Code == filter.ScopeNoticeManagedScopeMissing { + return nil, fmt.Errorf("%w: managed firewall scope is unavailable", filter.ErrProviderUnavailable) + } + } + } + return snapshots, nil +} + +func readFirewallRules(client filter.Adapter, ctx context.Context, scope filter.Scope) (filter.RuleSet, error) { + snapshot, err := client.ListRules(ctx, scope) + if err != nil { + return filter.RuleSet{}, err + } + ports, err := loadFirewallPortWhiteList() + if err != nil { + return filter.RuleSet{}, err + } + return filter.ProtectRuleSet(snapshot, ports) +} + +func firewallRulesByMarker(rules []filter.ObservedRule) map[string][]filter.ObservedRule { + index := make(map[string][]filter.ObservedRule) + for _, rule := range rules { + if rule.Marker != "" { + index[rule.Marker] = append(index[rule.Marker], rule) + } + } + return index +} + +func managedFirewallObserved(snapshot filter.RuleSet, desired filter.DesiredRule) (filter.ObservedRule, error) { + matches := make([]filter.ObservedRule, 0, 1) + for _, observed := range snapshot.Rules { + if desired.Marker != "" { + if observed.Marker != desired.Marker { + continue + } + } else if desired.ObservedInstanceKey != "" { + key, err := filter.InstanceKey(observed) + if err != nil || key != desired.ObservedInstanceKey { + continue + } + } else { + key, err := firewallInventoryRuleKey(observed.Rule) + wanted, wantErr := firewallInventoryRuleKey(desired.Rule) + if err != nil || wantErr != nil || key != wanted { + continue + } + } + matches = append(matches, observed) + } + if len(matches) != 1 { + return filter.ObservedRule{}, filter.ErrRuleStale + } + observed := matches[0] + if observed.Protected || desired.Protected { + return filter.ObservedRule{}, filter.ErrProtectedRule + } + if observed.Rule.Scope.Provider != filter.ProviderFirewalld && observed.Persistence != "" && observed.Persistence != filter.PersistenceStatusConverged { + return filter.ObservedRule{}, filter.ErrRuleStale + } + if desired.Marker != "" && observed.ParseStatus == filter.ParseStatusOpaque { + position := observed.Rule.OrderIndex + observed.Rule = desired.Rule + observed.Rule.OrderIndex = position + observed.ParseStatus = filter.ParseStatusSupported + observed.UncertainFields = nil + } else { + expected := desired.Rule + if expected.Scope.Provider == filter.ProviderFirewalld { + expected.Priority = observed.Rule.Priority + expected.NativeKind = observed.Rule.NativeKind + expected.OrderBucket = observed.Rule.OrderBucket + } + if !filter.ObservedRuleMatchesExpected(observed, expected) { + return filter.ObservedRule{}, filter.ErrRuleStale + } + } + return observed, nil +} + +func (s *FirewallService) updateRuleDescription(ctx context.Context, ruleUUID, description string) error { + stored, err := s.rules.GetByUUID(ctx, ruleUUID) + if err != nil { + return err + } + selected, err := s.selectedProviderForStoredRule(ctx, stored) + if err != nil { + return err + } + if err := checkFirewallRuleWhitelistProtection(selected, stored); err != nil { + return err + } + if stored.Origin != constant.FirewallRuleOriginCreated && stored.Origin != constant.FirewallRuleOriginAdopted { + return fmt.Errorf("%w: only created or adopted rules can be changed", filter.ErrInvalidRule) + } + description = strings.TrimSpace(description) + if stored.Description == description { + return nil + } + return s.rules.UpdateWithRevision(ctx, stored.UUID, stored.Revision, map[string]interface{}{"description": description}) +} + +func (s *FirewallService) replaceManagedRule(ctx context.Context, prepared preparedManagedUpdate) error { + runtime := prepared.Runtime + snapshot := filter.RuleSet{Scope: prepared.Before.Rule.Scope} + remove, err := runtime.BuildCommands(snapshot, []filter.RuleChange{{ + Operation: filter.ChangeDelete, Before: &prepared.Before.Rule, CommandOnly: true, + }}) + if err != nil { + return err + } + create, err := runtime.BuildCommands(snapshot, []filter.RuleChange{{ + Operation: filter.ChangeCreate, After: &prepared.After, CommandOnly: true, Append: prepared.After.OrderIndex == nil, + }}) + if err != nil { + return err + } + updates, err := firewallRuleSemanticUpdates(prepared.After) + if err != nil { + return err + } + if runtime.Provider() == filter.ProviderFirewalld { + updates["priority"] = prepared.After.Priority + } + remove.CommandOnly, create.CommandOnly = true, true + if err := runtime.RunCommands(ctx, remove); err != nil { + return err + } + saveCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + err = s.rules.UpdateWithRevision(saveCtx, prepared.Stored.UUID, prepared.Stored.Revision, updates) + cancel() + if err != nil { + return err + } + runErr := runtime.RunCommands(ctx, create) + if err := errors.Join(runErr, persistFirewallRules(ctx, runtime, create)); err != nil { + return buserr.WithDetail("ErrFirewallRuleSavedApplyFailed", err.Error(), err) + } + return nil +} + +func (s *FirewallService) executeManagedMutation(ctx context.Context, request managedMutationRequest) (err error) { + defer func() { + if err != nil && global.LOG != nil { + global.LOG.Errorf("update firewall rule %s failed: %v", request.Stored.UUID, err) + } + }() + before, after := request.Before, request.After + appendRule := after.Scope.Provider == filter.ProviderUFW && after.OrderIndex != nil && *after.OrderIndex == maxObservedFirewallPosition(request.RuleSet) + backendPlan, err := request.Runtime.BuildCommands(request.RuleSet, []filter.RuleChange{{ + Operation: request.AdapterOperation, Before: &before, After: &after, + Locator: &request.Locator, Append: appendRule, CommandOnly: true, + }}) + if err != nil { + return err + } + backendPlan.CommandOnly = true + updates, err := firewallRuleSemanticUpdates(request.After) + if err != nil { + return err + } + if len(backendPlan.Rules) == 1 && len(backendPlan.Rules[0].Commands) == 2 { + commands := backendPlan.Rules[0] + backendPlan.Rules[0].Commands = commands.Commands[:1] + backendPlan.Rules[0].RollbackCommands = commands.RollbackCommands[:1] + if err := request.Runtime.RunCommands(ctx, backendPlan); err != nil { + return err + } + if after.Scope.Provider == filter.ProviderFirewalld { + updates["priority"] = after.Priority + } + saveCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + err = s.rules.UpdateWithRevision(saveCtx, request.Stored.UUID, request.Stored.Revision, updates) + cancel() + if err != nil { + return err + } + backendPlan.Rules[0].Commands = commands.Commands[1:] + backendPlan.Rules[0].RollbackCommands = commands.RollbackCommands[1:] + runErr := request.Runtime.RunCommands(ctx, backendPlan) + if err := errors.Join(runErr, persistFirewallRules(ctx, request.Runtime, backendPlan)); err != nil { + return buserr.WithDetail("ErrFirewallRuleSavedApplyFailed", err.Error(), err) + } + return nil + } + runErr := request.Runtime.RunCommands(ctx, backendPlan) + if runErr != nil && request.Runtime.Provider() != filter.ProviderFirewalld { + return runErr + } + if err := errors.Join(runErr, persistFirewallRules(ctx, request.Runtime, backendPlan)); err != nil { + return err + } + sameContent, err := filter.SameRuleContent(request.Before, request.After) + if err != nil { + return err + } + if sameContent && request.Stored.Description == request.After.Description { + return nil + } + return s.rules.UpdateWithRevision(ctx, request.Stored.UUID, request.Stored.Revision, updates) +} + +func isFirewallPolicyIncompatible(err error) bool { + return errors.Is(err, filter.ErrInvalidRule) || errors.Is(err, filter.ErrUnsupportedScope) || + errors.Is(err, filter.ErrInvalidScope) || errors.Is(err, filter.ErrCompositeRule) +} + +func maxObservedFirewallPosition(snapshot filter.RuleSet) int64 { + maximum := int64(snapshot.LastPosition) + for _, observed := range snapshot.Rules { + if observed.Locator.Position != nil && int64(*observed.Locator.Position) > maximum { + maximum = int64(*observed.Locator.Position) + } + } + return maximum +} + +func applyFirewallChanges(client filter.Adapter, ctx context.Context, snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, firewallVerification, error) { + plan, err := client.BuildCommands(snapshot, changes) + if err != nil { + return filter.CommandBatch{}, firewallVerification{}, err + } + err = client.RunCommands(ctx, plan) + if err == nil { + err = persistFirewallRules(ctx, client, plan) + } + if err != nil { + return plan, firewallVerification{}, err + } + verification, err := verifyFirewallCommands(ctx, client, plan) + if plan.Provider == filter.ProviderUFW && len(plan.Rules) == 1 && plan.Rules[0].Operation == filter.ChangeAdopt && (err != nil || !verification.Matched) { + if err == nil { + err = filter.ErrVerificationFailed + } + err = buserr.WithDetail("ErrUFWRuleAdopt", err.Error(), err) + } + if err != nil { + if plan.CreatesOnly() { + return plan, verification, err + } + return plan, verification, rollbackFirewallPlan(ctx, client, plan, err) + } + if !verification.Matched && !plan.CreatesOnly() { + if rollbackErr := restoreFirewallCommands(client, ctx, plan); rollbackErr != nil { + return plan, verification, errors.Join(filter.ErrVerificationFailed, rollbackErr) + } + } + return plan, verification, nil +} + +func persistFirewallRules(ctx context.Context, client filter.Adapter, commands filter.CommandBatch) error { + saver, ok := client.(filter.RuleSaver) + if !ok { + return nil + } + err := saver.SaveRules(ctx, commands.Scope) + if err == nil || commands.CreatesOnly() || commands.CommandOnly { + return err + } + if rollbackErr := restoreFirewallCommands(client, ctx, commands); rollbackErr != nil { + return errors.Join(err, rollbackErr) + } + return err +} + +func persistFirewallRuleBatches(ctx context.Context, client filter.Adapter, plans []filter.CommandBatch) []error { + failures := make([]error, len(plans)) + saver, ok := client.(filter.RuleSaver) + if !ok { + return failures + } + saved := make(map[string]error) + for index, plan := range plans { + key := plan.Scope.Key() + if plan.Provider == filter.ProviderNftables { + key = string(plan.Provider) + } + err, exists := saved[key] + if !exists { + saveCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + err = saver.SaveRules(saveCtx, plan.Scope) + cancel() + saved[key] = err + } + failures[index] = err + } + return failures +} + +func restoreFirewallCommands(client filter.Adapter, ctx context.Context, plan filter.CommandBatch) error { + ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + return client.Rollback(ctx, plan) +} + +func readFirewallCommandResults(ctx context.Context, client filter.Adapter, plans ...filter.CommandBatch) ([]filter.RuleSet, map[string]map[string][]filter.ObservedRule, error) { + scopes := make([]filter.Scope, 0, len(plans)) + for _, plan := range plans { + scopes = append(scopes, plan.Scope) + if plan.Provider == filter.ProviderUFW { + related := plan.Scope + if related.Family == filter.FamilyIPv4 { + related.Family = filter.FamilyIPv6 + } else { + related.Family = filter.FamilyIPv4 + } + scopes = append(scopes, related) + } + } + snapshots, err := listFirewallRuleScopes(client, ctx, scopes) + if err != nil { + return nil, nil, err + } + byIdentity := make(map[string]map[string][]filter.ObservedRule, len(snapshots)) + for _, snapshot := range snapshots { + if snapshot.Scope.Provider == filter.ProviderFirewalld { + canonical := make(map[string][]filter.ObservedRule) + for _, observed := range snapshot.Rules { + canonical[observed.Locator.Canonical] = append(canonical[observed.Locator.Canonical], observed) + } + byIdentity[snapshot.Scope.Key()] = canonical + } else { + byIdentity[snapshot.Scope.Key()] = firewallRulesByMarker(snapshot.Rules) + } + } + return snapshots, byIdentity, nil +} + +func verifyFirewallCommands(ctx context.Context, client filter.Adapter, plans ...filter.CommandBatch) (firewallVerification, error) { + snapshots, byIdentity, err := readFirewallCommandResults(ctx, client, plans...) + if err != nil { + return firewallVerification{}, err + } + result := firewallVerification{Matched: true} + for _, plan := range plans { + result.Matched = false + for _, snapshot := range snapshots { + if snapshot.Scope.Key() != plan.Scope.Key() { + continue + } + result.RuleSet = snapshot + result.Matched = firewallCommandsMatch(plan, snapshot, snapshots, byIdentity) + break + } + if !result.Matched { + return result, nil + } + } + return result, nil +} + +func firewallCommandsMatch(commands filter.CommandBatch, current filter.RuleSet, all []filter.RuleSet, byIdentity map[string]map[string][]filter.ObservedRule) bool { + for _, command := range commands.Rules { + if commands.Provider == filter.ProviderFirewalld { + previous, matched := 0, 0 + canonical := byIdentity[current.Scope.Key()] + if command.Previous != nil { + previous = len(canonical[command.Previous.Locator.Canonical]) + } + for _, observed := range canonical[command.Expected.Locator.Canonical] { + if observed.Persistence == filter.PersistenceStatusConverged { + want, wantErr := filter.RuleKey(command.Expected.Rule) + got, gotErr := filter.RuleKey(observed.Rule) + if wantErr == nil && gotErr == nil && want == got { + matched++ + } + } + } + if command.Operation == filter.ChangeDelete { + if previous != 0 { + return false + } + continue + } + if command.Operation == filter.ChangeUpdate && command.Previous != nil && command.Previous.Locator.Canonical != command.Expected.Locator.Canonical && previous != 0 { + return false + } + if matched != 1 { + return false + } + continue + } + markerMatches, semanticMatches := 0, 0 + positionMatches := true + requiresPosition := command.Expected.Locator.Position != nil && (commands.Provider == filter.ProviderUFW || command.Operation == filter.ChangeCreate || commands.Provider == filter.ProviderIptables && (command.Operation == filter.ChangeReorder || command.Operation == filter.ChangeUpdate && command.Expected.Rule.OrderIndex != nil)) + if requiresPosition { + positionMatches = false + } + for _, rules := range all { + if commands.Provider != filter.ProviderUFW && rules.Scope.Key() != commands.Scope.Key() { + continue + } + candidates := byIdentity[rules.Scope.Key()][command.Expected.Marker] + if command.Expected.Marker == "" { + candidates = rules.Rules + } + for _, observed := range candidates { + if observed.Marker != command.Expected.Marker { + continue + } + markerMatches++ + if rules.Scope.Key() != commands.Scope.Key() { + continue + } + same := false + if commands.Provider == filter.ProviderUFW { + same = observed.ParseStatus == filter.ParseStatusOpaque || filter.ObservedRuleMatchesExpected(observed, command.Expected.Rule) + } else { + want, wantErr := filter.RuleKey(command.Expected.Rule) + got, gotErr := filter.RuleKey(observed.Rule) + same = wantErr == nil && gotErr == nil && want == got + } + if same { + semanticMatches++ + } + if observed.Locator.Position != nil && command.Expected.Locator.Position != nil && *observed.Locator.Position == *command.Expected.Locator.Position { + positionMatches = true + } + } + } + if command.Operation == filter.ChangeDelete { + if markerMatches != 0 { + return false + } + continue + } + if markerMatches != 1 || semanticMatches != 1 || !positionMatches { + return false + } + } + return true +} + +func rollbackFirewallPlan(ctx context.Context, runtime filter.Adapter, plan filter.CommandBatch, cause error) error { + if runtime == nil { + return cause + } + if err := restoreFirewallCommands(runtime, ctx, plan); err != nil { + return errors.Join(cause, fmt.Errorf("rollback applied firewall plan: %w", err)) + } + return cause +} + +func firewallRuleSemanticUpdates(rule filter.FirewallRule) (map[string]interface{}, error) { + record, err := firewallRuleFromDomain(rule) + if err != nil { + return nil, err + } + return map[string]interface{}{ + "family": record.Family, "protocol": record.Protocol, + "source_address": record.SourceAddress, "source_port": record.SourcePort, + "destination_address": record.DestinationAddress, "destination_port": record.DestinationPort, + "interface": record.Interface, "connection_states": record.ConnectionStates, "action": record.Action, + "description": record.Description, "compatibility_error": "", + }, nil +} + +func validateFirewallRulePosition(snapshot filter.RuleSet, rule filter.FirewallRule, target int64) error { + if target < 1 { + return fmt.Errorf("%w: target position must be positive", filter.ErrInvalidRule) + } + 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 := maxObservedFirewallPosition(snapshot) + if target > maximum { + return fmt.Errorf("%w: target position %d is out of range 1-%d", filter.ErrInvalidRule, target, maximum) + } + return nil +} + +func positionBounds(snapshot filter.RuleSet) (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 isFirewallMetadataOnlyUpdate(before, after filter.FirewallRule, locator filter.Locator) (bool, error) { + beforeKey, err := filter.RuleKey(before) + if err != nil { + return false, err + } + afterKey, err := filter.RuleKey(after) + if err != nil { + return false, err + } + if beforeKey != afterKey { + return false, nil + } + if after.Scope.Provider == filter.ProviderFirewalld { + return true, nil + } + return locator.Position != nil && after.OrderIndex != nil && *after.OrderIndex == int64(*locator.Position), nil +} + +func (index firewallRuleCollisionIndex) Check(rule filter.FirewallRule) error { + key, err := filter.RuleMatchKey(rule) + if err != nil { + return err + } + for _, action := range index[key] { + if err := checkCollisionActions(rule.Action, action); err != nil { + return err + } + } + return nil +} + +func checkCollisionActions(requested, existing filter.Action) error { + if requested == existing { + return fmt.Errorf("%w: equivalent rule already exists", filter.ErrRuleOperation) + } + if filter.OppositeActions(requested, existing) { + return filter.ErrRuleConflict + } + return nil +} + +func firewallRuleCollisions(stored []model.FirewallRule, provider filter.Provider) (firewallRuleCollisionIndex, error) { + identities := make(firewallRuleCollisionIndex, len(stored)) + for _, candidate := range stored { + rules, err := expandStoredFirewallRule(candidate, provider) + if err != nil { + continue + } + for _, rule := range rules { + if err := identities.Add(rule); err != nil { + return nil, err + } + } + } + return identities, nil +} + +func (index firewallRuleCollisionIndex) Add(rule filter.FirewallRule) error { + key, err := filter.RuleMatchKey(rule) + if err != nil { + return err + } + index[key] = append(index[key], rule.Action) + return nil +} + +func (s *FirewallService) restoreStoredFirewallRules(ctx context.Context, provider filter.Provider, t *task.Task) error { + result, err := s.syncRules(ctx, "", dto.FirewallRuleSyncRequest{TargetProvider: provider}, t) + if err != nil { + return fmt.Errorf("restore database firewall rules: %w", err) + } + failures := make([]error, 0, len(result.Errors)) + for _, failure := range result.Errors { + failures = append(failures, fmt.Errorf("rule %s: %s", failure.SourceUUID, failure.Error)) + } + return errors.Join(failures...) +} + +func executeFirewallRuleBatches(ctx context.Context, runtime filter.Adapter, items []firewallRuleBatchItem, t *task.Task, record func(int, error)) { + groups := make([][]int, 0) + byScope := make(map[string]int) + for index, item := range items { + key := string(item.change.Operation) + ":" + item.snapshot.Scope.Key() + if runtime.Provider() == filter.ProviderUFW && item.change.Operation == filter.ChangeDelete { + key = string(item.change.Operation) + } + group, exists := byScope[key] + if !exists { + group = len(groups) + byScope[key] = group + groups = append(groups, nil) + } + groups[group] = append(groups[group], index) + } + var plans []filter.CommandBatch + var completed [][]int + _, savesRules := runtime.(filter.RuleSaver) + for _, group := range groups { + if items[group[0]].change.Operation == filter.ChangeDelete { + sort.SliceStable(group, func(i, j int) bool { + left, right := items[group[i]].change.Locator, items[group[j]].change.Locator + if left == nil || right == nil || left.Position == nil || right.Position == nil { + return false + } + return *left.Position > *right.Position + }) + } + for start := 0; start < len(group); { + end := start + 1 + if runtime.Provider() != filter.ProviderUFW { + end = min(start+filter.MaxAtomicExpansion, len(group)) + } + batch := group[start:end] + start = end + changes := make([]filter.RuleChange, 0, len(batch)) + for _, index := range batch { + change := items[index].change + change.CommandOnly = true + changes = append(changes, change) + } + err := ctx.Err() + if err == nil { + var plan filter.CommandBatch + plan, err = runtime.BuildCommands(items[batch[0]].snapshot, changes) + if err == nil { + if len(batch) > 1 && changes[0].Operation == filter.ChangeCreate && t != nil { + t.Log(i18n.GetMsgWithMap("FirewallCreateBatchStep", map[string]interface{}{"backend": runtime.Provider(), "count": len(batch)})) + } + plan.CommandOnly = true + err = runtime.RunCommands(ctx, plan) + if err == nil && savesRules { + plans = append(plans, plan) + completed = append(completed, batch) + continue + } + } + } + for _, index := range batch { + record(index, err) + } + } + } + failures := persistFirewallRuleBatches(ctx, runtime, plans) + for index, batch := range completed { + for _, item := range batch { + record(item, failures[index]) + } + } +} + +func (s *FirewallService) syncRules(ctx context.Context, _ string, request dto.FirewallRuleSyncRequest, t *task.Task) (result dto.FirewallRuleSyncResult, err error) { + if request.SourceProvider != "" || request.ResetSource { + return result, fmt.Errorf("%w: system synchronization reads rules from the database", filter.ErrInvalidRule) + } + if err := s.checkSelectedProvider(ctx, request.TargetProvider); err != nil { + return result, err + } + whitelistErr := s.SyncPortWhitelist(ctx) + if t != nil { + t.LogWithStatus(i18n.GetMsgByKey("FirewallSyncWhitelistStep"), whitelistErr) + } + if whitelistErr != nil { + return result, whitelistErr + } + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + result = dto.FirewallRuleSyncResult{Subsystem: "system", TargetProvider: request.TargetProvider} + if err := ctx.Err(); err != nil { + return result, err + } + runtime, rules, snapshots, err := s.loadFirewallSyncRules(ctx, request) + if err != nil { + return result, err + } + + created, removed, unexecuted := 0, 0, 0 + failedRemovals := make(map[string]error) + record := func(operation string, rule *firewallSyncRule, cause error, skipped bool) { + item := rule.FirewallRuleSyncItem + if operation == "TaskDelete" && rule.observed != nil { + item.Rule = &rule.observed.Rule + } + switch { + case skipped: + result.Skipped++ + unexecuted++ + rule.done = true + case cause != nil: + appendDatabaseSyncFailure(&result, item, cause) + rule.done = true + if operation == "TaskDelete" { + failedRemovals[rule.SourceUUID] = cause + } + case operation == "TaskDelete": + removed++ + if rule.Status == firewallsync.StatusRemove { + result.Removed++ + rule.done = true + } + case operation == "TaskCreate": + created++ + result.Succeeded++ + rule.done = true + default: + result.Skipped++ + rule.done = true + } + if t == nil { + return + } + label := fmt.Sprintf("%s %s", i18n.GetMsgByKey(operation), item.SourceUUID) + if r := item.Rule; r != nil { + label += fmt.Sprintf(" [%s] %s %s:%s -> %s:%s %s", r.Scope.Key(), r.Protocol, r.SourceAddress, r.SourcePort, r.DestinationAddress, r.DestinationPort, r.Action) + } + if skipped { + t.Logf("%s %s: %v", label, i18n.GetMsgByKey("FirewallCreateRuleSkipped"), cause) + } else if operation == task.TaskSync && cause == nil { + t.Logf("%s %s", label, i18n.GetMsgByKey("FirewallSyncRuleUnchanged")) + } else { + t.LogWithStatus(label, cause) + } + } + defer func() { + if t != nil { + t.Log(i18n.GetMsgWithMap("FirewallSyncOperationsResult", map[string]interface{}{"created": created, "removed": removed, "failed": result.Failed, "skipped": unexecuted, "unchanged": result.Skipped - unexecuted})) + } + }() + for _, rule := range rules { + if rule.Status != firewallsync.StatusRemove { + result.Total++ + } + if rule.Status == firewallsync.StatusBlocked { + record("TaskSync", rule, errors.New(rule.Reason), false) + } + } + for _, rule := range rules { + if rule.Status == firewallsync.StatusExisting { + record("TaskSync", rule, nil, false) + } + } + byScope := make(map[string]filter.RuleSet, len(snapshots)) + for _, snapshot := range snapshots { + byScope[snapshot.Scope.Key()] = snapshot + } + for _, operation := range []filter.ChangeOperation{filter.ChangeDelete, filter.ChangeCreate} { + name := "TaskCreate" + if operation == filter.ChangeDelete { + name = "TaskDelete" + } + pending := make([]*firewallSyncRule, 0) + items := make([]firewallRuleBatchItem, 0) + for _, rule := range rules { + if rule.done { + continue + } + if operation == filter.ChangeDelete && rule.observed == nil || operation == filter.ChangeCreate && rule.Status != firewallsync.StatusReady { + continue + } + if operation == filter.ChangeCreate { + var cause error + for id, err := range failedRemovals { + if id == rule.SourceUUID || strings.HasPrefix(id, rule.SourceUUID+"-") { + cause = err + break + } + } + if cause != nil { + record(name, rule, cause, true) + continue + } + } + item := firewallRuleBatchItem{snapshot: filter.RuleSet{Scope: rule.Rule.Scope}} + if operation == filter.ChangeDelete { + item.snapshot = byScope[rule.Rule.Scope.Key()] + item.change, err = firewallDeleteChange(*rule.observed, rule.desired) + if err != nil { + record(name, rule, err, false) + continue + } + if runtime.Provider() == filter.ProviderUFW { + item.snapshot.Rules = []filter.ObservedRule{*rule.observed} + } + } else { + after := *rule.Rule + after.OrderIndex = nil + item.change = filter.RuleChange{Operation: operation, After: &after, Append: true} + } + pending = append(pending, rule) + items = append(items, item) + } + executeFirewallRuleBatches(ctx, runtime, items, t, func(index int, failure error) { + record(name, pending[index], failure, false) + }) + } + return result, nil +} + +func (s *FirewallService) checkSelectedProvider(ctx context.Context, requested filter.Provider) error { + selected, err := s.selectedProvider(ctx) + if err != nil { + return err + } + if selected != requested { + return fmt.Errorf("%w: selected provider is %s, requested %s", filter.ErrProviderUnavailable, selected, requested) + } + return nil +} + +func (s *FirewallService) saveFirewallRule(ctx context.Context, record *model.FirewallRule) error { + if err := s.rules.Create(ctx, record); err != nil { + message := "FirewallCreateRulePersistenceFailed" + if record.Origin == constant.FirewallRuleOriginAdopted { + message = "FirewallAdoptRulePersistenceFailed" + } + return fmt.Errorf("%s: %w", i18n.GetMsgByKey(message), err) + } + return nil +} + +func (s *FirewallService) createRules(ctx context.Context, request dto.FirewallRuleCreate, t *task.Task) (result dto.FirewallRuleCreateResponse, taskErr error) { + var firstFailure error + defer func() { + if t != nil { + t.Log(i18n.GetMsgWithMap("FirewallCreateRulesResult", map[string]interface{}{ + "succeeded": result.Succeeded, "failed": result.Failed, "skipped": result.Skipped, + })) + } + if taskErr == nil { + taskErr = firstFailure + } + }() + type createItem struct { + index, part, count int + request dto.FirewallRuleCreateItem + stored model.FirewallRule + } + describe := func(rule filter.FirewallRule) string { + return fmt.Sprintf("%s %s %s:%s -> %s:%s %s", rule.Scope.Family, rule.Protocol, + rule.SourceAddress, rule.SourcePort, rule.DestinationAddress, rule.DestinationPort, rule.Action) + } + record := func(item createItem, status string, err error) { + rule := item.request.Rule + label := fmt.Sprintf("[%d/%d]", item.index+1, len(request.Items)) + if item.count > 1 { + label += fmt.Sprintf("[%d/%d]", item.part+1, item.count) + } + label += fmt.Sprintf(" %s %s", rule.Scope.Provider, describe(rule)) + switch status { + case "succeeded": + result.Succeeded++ + if t != nil { + t.LogSuccess(label) + } + case "failed": + if firstFailure == nil { + firstFailure = err + } + result.Failed++ + if t != nil { + t.LogFailedWithErr(label, err) + } + case "skipped": + result.Skipped++ + if t != nil { + t.Logf("%s %s: %v", label, i18n.GetMsgByKey("FirewallCreateRuleSkipped"), err) + } + } + if err != nil { + result.Errors = append(result.Errors, dto.FirewallRuleCreateFailure{ + Index: item.index, Status: status, Rule: rule, Error: err.Error(), + }) + } + } + selected, err := s.selectedProvider(ctx) + if err != nil { + for index, item := range request.Items { + record(createItem{index: index, request: item}, "skipped", err) + } + return result, err + } + runtime, err := s.firewallAdapter(selected) + if err != nil { + for index, item := range request.Items { + record(createItem{index: index, request: item}, "skipped", err) + } + return result, err + } + items := make([]createItem, 0, len(request.Items)) + var stop error + for index, item := range request.Items { + if stop == nil { + stop = ctx.Err() + } + if stop != nil { + record(createItem{index: index, request: item}, "skipped", stop) + continue + } + item.Rule.OrderIndex = nil + rules, err := expandFirewallCreateRule(item, selected) + if err != nil { + record(createItem{index: index, request: item}, "failed", err) + continue + } + if item.SourceKind == constant.FirewallRuleSourceImported && t != nil { + t.Log(i18n.GetMsgWithMap("FirewallImportRuleConversion", map[string]interface{}{ + "index": index + 1, "total": len(request.Items), "source": item.Rule.Scope.Provider, + "target": selected, "rule": describe(item.Rule), "count": len(rules), + })) + } + for part, rule := range rules { + child := item + child.Rule = rule + entry := createItem{index: index, part: part, count: len(rules), request: child} + prepared, err := prepareFirewallCreateRule(ctx, runtime, child) + if err != nil { + record(entry, "failed", err) + continue + } + entry.request = prepared + items = append(items, entry) + } + } + stored, err := s.rules.List(ctx) + if err != nil { + for _, item := range items { + record(item, "failed", err) + } + return result, err + } + identities, err := firewallRuleCollisions(stored, selected) + if err != nil { + return result, err + } + valid := items[:0] + for _, item := range items { + if stop == nil { + stop = ctx.Err() + } + if stop != nil { + record(item, "skipped", stop) + continue + } + rule := item.request.Rule + if err := identities.CheckDuplicate(rule); errors.Is(err, filter.ErrRuleOperation) { + record(item, "skipped", err) + continue + } else if err != nil { + record(item, "failed", err) + continue + } + item.stored, err = firewallRuleModelForCreate(rule, item.request, constant.FirewallRuleOriginCreated) + if err != nil { + record(item, "failed", err) + continue + } + item.stored.UUID = uuid.NewString() + rule.UUID = item.stored.UUID + item.request.Rule = rule + if err := identities.Add(rule); err != nil { + record(item, "failed", err) + continue + } + valid = append(valid, item) + } + changes := make([]firewallRuleBatchItem, 0, len(valid)) + for index := range valid { + rule := &valid[index].request.Rule + changes = append(changes, firewallRuleBatchItem{ + snapshot: filter.RuleSet{Scope: rule.Scope}, + change: filter.RuleChange{Operation: filter.ChangeCreate, After: rule, Append: true}, + }) + } + executeFirewallRuleBatches(ctx, runtime, changes, t, func(index int, err error) { + item := valid[index] + if err != nil { + if errors.Is(err, filterfirewalld.ErrAlreadyEnabled) { + record(item, "skipped", err) + } else if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + record(item, "skipped", err) + stop = err + } else { + record(item, "failed", fmt.Errorf("%s: %w", i18n.GetMsgByKey("FirewallCreateRuleExecutionFailed"), err)) + } + return + } + persistCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + if err := s.saveFirewallRule(persistCtx, &item.stored); err != nil { + record(item, "failed", err) + } else { + record(item, "succeeded", nil) + } + }) + return result, stop +} + +func validateFirewallCreateBatch(request dto.FirewallRuleCreate, provider filter.Provider) error { + count := 0 + if len(request.Items) > filter.MaxAtomicExpansion { + return fmt.Errorf("create or import at most %d rules per batch (after expansion)", filter.MaxAtomicExpansion) + } + for _, item := range request.Items { + rules, expandErr := expandFirewallCreateRule(item, provider) + if errors.Is(expandErr, filter.ErrExpansionLimit) { + return fmt.Errorf("create or import at most %d rules per batch (after expansion)", filter.MaxAtomicExpansion) + } + count += len(rules) + if count > filter.MaxAtomicExpansion { + return fmt.Errorf("create or import at most %d rules per batch (after expansion)", filter.MaxAtomicExpansion) + } + } + return nil +} + +func expandFirewallCreateRule(item dto.FirewallRuleCreateItem, provider filter.Provider) ([]filter.FirewallRule, error) { + if item.SourceKind == constant.FirewallRuleSourceImported { + source := item.Rule.Scope.Normalize().Provider + if source == "" { + source = provider + } + rules, err := filter.ExpandAtomicRules(applySelectedProviderScopeDefaults(item.Rule, source)) + if err != nil { + return nil, err + } + var converted []filter.FirewallRule + for _, sourceRule := range rules { + policy, err := firewallRuleFromDomain(sourceRule) + if err != nil { + return nil, err + } + targetRules, err := expandStoredFirewallRule(policy, provider) + if err != nil { + return nil, err + } + converted = append(converted, targetRules...) + } + return converted, nil + } + rule := applySelectedProviderScopeDefaults(item.Rule, provider) + if provider == filter.ProviderUFW && strings.TrimSpace(rule.DestinationPort) != "" { + protocol := strings.ToLower(strings.TrimSpace(rule.Protocol)) + if protocol == "" || protocol == "all" || protocol == "any" { + return filter.ExpandAtomicRules(rule) + } + } + return []filter.FirewallRule{rule}, nil +} + +func applySelectedProviderScopeDefaults(rule filter.FirewallRule, selected filter.Provider) filter.FirewallRule { + scope := rule.Scope + if scope.Provider == "" { + scope.Provider = selected + } + if scope.Provider != selected { + return rule + } + if scope.Direction == "" { + scope.Direction = filter.DirectionInput + } + if scope.Family == "" { + switch { + case strings.EqualFold(strings.TrimSpace(rule.Protocol), "icmpv6"), strings.Contains(rule.SourceAddress, ":"), strings.Contains(rule.DestinationAddress, ":"): + scope.Family = filter.FamilyIPv6 + case selected == filter.ProviderFirewalld: + scope.Family = filter.FamilyInet + default: + scope.Family = filter.FamilyIPv4 + } + } + switch selected { + case filter.ProviderIptables, filter.ProviderNftables: + if scope.Table == "" { + scope.Table = "filter" + } + if scope.Chain == "" { + scope.Chain = filter.IptablesInputChain + } + case filter.ProviderFirewalld: + if scope.Zone == "" { + scope.Zone = filter.FirewalldInputZone + } + case filter.ProviderUFW: + if scope.Chain == "" { + scope.Chain = filter.UFWInputChain + } + } + rule.Scope = scope + return rule +} + +func prepareFirewallCreateRule(ctx context.Context, runtime filter.Adapter, request dto.FirewallRuleCreateItem) (dto.FirewallRuleCreateItem, error) { + selected := runtime.Provider() + rule, err := filter.NormalizeRule(applySelectedProviderScopeDefaults(request.Rule, selected)) + if err != nil { + return dto.FirewallRuleCreateItem{}, err + } + if rule.Scope.Provider != selected { + return dto.FirewallRuleCreateItem{}, fmt.Errorf("%w: selected provider is %s", filter.ErrInvalidRule, selected) + } + rule, err = prepareFirewallBackendRule(ctx, runtime, rule) + if err != nil { + return dto.FirewallRuleCreateItem{}, err + } + rule.UUID = "" + request.Rule = rule + if request.SourceKind == "" { + request.SourceKind = constant.FirewallRuleSourceUser + } + return request, nil +} + +func readFirewallRuleScopes(client filter.Adapter, ctx context.Context, scopes []filter.Scope) ([]filter.RuleSet, error) { + snapshots, err := listFirewallRuleScopes(client, ctx, scopes) + if err != nil || len(snapshots) == 0 { + return snapshots, err + } + ports, err := loadFirewallPortWhiteList() + if err != nil { + return nil, err + } + for index := range snapshots { + snapshots[index], err = filter.ProtectRuleSet(snapshots[index], ports) + if err != nil { + return nil, err + } + } + return snapshots, nil +} + +func observedFirewallRuleCollisionIndex(snapshot filter.RuleSet) (firewallRuleCollisionIndex, error) { + index := make(firewallRuleCollisionIndex, len(snapshot.Rules)) + for _, observed := range snapshot.Rules { + if observed.ParseStatus != filter.ParseStatusSupported { + continue + } + if err := index.Add(observed.Rule); err != nil { + return nil, err + } + } + return index, nil +} + +func firewallRuleFromDomain(rule filter.FirewallRule) (model.FirewallRule, error) { + normalized, err := filter.NormalizeRule(rule) + if err != nil { + return model.FirewallRule{}, err + } + switch normalized.NativeKind { + case "", filter.NativeKindRule, filter.NativeKindZonePort, filter.NativeKindRichRule, filter.NativeKindUFWRule: + default: + return model.FirewallRule{}, fmt.Errorf("%w: native rule %q cannot be stored as a provider-neutral policy", filter.ErrUnsupportedScope, normalized.NativeKind) + } + record := model.FirewallRule{ + Family: string(normalized.Scope.Family), + Protocol: normalized.Protocol, + SourceAddress: normalized.SourceAddress, + SourcePort: normalized.SourcePort, + DestinationAddress: normalized.DestinationAddress, + DestinationPort: normalized.DestinationPort, + Interface: normalized.Interface, + ConnectionStates: strings.Join(normalized.ConnectionStates, ","), + Action: string(normalized.Action), + Description: normalized.Description, + } + return record, nil +} + +func sortFirewallRules(rules []model.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 firewallRuleModelForCreate(rule filter.FirewallRule, request dto.FirewallRuleCreateItem, origin string) (model.FirewallRule, error) { + record, err := firewallRuleFromDomain(rule) + if err != nil { + return model.FirewallRule{}, err + } + record.Origin = origin + record.Owner = strings.TrimSpace(request.SourceKind) + if sourceID := strings.TrimSpace(request.SourceID); sourceID != "" { + record.Owner += ":" + sourceID + } + return record, nil +} + +func firewallTaskName(operation, subsystem, backend string) string { + name := i18n.GetMsgByKey(subsystem) + if backend != "" { + name += " · " + backend + } + key := "FirewallRule" + operation + if operation == task.TaskExec { + key = "FirewallTaskInitialize" + } + return i18n.GetMsgWithMap(key, map[string]interface{}{"name": name}) +} + +func closeUnstartedFirewallTask(t *task.Task) { + if cancel, ok := global.LoadTaskCancel(t.TaskID); ok { + cancel() + } + global.RemoveTaskCancel(t.TaskID) + if closer, ok := t.Logger.Out.(io.Closer); ok { + _ = closer.Close() + } +} + +func whitelistRules(provider filter.Provider, ports, required []firewall.SystemPort) []filter.FirewallRule { + rules := make([]filter.FirewallRule, 0, len(ports)+len(required)) + for _, port := range required { + rule := firewall.RuleForSystemPort(provider, firewall.SystemPort(port)) + if provider == filter.ProviderIptables || provider == filter.ProviderNftables { + rule.Scope.Chain = filter.BasicBeforeChain + } + rules = append(rules, rule) + } + for _, port := range ports { + rules = append(rules, firewall.RuleForSystemPort(provider, firewall.SystemPort(port))) + } + return rules +} + +func customWhitelist(entries []firewall.PortWhitelist) []firewall.PortWhitelist { + result := make([]firewall.PortWhitelist, 0, len(entries)) + for _, entry := range entries { + if entry.Type == "" { + result = append(result, entry) + } + } + return result +} + +func (s *FirewallService) loadFirewallSyncRules(ctx context.Context, request dto.FirewallRuleSyncRequest) (filter.Adapter, []*firewallSyncRule, []filter.RuleSet, error) { + if request.SourceProvider != "" || request.ResetSource { + return nil, nil, nil, fmt.Errorf("%w: system synchronization reads rules from the database", filter.ErrInvalidRule) + } + if err := s.checkSelectedProvider(ctx, request.TargetProvider); err != nil { + return nil, nil, nil, err + } + runtime, err := s.firewallAdapter(request.TargetProvider) + if err != nil { + return nil, nil, nil, err + } + stored, err := s.rules.List(ctx) + if err != nil { + return nil, nil, nil, err + } + sortFirewallRules(stored, request.TargetProvider) + ports, err := loadFirewallPortWhiteList() + if err != nil { + return nil, nil, nil, err + } + required, err := firewall.RequiredPortWhitelist(ports) + if err != nil { + return nil, nil, nil, err + } + whitelistDesired := whitelistRules(request.TargetProvider, firewall.ExpandPortWhitelist(customWhitelist(ports)), firewall.ExpandPortWhitelist(required)) + rules := make([]*firewallSyncRule, 0, len(stored)) + preservedMarkers := make(map[string]bool) + compileFailed := false + whitelist := filter.NewPortWhitelistIndex(ports) + for _, record := range stored { + desired, preserved, err := s.compileRestorableFirewallRules(ctx, record, runtime, required) + if err != nil { + compileFailed = true + rules = append(rules, &firewallSyncRule{FirewallRuleSyncItem: dto.FirewallRuleSyncItem{SourceUUID: record.UUID, Status: firewallsync.StatusBlocked, Reason: err.Error()}}) + continue + } + for _, rule := range preserved { + preservedMarkers[rule.Marker] = true + } + for _, native := range desired { + rule := native.Rule + native.Protected = whitelist.Matches(rule) + rules = append(rules, &firewallSyncRule{desired: native, FirewallRuleSyncItem: dto.FirewallRuleSyncItem{SourceUUID: record.UUID, Rule: &rule}}) + } + } + identities := make(firewallRuleCollisionIndex, len(rules)) + for _, existing := range rules { + if existing.Rule != nil { + if err := identities.Add(*existing.Rule); err != nil { + return nil, nil, nil, err + } + } + } + for _, candidate := range whitelistDesired { + prepared, err := prepareFirewallCreateRule(ctx, runtime, dto.FirewallRuleCreateItem{Rule: candidate}) + if err != nil { + return nil, nil, nil, err + } + rule := prepared.Rule + if err := identities.CheckDuplicate(rule); errors.Is(err, filter.ErrRuleOperation) { + continue + } else if err != nil { + return nil, nil, nil, err + } + if err := identities.Add(rule); err != nil { + return nil, nil, nil, err + } + key, err := filter.RuleKey(rule) + if err != nil { + return nil, nil, nil, err + } + rule.UUID = uuid.NewSHA1(uuid.NameSpaceOID, []byte(key)).String() + rules = append(rules, &firewallSyncRule{ + desired: filter.DesiredRule{UUID: rule.UUID, Rule: rule, RuleKey: key, Origin: filter.RuleOriginCreated, Protected: true}, + FirewallRuleSyncItem: dto.FirewallRuleSyncItem{SourceUUID: rule.UUID, Rule: &rule}, + }) + } + scopes := filter.ManagedInputScopes(request.TargetProvider) + if compileFailed { + scopes = slices.DeleteFunc(scopes, func(scope filter.Scope) bool { + return !slices.ContainsFunc(rules, func(rule *firewallSyncRule) bool { + return rule.Rule != nil && rule.Rule.Scope.Key() == scope.Key() + }) + }) + } + snapshots := make([]filter.RuleSet, 0, len(scopes)) + for _, group := range firewallScopeReadGroups(scopes) { + current, err := readMutableFirewallRuleScopes(runtime, ctx, group) + if errors.Is(err, filter.ErrFamilyUnavailable) { + for _, rule := range rules { + if rule.Rule != nil && slices.ContainsFunc(group, func(scope filter.Scope) bool { return scope.Key() == rule.Rule.Scope.Key() }) { + rule.Status, rule.Reason = firewallsync.StatusBlocked, err.Error() + } + } + continue + } + if err != nil { + return nil, nil, nil, err + } + snapshots = append(snapshots, current...) + } + for _, snapshot := range snapshots { + scope := snapshot.Scope + desired := make([]filter.DesiredRule, 0) + byUUID := make(map[string]*firewallSyncRule) + for _, rule := range rules { + if rule.Rule != nil && rule.Rule.Scope.Key() == scope.Key() { + desired = append(desired, rule.desired) + byUUID[rule.desired.Rule.UUID] = rule + } + } + inventory, err := mergeFirewallInventory(filter.InventoryMergeInput{Observed: snapshot.Rules, Desired: desired}) + if err != nil { + return nil, nil, nil, err + } + for _, item := range inventory { + if item.Desired == nil { + if compileFailed || item.Observed == nil || !strings.HasPrefix(item.Observed.Marker, "1panel-rule:") || preservedMarkers[item.Observed.Marker] { + continue + } + observed := item.Observed + if observed.Rule.Scope.Chain == filter.BasicBeforeChain { + continue + } + rule := &firewallSyncRule{observed: observed, FirewallRuleSyncItem: dto.FirewallRuleSyncItem{ + SourceUUID: strings.TrimPrefix(observed.Marker, "1panel-rule:"), Rule: &observed.Rule, Status: firewallsync.StatusRemove, + ReasonCode: firewallsync.ReasonManagedOnlyInTarget, Reason: firewallSyncReasonMessage(firewallsync.ReasonManagedOnlyInTarget), + }} + if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque { + rule.Status, rule.ReasonCode = firewallsync.StatusBlocked, firewallsync.ReasonUnsafeRemoval + rule.Reason = firewallSyncReasonMessage(rule.ReasonCode) + } + rules = append(rules, rule) + continue + } + rule := byUUID[item.Desired.Rule.UUID] + rule.observed = item.Observed + if item.Match == filter.InventoryMatchExact && item.Observed != nil && scope.Provider == filter.ProviderFirewalld { + rule.Rule.Priority = item.Observed.Rule.Priority + rule.Rule.NativeKind = item.Observed.Rule.NativeKind + rule.Rule.OrderBucket = item.Observed.Rule.OrderBucket + rule.desired.Rule = *rule.Rule + rule.desired.RuleKey, err = filter.RuleKey(*rule.Rule) + if err != nil { + return nil, nil, nil, err + } + } + divergent := item.Observed != nil && item.Observed.Persistence != "" && item.Observed.Persistence != filter.PersistenceStatusConverged + switch { + case item.Match == filter.InventoryMatchExact && !divergent: + rule.Status, rule.Reason = firewallsync.StatusExisting, "rule already matches database policy" + case item.Match == filter.InventoryMatchMissing || item.Match == filter.InventoryMatchChanged || item.Match == filter.InventoryMatchExact: + rule.Status, rule.Reason = firewallsync.StatusReady, "target rule differs from database policy" + if item.Observed != nil && item.Observed.Protected { + rule.Status, rule.Reason = firewallsync.StatusBlocked, filter.ErrProtectedRule.Error() + } + default: + rule.Status, rule.Reason = firewallsync.StatusBlocked, fmt.Sprintf("target rule cannot be synchronized: %s", item.Match) + } + } + } + return runtime, rules, snapshots, nil +} + +func (s *FirewallService) compileRestorableFirewallRules(ctx context.Context, stored model.FirewallRule, client filter.Adapter, required []firewall.PortWhitelist) (restorable, preserved []filter.DesiredRule, err error) { + provider := client.Provider() + compiled, err := compileStoredFirewallRules(ctx, stored, client) + if err != nil { + return nil, nil, err + } + if (provider != filter.ProviderIptables && provider != filter.ProviderNftables) || !strings.HasPrefix(stored.Owner, constant.FirewallRuleSourceSecurity+":"+constant.FirewallSystemAcceptedPortSourcePrefix) { + return compiled, nil, nil + } + requiredPorts := firewall.ExpandPortWhitelist(required) + for _, desired := range compiled { + covered := false + for _, port := range requiredPorts { + covered, err = filter.SameRuleContent(desired.Rule, firewall.RuleForSystemPort(provider, firewall.SystemPort(port))) + if err != nil { + return nil, nil, err + } + if covered { + break + } + } + if covered { + preserved = append(preserved, desired) + } else { + restorable = append(restorable, desired) + } + } + return restorable, preserved, nil +} + +func firewallInventoryRuleKey(rule filter.FirewallRule) (string, error) { + if rule.Scope.Provider == filter.ProviderFirewalld { + key, err := filter.RuleMatchKey(rule) + return key + ":" + string(rule.Action), err + } + return filter.RuleKey(rule) +} + +func mergeFirewallInventory(input filter.InventoryMergeInput) ([]filter.InventoryItem, error) { + candidates := make([]observedInventoryCandidate, len(input.Observed)) + byRuleKey := make(map[string][]int) + byInstanceKey := make(map[string][]int) + byMarker := make(map[string][]int) + bySemanticKey := make(map[string][]int) + byPartialRuleKey := make(map[string][]int) + for index, observed := range input.Observed { + candidate := observedInventoryCandidate{rule: observed} + if marker := strings.TrimSpace(candidate.rule.Marker); marker != "" { + markerKey := candidate.rule.Rule.Scope.Key() + "\x00" + marker + byMarker[markerKey] = append(byMarker[markerKey], index) + } + if observed.ParseStatus == filter.ParseStatusSupported { + normalized, err := filter.NormalizeRule(observed.Rule) + if err != nil { + return nil, fmt.Errorf("normalize observed firewall rule %d: %w", index, err) + } + candidate.rule.Rule = normalized + candidate.ruleKey, err = filter.RuleKey(normalized) + if err != nil { + return nil, err + } + key, err := firewallInventoryRuleKey(normalized) + if err != nil { + return nil, err + } + byRuleKey[key] = append(byRuleKey[key], index) + if instanceKey, err := filter.InstanceKey(candidate.rule); err == nil { + candidate.instanceKey = instanceKey + candidate.rule.InstanceKey = instanceKey + byInstanceKey[instanceKey] = append(byInstanceKey[instanceKey], index) + } + } + if observed.ParseStatus != filter.ParseStatusOpaque { + rule := candidate.rule.Rule + protocolUnknown, supportedFields := false, true + for _, field := range observed.UncertainFields { + if field != filter.ObservedFieldProtocol { + supportedFields = false + break + } + protocolUnknown = true + } + if supportedFields { + key := candidate.ruleKey + var err error + if protocolUnknown { + rule.Protocol = "tcp" + key, err = filter.RuleKey(rule) + } else if key == "" { + key, err = filter.RuleKey(rule) + } + if err == nil { + key += "\x00" + strings.TrimSpace(observed.Marker) + if protocolUnknown { + key = "protocol\x00" + key + byPartialRuleKey[key] = append(byPartialRuleKey[key], index) + } else if observed.ParseStatus != filter.ParseStatusSupported { + key = "exact\x00" + key + byPartialRuleKey[key] = append(byPartialRuleKey[key], index) + } else { + bySemanticKey[key] = append(bySemanticKey[key], index) + } + } + } + } + candidates[index] = candidate + } + + normalizedDesired := make([]filter.DesiredRule, 0, len(input.Desired)) + desiredMatches := make(map[int]int) + desiredMatchStates := make([]filter.InventoryMatch, 0, len(input.Desired)) + for _, desired := range input.Desired { + normalized, err := filter.NormalizeRule(desired.Rule) + if err != nil { + return nil, fmt.Errorf("normalize desired firewall rule %q: %w", desired.UUID, err) + } + desired.Rule = normalized + calculatedKey, err := filter.RuleKey(normalized) + if err != nil { + return nil, err + } + if desired.RuleKey != "" && desired.RuleKey != calculatedKey { + return nil, fmt.Errorf("%w: desired rule %q key does not match its semantics", filter.ErrInvalidRule, desired.UUID) + } + desired.RuleKey = calculatedKey + + match, matchState := findObservedInventoryMatch(desired, candidates, byRuleKey, byInstanceKey, byMarker, bySemanticKey, byPartialRuleKey) + normalizedIndex := len(normalizedDesired) + normalizedDesired = append(normalizedDesired, desired) + desiredMatchStates = append(desiredMatchStates, matchState) + if match >= 0 { + candidates[match].claimed = true + desiredMatches[match] = normalizedIndex + } + } + + items := make([]filter.InventoryItem, 0, len(candidates)+len(normalizedDesired)) + matchedDesired := make(map[int]struct{}, len(desiredMatches)) + for index := range candidates { + candidate := &candidates[index] + if desiredIndex, exists := desiredMatches[index]; exists { + desired := normalizedDesired[desiredIndex] + observed := candidate.rule + match := desiredMatchStates[desiredIndex] + if match == filter.InventoryMatchExact && observed.ParseStatus != filter.ParseStatusSupported { + orderIndex := observed.Rule.OrderIndex + observed.Rule = desired.Rule + observed.Rule.OrderIndex = orderIndex + observed.ParseStatus = filter.ParseStatusSupported + observed.UncertainFields = nil + } + displayRule := observed.Rule + displayRule.Description = desired.Rule.Description + state := inventoryStateForDesired(desired, match) + if observed.Protected { + state = filter.InventoryStateProtected + } else if observed.Persistence != "" && observed.Persistence != filter.PersistenceStatusConverged { + state = filter.InventoryStateDrifted + } + items = append(items, filter.InventoryItem{ + Rule: displayRule, + Observed: &observed, + Desired: &desired, + State: state, + Match: match, + }) + matchedDesired[desiredIndex] = struct{}{} + continue + } + observed := candidate.rule + state := filter.InventoryStateExternal + if observed.Protected { + state = filter.InventoryStateProtected + } else if _, protected := input.ProtectedObservedKeys[candidate.ruleKey]; protected { + state = filter.InventoryStateProtected + } + match := filter.InventoryMatchNone + if observed.ParseStatus != filter.ParseStatusSupported { + match = filter.InventoryMatchOpaque + } + items = append(items, filter.InventoryItem{Rule: observed.Rule, Observed: &observed, State: state, Match: match}) + } + for index, desired := range normalizedDesired { + if _, matched := matchedDesired[index]; matched { + continue + } + desiredCopy := desired + match := desiredMatchStates[index] + items = append(items, filter.InventoryItem{ + Rule: desired.Rule, + Desired: &desiredCopy, + State: inventoryStateForDesired(desired, match), + Match: match, + }) + } + return items, nil +} + +func findObservedInventoryMatch(desired filter.DesiredRule, candidates []observedInventoryCandidate, byRuleKey, byInstanceKey, byMarker, bySemanticKey, byPartialRuleKey map[string][]int) (int, filter.InventoryMatch) { + if marker := strings.TrimSpace(desired.Marker); marker != "" { + match, status := matchUnclaimedSemanticCandidate(desired, marker, candidates, bySemanticKey, byPartialRuleKey) + if status != filter.InventoryMatchMissing { + return match, status + } + markerKey := desired.Rule.Scope.Key() + "\x00" + marker + match, status = uniqueUnclaimedCandidate(byMarker, markerKey, candidates) + if match >= 0 && candidates[match].rule.ParseStatus != filter.ParseStatusOpaque && + !filter.ObservedRuleMatchesExpected(candidates[match].rule, desired.Rule) { + return match, filter.InventoryMatchChanged + } + if status != filter.InventoryMatchMissing { + return match, status + } + if desired.Origin == filter.RuleOriginAdopted { + match, status = matchUnclaimedSemanticCandidate(desired, "", candidates, bySemanticKey, byPartialRuleKey) + if status != filter.InventoryMatchMissing { + if match >= 0 { + return match, filter.InventoryMatchChanged + } + return match, status + } + } + legacyMarker := "1panel-rule:" + strings.TrimSpace(desired.UUID) + if legacyMarker != "1panel-rule:" && legacyMarker != marker { + match, status = matchUnclaimedSemanticCandidate(desired, legacyMarker, candidates, bySemanticKey, byPartialRuleKey) + if status != filter.InventoryMatchMissing { + if match >= 0 { + return match, filter.InventoryMatchChanged + } + return match, status + } + } + return match, status + } + var match int + if desired.ObservedInstanceKey != "" { + match = firstUnclaimedCandidate(byInstanceKey, desired.ObservedInstanceKey, candidates) + } else { + key, err := firewallInventoryRuleKey(desired.Rule) + if err != nil { + return -1, filter.InventoryMatchMissing + } + match = firstUnclaimedCandidate(byRuleKey, key, candidates) + } + if match < 0 { + return -1, filter.InventoryMatchMissing + } + return match, filter.InventoryMatchExact +} + +func firstUnclaimedCandidate(byKey map[string][]int, key string, candidates []observedInventoryCandidate) int { + indices := byKey[key] + for len(indices) > 0 && candidates[indices[0]].claimed { + indices = indices[1:] + } + if len(indices) == 0 { + delete(byKey, key) + return -1 + } + byKey[key] = indices + return indices[0] +} + +func uniqueUnclaimedCandidate(byKey map[string][]int, key string, candidates []observedInventoryCandidate) (int, filter.InventoryMatch) { + match := firstUnclaimedCandidate(byKey, key, candidates) + if match < 0 { + return -1, filter.InventoryMatchMissing + } + for _, index := range byKey[key][1:] { + if !candidates[index].claimed { + return -1, filter.InventoryMatchAmbiguous + } + } + return match, filter.InventoryMatchExact +} + +func matchUnclaimedSemanticCandidate(desired filter.DesiredRule, marker string, candidates []observedInventoryCandidate, bySemanticKey, byPartialRuleKey map[string][]int) (int, filter.InventoryMatch) { + key := desired.RuleKey + "\x00" + marker + if match := firstUnclaimedCandidate(bySemanticKey, key, candidates); match >= 0 { + return match, filter.InventoryMatchExact + } + if len(byPartialRuleKey) == 0 { + return -1, filter.InventoryMatchMissing + } + match, status := uniqueUnclaimedCandidate(byPartialRuleKey, "exact\x00"+key, candidates) + if status == filter.InventoryMatchAmbiguous { + return match, status + } + protocolIndependent := desired.Rule + protocolIndependent.Protocol = "tcp" + partialKey, err := filter.RuleKey(protocolIndependent) + if err != nil { + return match, status + } + partial, partialStatus := uniqueUnclaimedCandidate(byPartialRuleKey, "protocol\x00"+partialKey+"\x00"+marker, candidates) + if partialStatus == filter.InventoryMatchAmbiguous || (match >= 0 && partial >= 0) { + return -1, filter.InventoryMatchAmbiguous + } + if match >= 0 { + return match, status + } + return partial, partialStatus +} + +func inventoryStateForDesired(desired filter.DesiredRule, match filter.InventoryMatch) filter.InventoryState { + if match != filter.InventoryMatchExact { + return filter.InventoryStateDrifted + } + if desired.Protected { + return filter.InventoryStateProtected + } + switch desired.Origin { + case filter.RuleOriginAdopted: + return filter.InventoryStateAdopted + default: + return filter.InventoryStateManaged + } +} + +func firewallSyncReasonMessage(code firewallsync.ReasonCode) string { + switch code { + case firewallsync.ReasonAlreadyExists: + return "rule already exists in target backend" + case firewallsync.ReasonOnlyExistsInTarget: + return "rule exists only in target backend" + case firewallsync.ReasonManagedOnlyInTarget: + return "managed rule exists only in target backend" + case firewallsync.ReasonUnsafeRemoval: + return "managed runtime rule cannot be safely removed" + case firewallsync.ReasonReadOnlyRule: + return "read-only runtime rule is preserved but cannot be synchronized" + default: + return "" + } +} + +func appendDatabaseSyncFailure(result *dto.FirewallRuleSyncResult, item dto.FirewallRuleSyncItem, err error) { + if err == nil { + err = errors.New("database synchronization failed") + } + result.Failed++ + result.Errors = append(result.Errors, dto.FirewallRuleSyncFailure{ + SourceUUID: item.SourceUUID, + Rule: item.Rule, ForwardRule: item.ForwardRule, DockerRule: item.DockerRule, + Error: err.Error(), + }) +} + +func firewallDeleteChange(current filter.ObservedRule, desired filter.DesiredRule) (filter.RuleChange, error) { + if err := filter.GuardMutation(current); err != nil { + return filter.RuleChange{}, err + } + before := current.Rule + if before.UUID == "" && strings.HasPrefix(current.Marker, "1panel-rule:") { + before.UUID = strings.TrimSpace(strings.TrimPrefix(current.Marker, "1panel-rule:")) + } + if before.UUID == "" { + before.UUID = desired.Rule.UUID + } + return filter.RuleChange{Operation: filter.ChangeDelete, Before: &before, Locator: ¤t.Locator, UnmarkedAdopted: current.Marker == "" && desired.Origin == filter.RuleOriginAdopted}, nil +} + +func loadForwardingFirewallOverview(manager forwarding.Adapter) (dto.FirewallSubsystemStatus, error) { + ipv4, ipv4Err := loadForwardingFamilyInfo(manager, manager.Name(), constant.FirewallFamilyIPv4) + ipv6, ipv6Err := loadForwardingFamilyInfo(manager, manager.Name(), constant.FirewallFamilyIPv6) + initialized, bound, statusErr := ipv4.Initialized, ipv4.Bound, ipv4Err + if manager.Name() == constant.FirewallProviderNftables { + if statusErr == nil && initialized && bound { + initialized, bound, statusErr = ipv6.Initialized, ipv6.Bound, ipv6Err + } + } else { + statusErr = errors.Join(statusErr, ipv6Err) + if ipv6.Available { + initialized, bound = initialized && ipv6.Initialized, bound && ipv6.Bound + } + } + return dto.FirewallSubsystemStatus{IsInit: initialized, IsBind: bound, IPv4: ipv4, IPv6: ipv6}, statusErr +} + +func loadForwardingFamilyInfo(manager forwarding.Adapter, backend, family string) (dto.FirewallBackendFamilyStatus, error) { + initialized, bound, err := manager.FamilyStatus(family) + available := err == nil + if backend == constant.FirewallProviderIptables && family == constant.FirewallFamilyIPv6 { + commands, commandErr := lifecycle.ResolveIptablesCommands() + available = available && commandErr == nil && commands.IPv6Available() + } + return dto.FirewallBackendFamilyStatus{Available: available, Initialized: initialized, Bound: bound}, err +} + +func (s *ForwardingService) forwardingEnabled() (bool, error) { + if s.enabled != nil { + return s.enabled() + } + status, err := settingRepo.GetValueByKey(constant.FirewallForwardingInitializedKey) + return status == constant.StatusEnable, err +} + +func (s *ForwardingService) initializeForwarding(manager forwarding.Adapter) error { + if err := s.saveForwardingBackend(manager.Name()); err != nil { + return err + } + return manager.Enable() +} + +func (s *ForwardingService) saveForwardingBackend(backend string) error { + if s.persistBackend != nil { + return s.persistBackend(backend) + } + return settingRepo.UpdateOrCreate(constant.FirewallForwardingBackendKey, backend) +} + +func (s *ForwardingService) persistForwardingEnabled() error { + if s.markEnabled != nil { + return s.markEnabled() + } + return settingRepo.UpdateOrCreate(constant.FirewallForwardingInitializedKey, constant.StatusEnable) +} + +func recordForwardingSyncError(err error) { + forwardingSyncStateMu.Lock() + forwardingLastSyncErr = err + forwardingSyncStateMu.Unlock() +} + +func forwardingRulesFromModels(stored []model.ForwardingRule) []forwarding.Rule { + rules := make([]forwarding.Rule, 0, len(stored)) + for _, rule := range stored { + rules = append(rules, forwarding.Rule{ + Family: rule.Family, Protocol: rule.Protocol, Port: rule.Port, TargetIP: rule.TargetIP, + TargetPort: rule.TargetPort, Interface: rule.Interface, + }) + } + return rules +} + +func newForwardingService() *ForwardingService { + return &ForwardingService{ + clientFactory: newForwardingAdapter, + rules: forwardingRuleRepo, + markEnabled: func() error { + return settingRepo.UpdateOrCreate(constant.FirewallForwardingInitializedKey, constant.StatusEnable) + }, + persistBackend: func(backend string) error { + return settingRepo.UpdateOrCreate(constant.FirewallForwardingBackendKey, backend) + }, + } +} + +func newForwardingAdapter() (forwarding.Adapter, error) { + selected, _ := settingRepo.GetValueByKey(constant.FirewallForwardingBackendKey) + selected = strings.TrimSpace(selected) + if selected == "" { + selected = constant.FirewallProviderIptables + } + return newForwardingAdapterFor(selected) +} + +func newForwardingAdapterFor(backend string) (forwarding.Adapter, error) { + client, err := lifecycle.NewClient(backend) + if err != nil { + return nil, fmt.Errorf( + "%w: selected forwarding backend %s: %w", + errForwardingBackendUnavailable, backend, err, + ) + } + switch client.Name() { + case constant.FirewallProviderIptables: + return forwarding.NewIptables(client.Name()), nil + case constant.FirewallProviderNftables: + return forwarding.NewNftables(), nil + default: + return nil, errForwardingBackendUnavailable + } +} + +func ReconcileDockerPortGuard(ctx context.Context) error { + return NewIDockerPortGuardService().Reconcile(ctx) +} + +func (s *DockerPortGuardService) reconcileLocked(ctx context.Context) error { + persistedEnabled, err := dockerPortGuardPersistedEnabled() + if err != nil { + return fmt.Errorf("load Docker port guard persisted status: %w", err) + } + backends := []string{constant.FirewallProviderIptables, constant.FirewallProviderNftables} + if s.runtime != nil { + backends = []string{selectedDockerFirewallBackend(constant.FirewallProviderIptables)} + } + initializedByBackend := make(map[string]bool, len(backends)) + initialized := false + for _, backend := range backends { + initialized, err = s.guardRuntime(backend).Initialized(dockerfirewall.FamilyIPv4) + if err != nil { + return &dockerfirewall.FamilyError{Family: dockerfirewall.FamilyIPv4, Err: fmt.Errorf("inspect initialization: %w", err)} + } + initializedByBackend[backend] = initialized + if initialized { + break + } + } + if !initialized && !persistedEnabled { + return nil + } + runtime, backend, err := s.runtimeForDocker(ctx) + if err != nil { + return err + } + initialized, inspected := initializedByBackend[backend] + if !inspected { + initialized, err = runtime.Initialized(dockerfirewall.FamilyIPv4) + if err != nil { + return &dockerfirewall.FamilyError{Family: dockerfirewall.FamilyIPv4, Err: fmt.Errorf("inspect initialization: %w", err)} + } + } + if !initialized && !persistedEnabled { + return nil + } + policies, err := s.runtimePolicies(ctx) + if err != nil { + return err + } + inventory, err := runtime.ListPolicies() + if err != nil { + return err + } + if !initialized { + err = runtime.Initialize(policies, inventory) + } else { + err = runtime.ReplacePolicies(policies, inventory) + } + if err != nil { + return err + } + return verifyDockerFirewall(runtime, policies, inventory.ReadOnly) +} + +func (s *DockerPortGuardService) runtimeForDocker(ctx context.Context) (dockerfirewall.Runtime, string, error) { + if s.runtime != nil { + return s.runtime, selectedDockerFirewallBackend(constant.FirewallProviderIptables), nil + } + cli, err := s.client() + if err != nil { + return nil, "", buserr.WithDetail("ErrDockerFailed", err.Error(), err) + } + defer cli.Close() + info, err := cli.Info(ctx) + if err != nil { + return nil, "", buserr.WithDetail("ErrDockerFailed", err.Error(), err) + } + backend := selectedDockerFirewallBackend(dockerFirewallBackend(info)) + if backend != constant.FirewallProviderIptables && backend != constant.FirewallProviderNftables { + return nil, backend, fmt.Errorf("Docker firewall backend %q is not supported", backend) + } + return s.guardRuntime(backend), backend, nil +} + +func (s *DockerPortGuardService) guardRuntime(backend string) dockerfirewall.Runtime { + if s.runtimeForBackend != nil { + return s.runtimeForBackend(backend) + } + if s.runtime != nil { + return s.runtime + } + return newDockerFirewallRuntime(backend) +} + +func newDockerFirewallRuntime(backend string) dockerfirewall.Runtime { + if backend == constant.FirewallProviderNftables { + return dockerfirewall.NewNftables() + } + return dockerfirewall.NewIptables() +} + +func selectedDockerFirewallBackend(fallback string) string { + selected, _ := settingRepo.GetValueByKey(constant.FirewallDockerBackendKey) + selected = strings.ToLower(strings.TrimSpace(selected)) + if selected == constant.FirewallProviderIptables || selected == constant.FirewallProviderNftables { + return selected + } + fallback = strings.ToLower(strings.TrimSpace(fallback)) + if fallback == constant.FirewallProviderNftables { + return fallback + } + return constant.FirewallProviderIptables +} + +func dockerFirewallBackend(info system.Info) string { + if info.FirewallBackend == nil || info.FirewallBackend.Driver == "" { + return constant.FirewallProviderIptables + } + return strings.ToLower(info.FirewallBackend.Driver) +} + +func (s *DockerPortGuardService) runtimePolicies(ctx context.Context) ([]dockerfirewall.Policy, error) { + stored, err := s.policies.ListManaged(ctx) + if err != nil { + return nil, err + } + policies := make([]dockerfirewall.Policy, 0, len(stored)) + for _, policy := range stored { + sources := []string{} + _ = json.Unmarshal([]byte(policy.Sources), &sources) + policies = append(policies, dockerfirewall.Policy{UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, Mode: policy.Mode, Sources: sources}) + } + return policies, nil +} + +func dockerGuardReadOnlyPolicyUUID(policy dockerfirewall.ReadOnlyPolicy) string { + nativeRules, _ := json.Marshal(policy.NativeRules) + fingerprint := strings.Join([]string{ + policy.Policy.Family, policy.Policy.HostIP, strconv.Itoa(int(policy.Policy.HostPort)), + policy.Policy.Protocol, policy.Action, string(nativeRules), + }, "\x00") + return uuid.NewSHA1(uuid.NameSpaceOID, []byte(fingerprint)).String() +} + +func dockerPortGuardPersistedEnabled() (bool, error) { + status, err := settingRepo.GetValueByKey(constant.FirewallDockerPortGuardStatusKey) + if errors.Is(err, gorm.ErrRecordNotFound) { + return false, nil + } + return status == constant.StatusEnable, err +} + +func verifyDockerFirewall(runtime dockerfirewall.Runtime, desired []dockerfirewall.Policy, preserved []dockerfirewall.ReadOnlyPolicy) error { + inventory, err := runtime.ListPolicies() + if err != nil { + return fmt.Errorf("verify synchronized Docker firewall policies: %w", err) + } + if !dockerFirewallPoliciesEqual(inventory.Policies, desired) { + return fmt.Errorf("verify synchronized Docker firewall policies: target policies do not match the database") + } + if !readOnlyStatesEqual(inventory.ReadOnly, preserved) { + return fmt.Errorf("verify synchronized Docker firewall policies: read-only runtime rules changed") + } + return nil +} + +func dockerFirewallPoliciesEqual(left, right []dockerfirewall.Policy) bool { + if len(left) != len(right) { + return false + } + counts := make(map[string]int, len(left)) + for _, policy := range left { + counts[dockerFirewallPolicyKey(policy)]++ + } + for _, policy := range right { + key := dockerFirewallPolicyKey(policy) + if counts[key] == 0 { + return false + } + counts[key]-- + } + return true +} + +func dockerFirewallPolicyKey(policy dockerfirewall.Policy) string { + mode := policy.Mode + if mode == dockerfirewall.ModeAllow && len(policy.Sources) == 0 { + mode = dockerfirewall.ModeAll + } + sources := make([]string, 0, len(policy.Sources)) + for _, source := range policy.Sources { + source = strings.TrimSpace(source) + if prefix, err := netip.ParsePrefix(source); err == nil { + source = prefix.Masked().String() + } else if address, err := netip.ParseAddr(source); err == nil { + address = address.Unmap() + source = netip.PrefixFrom(address, address.BitLen()).String() + } + sources = append(sources, source) + } + sort.Strings(sources) + host := policy.HostIP + if address, err := netip.ParseAddr(host); err == nil { + host = address.String() + } + return strings.Join([]string{ + policy.UUID, policy.Family, host, strconv.Itoa(int(policy.HostPort)), + policy.Protocol, mode, strings.Join(sources, ","), + }, "\x00") +} + +func readOnlyStatesEqual(left, right []dockerfirewall.ReadOnlyPolicy) bool { + if len(left) != len(right) { + return false + } + leftRules := flattenNativeRules(left) + rightRules := flattenNativeRules(right) + if len(leftRules) != len(rightRules) { + return false + } + for index := range leftRules { + if leftRules[index].Family != rightRules[index].Family || !slices.Equal(leftRules[index].Tokens, rightRules[index].Tokens) { + return false + } + } + return true +} + +func flattenNativeRules(policies []dockerfirewall.ReadOnlyPolicy) []dockerfirewall.NativeRule { + rules := make([]dockerfirewall.NativeRule, 0) + for _, policy := range policies { + rules = append(rules, policy.NativeRules...) + } + slices.SortStableFunc(rules, func(left, right dockerfirewall.NativeRule) int { + if left.Family < right.Family { + return -1 + } + if left.Family > right.Family { + return 1 + } + if left.Order < right.Order { + return -1 + } + if left.Order > right.Order { + return 1 + } + return 0 + }) + return rules +} + +func newDockerPortGuardService() *DockerPortGuardService { + return &DockerPortGuardService{ + policies: repo.NewIDockerPortGuardRepo(), + client: docker.NewDockerClient, + version: dockerFirewallVersion, + } +} + +func dockerFirewallVersion(backend string) string { + client, err := lifecycle.NewClient(backend) + if err != nil { + return "-" + } + version, err := client.Version() + if err != nil || strings.TrimSpace(version) == "" { + return "-" + } + return version +} + +func firewallDockerActive() (bool, error) { + if !cmd.Which("docker") { + return false, nil + } + return controller.CheckActive("docker") +} + +func restoreFirewalldDependents(ctx context.Context, reason string, restoreDocker bool, restoreForwarding func(context.Context) error, restoreDockerGuard func(context.Context) error) error { + var errs []error + if err := restoreForwarding(ctx); err != nil { + errs = append(errs, fmt.Errorf("restore port forwarding %s: %w", reason, err)) + } + if restoreDocker { + if err := restoreDockerGuard(ctx); err != nil { + errs = append(errs, fmt.Errorf("restore Docker port guard %s: %w", reason, err)) + } + } + return errors.Join(errs...) +} + +func operateFirewallLifecycle(client lifecycle.Client, operation string, withDockerRestart bool, prepareStart func(lifecycle.Client) error, t *task.Task) error { + run := func(operation, name string, action func() error) error { + if t != nil { + return runFirewallLifecycleAction(t, task.GetTaskName(name, operation, ""), action) + } + return action() + } + switch operation { + case string(lifecycle.OperationStart): + if err := run("Start", client.Name(), client.Start); err != nil { + return err + } + case string(lifecycle.OperationRestart): + if err := run("TaskRestart", client.Name(), client.Restart); err != nil { + return err + } + case string(lifecycle.OperationStop): + return stopFirewallLifecycle(client, withDockerRestart, nil, t) + default: + return fmt.Errorf("not supported operation: %s", operation) + } + var recoveryErrors []error + if prepareStart != nil { + err := prepareStart(client) + if err == nil && client.Name() == constant.FirewallProviderFirewalld { + err = lifecycleproviders.RemoveFirewalldSSHService() + } + if err != nil { + recoveryErrors = append(recoveryErrors, fmt.Errorf("prepare firewall after %s: %w", operation, err)) + } + } + if withDockerRestart { + if err := run("TaskRestart", "Docker", func() error { return controller.HandleRestart("docker") }); err != nil { + recoveryErrors = append(recoveryErrors, &firewallDockerRestartError{Err: err}) + } + } + if client.Name() == constant.FirewallProviderFirewalld && operation == string(lifecycle.OperationStart) { + if err := run("TaskRecover", "Fail2Ban", restoreFail2BanAfterFirewallStart); err != nil { + recoveryErrors = append(recoveryErrors, err) + } + } + if err := errors.Join(recoveryErrors...); err != nil { + return &firewallCompletedOperationError{Operation: operation, Err: err} + } + return nil +} + +func (c firewallLifecycleClient) Start() error { + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + return c.Client.Start() +} + +func (c firewallLifecycleClient) Restart() error { + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + return c.Client.Restart() +} + +func runFirewallLifecycleAction(t *task.Task, name string, action func() error) error { + t.Log(i18n.GetWithName("TaskStart", name)) + started := time.Now() + err := t.TaskCtx.Err() + if err == nil { + err = action() + } + t.LogWithStatus(fmt.Sprintf("%s (%.2fs)", name, time.Since(started).Seconds()), err) + return err +} + +func stopFirewallLifecycle(client lifecycle.Client, withDockerRestart bool, prepareStop func() error, t *task.Task) error { + if client.Name() == constant.FirewallProviderFirewalld { + if err := rememberFail2BanBeforeFirewallStop(); err != nil { + return err + } + } + if prepareStop != nil { + if err := prepareStop(); err != nil { + return err + } + } + stop := client.Stop + if t != nil { + stop = func() error { + return runFirewallLifecycleAction(t, task.GetTaskName(client.Name(), "Stop", ""), client.Stop) + } + } + if err := stop(); err != nil { + return err + } + if withDockerRestart { + restart := func() error { return controller.HandleRestart("docker") } + var err error + if t != nil { + err = runFirewallLifecycleAction(t, task.GetTaskName("Docker", "TaskRestart", ""), restart) + } else { + err = restart() + } + if err != nil { + return &firewallDockerRestartError{Err: err} + } + } + return nil +} + +func (c firewallLifecycleClient) Stop() error { + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + return c.Client.Stop() +} + +func rememberFail2BanBeforeFirewallStop() error { + exists, err := controller.CheckExist("fail2ban.service") + if err != nil { + global.LOG.Warnf("check fail2ban.service installation before stopping the firewall failed: %v", err) + } + if !exists { + return nil + } + active, err := controller.CheckActive("fail2ban.service") + if err != nil { + global.LOG.Warnf("check fail2ban.service status before stopping the firewall failed: %v", err) + } + if !active { + return nil + } + if err := os.WriteFile(fail2BanRestoreWithFirewallMarker, nil, 0600); err != nil { + return fmt.Errorf("mark Fail2Ban for restoration with the firewall: %w", err) + } + return nil +} + +func restoreFail2BanAfterFirewallStart() error { + if _, err := os.Stat(fail2BanRestoreWithFirewallMarker); err != nil { + if os.IsNotExist(err) { + return nil + } + return fmt.Errorf("load Fail2Ban restore marker after starting the firewall: %w", err) + } + if err := controller.HandleStart("fail2ban.service"); err != nil { + return fmt.Errorf("restore Fail2Ban after starting the firewall: %w", err) + } + if err := os.Remove(fail2BanRestoreWithFirewallMarker); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("clear Fail2Ban firewall restore status: %w", err) + } + return nil +} + +func currentFirewallRuleSyncTaskLocked() (dto.FirewallRuleSyncTask, error) { + if firewallRuleSyncTaskID != "" { + return dto.FirewallRuleSyncTask{TaskID: firewallRuleSyncTaskID, Executing: true}, nil + } + if global.TaskDB == nil { + return dto.FirewallRuleSyncTask{}, nil + } + taskRepo := repo.NewITaskRepo() + record, err := taskRepo.GetFirst( + repo.WithByStatus(constant.StatusExecuting), + repo.WithByType(task.TaskScopeFirewall), + taskRepo.WithOperate(task.TaskSync), + ) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return dto.FirewallRuleSyncTask{}, nil + } + return dto.FirewallRuleSyncTask{}, err + } + return dto.FirewallRuleSyncTask{TaskID: record.ID, Executing: true}, nil +} + +func (s *FirewallService) operateFilterChainBase(provider string, request dto.FilterChainOperation) error { + firewallRuleMutationMu.Lock() + defer firewallRuleMutationMu.Unlock() + return s.operateFilterChainBaseLocked(provider, request) +} + +func (s *FirewallService) operateFilterChainBaseLocked(provider string, request dto.FilterChainOperation) error { + if err := s.checkSelectedProvider(context.Background(), filter.Provider(provider)); err != nil { + return err + } + if provider != constant.FirewallProviderIptables && provider != constant.FirewallProviderNftables { + return fmt.Errorf("filter chain operations are not supported for %s", provider) + } + operation := firewall.BaseOperation(request.Operate) + var ports []firewall.PortWhitelist + if operation == firewall.BaseOperationInit || operation == firewall.BaseOperationBind { + var err error + ports, err = LoadRequiredFirewallPortWhiteList() + if err != nil { + return err + } + } + if provider == constant.FirewallProviderNftables { + if err := nftables_helper.Operate(operation, ports); err != nil { + return err + } + } else if err := iptables_helper.Operate(operation, ports); err != nil { + return err + } + status := constant.StatusEnable + if operation == firewall.BaseOperationUnbind { + status = constant.StatusDisable + } + return settingRepo.Update("IptablesStatus", status) +} + +func LoadRequiredFirewallPortWhiteList() ([]firewall.PortWhitelist, error) { + ports, err := loadFirewallPortWhiteList() + if err != nil { + return nil, err + } + return firewall.RequiredPortWhitelist(ports) +} + +func lockFirewallLifecycleIdle() error { + if !firewallLifecycleTaskMu.TryLock() { + return buserr.New("TaskIsExecuting") + } + if firewallLifecycleTaskID != "" { + firewallLifecycleTaskMu.Unlock() + return buserr.New("TaskIsExecuting") + } + return nil +} + +func cleanupInactiveSystemBackend(backend string) error { + switch backend { + case constant.FirewallProviderIptables: + return iptables_helper.Cleanup() + case constant.FirewallProviderNftables: + return nftables_helper.Cleanup() + default: + return fmt.Errorf("cleanup is only available for 1Panel-owned iptables and nftables resources") + } +} + +func cleanupSystemBackend(backend string) error { + switch backend { + case constant.FirewallProviderIptables: + if err := iptables_helper.Cleanup(); err != nil { + return err + } + case constant.FirewallProviderNftables: + if err := nftables_helper.Cleanup(); err != nil { + return err + } + default: + return fmt.Errorf("cleanup is only available for 1Panel-owned iptables and nftables resources") + } + return settingRepo.Update("IptablesStatus", constant.StatusDisable) +} + +func resetServiceFirewallBackend(provider string, withDockerRestart bool) error { + client, err := lifecycle.NewClient(provider) + if err != nil { + return err + } + return resetServiceFirewallClient(client, withDockerRestart, func( + client lifecycle.Client, + restartDocker bool, + prepareStop func() error, + ) error { + return stopFirewallLifecycle(client, restartDocker, prepareStop, nil) + }) +} + +func resetServiceFirewallClient(client lifecycle.Client, withDockerRestart bool, stop func(lifecycle.Client, bool, func() error) error) error { + resetter, ok := client.(lifecycle.Resetter) + if !ok { + return fmt.Errorf("firewall provider %s does not support reset", client.Name()) + } + if resetBeforeStop, ok := client.(lifecycle.PreStopResetter); ok { + if err := stop(client, withDockerRestart, resetBeforeStop.ResetBeforeStop); err != nil { + return err + } + return nil + } + return resetter.Reset() +} + +func (s *FirewallService) desiredFirewallRulesByScope(ctx context.Context, stored []model.FirewallRule, runtime filter.Adapter) (map[string][]filter.DesiredRule, []filter.InventoryItem) { + provider := runtime.Provider() + desired := make(map[string][]filter.DesiredRule) + var failures []filter.InventoryItem + ports, protectionErr := loadFirewallPortWhiteList() + var required []firewall.PortWhitelist + if protectionErr == nil { + loadRequired := s.requiredPorts + if loadRequired != nil { + required, protectionErr = loadRequired() + } else { + required, protectionErr = firewall.RequiredPortWhitelist(ports) + } + } + whitelist := filter.NewPortWhitelistIndex(ports) + for _, record := range stored { + compiled, _, err := s.compileRestorableFirewallRules(ctx, record, runtime, required) + if err == nil { + err = protectionErr + } + if err != nil { + rule := filter.FirewallRule{ + UUID: record.UUID, + Scope: filter.Scope{Provider: provider, Family: filter.Family(record.Family), Direction: filter.DirectionInput}.Normalize(), + Protocol: record.Protocol, SourceAddress: record.SourceAddress, SourcePort: record.SourcePort, + DestinationAddress: record.DestinationAddress, DestinationPort: record.DestinationPort, + Interface: record.Interface, ConnectionStates: strings.FieldsFunc(record.ConnectionStates, func(r rune) bool { return r == ',' }), + Action: filter.Action(record.Action), Description: record.Description, Priority: record.Priority, + } + failures = append(failures, filter.InventoryItem{ + Incompatible: isFirewallPolicyIncompatible(err), + Rule: rule, State: filter.InventoryStateDrifted, Match: filter.InventoryMatchNone, + Desired: &filter.DesiredRule{UUID: record.UUID, Rule: rule, Origin: filter.RuleOrigin(record.Origin), Protected: protectionErr != nil || whitelist.Matches(rule)}, + Error: fmt.Sprintf("policy %s: %v", record.UUID, err), + }) + continue + } + for _, rule := range compiled { + rule.Protected = whitelist.Matches(rule.Rule) + rule.Expanded = len(compiled) > 1 + key := rule.Rule.Scope.Key() + desired[key] = append(desired[key], rule) + } + } + return desired, failures +} + +func firewallInventoryPositionRanges(provider filter.Provider, items []filter.InventoryItem) (ipv4, ipv6 filter.PositionRange) { + if provider == filter.ProviderFirewalld { + return filter.PositionRange{Min: -32768, Max: 32767}, filter.PositionRange{Min: -32768, Max: 32767} + } + for _, item := range items { + if item.Observed == nil || item.Observed.Locator.Position == nil { + continue + } + scope := item.Observed.Rule.Scope + if scope.Provider != provider || scope.Direction != filter.DirectionInput { + continue + } + if (provider == filter.ProviderIptables || provider == filter.ProviderNftables) && + (scope.Table != "filter" || scope.Chain != filter.IptablesInputChain) { + continue + } + bounds := &ipv4 + if scope.Family == filter.FamilyIPv6 { + bounds = &ipv6 + } else if scope.Family != filter.FamilyIPv4 { + continue + } + position := *item.Observed.Locator.Position + if position < 1 { + continue + } + if bounds.Min == 0 || position < bounds.Min { + bounds.Min = position + } + bounds.Max = max(bounds.Max, position) + if provider != filter.ProviderUFW { + bounds.Min = 1 + } + } + return +} + +func (s *FirewallService) adoptRule(ctx context.Context, runtime filter.Adapter, snapshot filter.RuleSet, observed filter.ObservedRule, source dto.FirewallRuleCreateItem) error { + if (observed.Rule.Scope.Provider == filter.ProviderIptables || observed.Rule.Scope.Provider == filter.ProviderNftables) && + (observed.Rule.Scope.Chain == filter.BasicBeforeChain || observed.Rule.Scope.Chain == filter.BasicAfterChain) { + return fmt.Errorf("%w: system preset chains cannot be adopted", filter.ErrUnsupportedScope) + } + if observed.Protected { + return filter.ErrProtectedRule + } + if observed.ParseStatus != filter.ParseStatusSupported || + (observed.Persistence != "" && observed.Persistence != filter.PersistenceStatusConverged) { + return fmt.Errorf("%w: rule cannot be managed", filter.ErrRuleOperation) + } + rule, err := prepareFirewallBackendRule(ctx, runtime, observed.Rule) + if err != nil { + return err + } + if rule.Scope.Provider != filter.ProviderFirewalld && observed.Locator.Position != nil { + position := int64(*observed.Locator.Position) + rule.OrderIndex = &position + } + record, err := firewallRuleModelForCreate(rule, source, constant.FirewallRuleOriginAdopted) + if err != nil { + return err + } + stored, err := s.rules.List(ctx) + if err != nil { + return err + } + if runtime.Provider() == filter.ProviderNftables { + duplicates := 0 + for _, candidate := range snapshot.Rules { + if candidate.ParseStatus != filter.ParseStatusSupported { + continue + } + same, err := filter.SameRuleContent(candidate.Rule, rule) + if err != nil { + return err + } + if same { + duplicates++ + if duplicates > 1 { + return buserr.WithDetail("ErrInvalidParams", "duplicate firewall rules prevent adoption; manually delete duplicate rules and retry", nil) + } + } + } + } + identities, err := firewallRuleCollisions(stored, rule.Scope.Provider) + if err != nil { + return err + } + for _, existing := range stored { + marker := "1panel-rule:" + existing.UUID + if existing.UUID != "" && (observed.Marker == marker || strings.HasPrefix(observed.Marker, marker+"-")) { + return fmt.Errorf("%w: rule is already managed", filter.ErrRuleOperation) + } + } + if err := identities.CheckDuplicate(rule); err != nil { + if errors.Is(err, filter.ErrRuleOperation) { + return buserr.WithDetail("ErrInvalidParams", "duplicate firewall rules prevent adoption; manually delete duplicate rules and retry", nil) + } + return err + } + record.UUID = uuid.NewString() + rule.UUID = record.UUID + if runtime.Provider() != filter.ProviderNftables { + ports, err := loadFirewallPortWhiteList() + if err != nil { + return err + } + if filter.RuleMatchesPortWhitelist(rule, ports) { + return filter.ErrProtectedRule + } + before := observed.Rule + before.UUID = record.UUID + scope := filter.RuleSet{Scope: rule.Scope} + remove, err := runtime.BuildCommands(scope, []filter.RuleChange{{ + Operation: filter.ChangeDelete, Before: &before, CommandOnly: true, + UnmarkedAdopted: observed.Marker == "", PreviousMarker: observed.Marker, + }}) + if err != nil { + return err + } + create, err := runtime.BuildCommands(scope, []filter.RuleChange{{ + Operation: filter.ChangeCreate, After: &rule, CommandOnly: true, Append: rule.OrderIndex == nil, + }}) + if err != nil { + return err + } + remove.CommandOnly, create.CommandOnly = true, true + if err := runtime.RunCommands(ctx, remove); err != nil { + return err + } + record.Priority = rule.Priority + saveCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + err = s.saveFirewallRule(saveCtx, &record) + cancel() + if err != nil { + return err + } + runErr := runtime.RunCommands(ctx, create) + if err := errors.Join(runErr, persistFirewallRules(ctx, runtime, create)); err != nil { + return buserr.WithDetail("ErrFirewallRuleSavedApplyFailed", err.Error(), err) + } + return nil + } + plan, verification, err := applyFirewallChanges(runtime, ctx, snapshot, []filter.RuleChange{{ + Operation: filter.ChangeAdopt, After: &rule, Locator: &observed.Locator, PreviousMarker: observed.Marker, + }}) + if err != nil { + return err + } + if !verification.Matched { + return filter.ErrVerificationFailed + } + if _, err := filter.FindCommittedObserved(verification.RuleSet, rule, plan); err != nil { + return rollbackFirewallPlan(ctx, runtime, plan, err) + } + return s.saveFirewallRule(ctx, &record) +} + +func (index firewallRuleCollisionIndex) CheckDuplicate(rule filter.FirewallRule) error { + key, err := filter.RuleMatchKey(rule) + if err != nil { + return err + } + for _, action := range index[key] { + if action == rule.Action { + return checkCollisionActions(rule.Action, action) + } + } + return nil +} + +func findPortWhitelistRule(rules []firewall.PortWhitelist, target firewall.PortWhitelist) (int, error) { + index := slices.IndexFunc(rules, func(rule firewall.PortWhitelist) bool { + return samePortWhitelistRule(rule, target) + }) + if index < 0 { + return -1, fmt.Errorf("firewall port whitelist rule has changed or no longer exists; refresh and retry") + } + return index, nil +} + +func samePortWhitelistRule(left, right firewall.PortWhitelist) bool { + if reflect.DeepEqual(left, right) { + return true + } + normalizedLeft, err := firewall.ValidatePortWhitelist([]firewall.PortWhitelist{left}) + if err != nil { + return false + } + normalizedRight, err := firewall.ValidatePortWhitelist([]firewall.PortWhitelist{right}) + if err != nil { + return false + } + slices.Sort(normalizedLeft[0].Sources) + slices.Sort(normalizedRight[0].Sources) + return reflect.DeepEqual(normalizedLeft[0], normalizedRight[0]) +} + +func loadSSHWhitelistPortFrom(path string) (string, error) { + directives, _, err := parseSSHConfigTree(path) + if errors.Is(err, os.ErrNotExist) { + return defaultSSHPort, nil + } + if err != nil { + return "", err + } + return loadSSHPortValues(directives)[0], nil +} + +func newFirewallService() *FirewallService { + return &FirewallService{ + rules: repo.NewIFirewallRuleRepo(), + adapters: nil, + forwardingSync: newForwardingService(), + dockerSync: newDockerPortGuardService(), + selectedProvider: firewallRuleSelectedProvider, + requiredPorts: LoadRequiredFirewallPortWhiteList, + cleanupBackend: cleanupSystemBackend, + cleanupInactiveBackend: cleanupInactiveSystemBackend, + resetBackend: resetServiceFirewallBackend, + dockerActive: firewallDockerActive, + restoreForwarding: func(ctx context.Context) error { + return newForwardingService().Restore(ctx) + }, + restoreDockerGuard: ReconcileDockerPortGuard, + baseClient: NewSelectedSystemFirewallClient, + } +} + +func firewallRuleSelectedProvider(context.Context) (filter.Provider, error) { + client, err := NewSelectedSystemFirewallClient() + if err != nil { + return "", fmt.Errorf("%w: %v", filter.ErrProviderUnavailable, err) + } + return filter.Provider(client.Name()), nil +} + +func (s *ForwardingService) loadRuleSyncCandidates(ctx context.Context, targetProvider filter.Provider) (forwarding.Adapter, []forwardingRuleSyncCandidate, []forwarding.Rule, bool, error) { + target, err := s.clientFactory() + if err != nil { + return nil, nil, nil, false, err + } + if target.Name() != string(targetProvider) { + return nil, nil, nil, false, fmt.Errorf( + "%w: selected forwarding backend is %s, requested target is %s", + filter.ErrProviderUnavailable, target.Name(), targetProvider, + ) + } + stored, err := s.rules.List(ctx) + if err != nil { + return nil, nil, nil, false, err + } + candidates := make([]forwardingRuleSyncCandidate, 0, len(stored)) + for _, record := range stored { + rule := forwarding.Rule{ + Family: record.Family, Protocol: record.Protocol, Port: record.Port, TargetIP: record.TargetIP, + TargetPort: record.TargetPort, Interface: record.Interface, + } + normalized, normalizeErr := forwarding.NormalizeRule(rule) + candidates = append(candidates, forwardingRuleSyncCandidate{rule: normalized, err: normalizeErr}) + } + initialized, _, err := target.InitStatus() + if err != nil { + return nil, nil, nil, false, err + } + targetRules := make([]forwarding.Rule, 0) + if initialized { + targetRules, err = target.List() + if err != nil { + return nil, nil, nil, false, err + } + targetRules, err = normalizeForwardingRuntimeRules(targetRules) + if err != nil { + return nil, nil, nil, false, err + } + } + return target, candidates, targetRules, initialized, nil +} + +func normalizeForwardingRuntimeRules(rules []forwarding.Rule) ([]forwarding.Rule, error) { + normalized := make([]forwarding.Rule, 0, len(rules)) + for _, rule := range rules { + item, err := forwarding.NormalizeRule(rule) + if err != nil { + return nil, fmt.Errorf("normalize target forwarding rule %s: %w", rule.Identity(), err) + } + normalized = append(normalized, item) + } + return normalized, nil +} + +func databaseRuleSyncTarget(request dto.FirewallRuleSyncRequest, subsystem string) (filter.Provider, error) { + if request.SourceProvider != "" { + return "", fmt.Errorf("%w: %s synchronization reads rules from the database and does not accept a source provider", filter.ErrInvalidRule, subsystem) + } + if request.ResetSource { + return "", fmt.Errorf("%w: %s synchronization does not have a source firewall to reset", filter.ErrInvalidRule, subsystem) + } + if request.TargetProvider != filter.ProviderIptables && request.TargetProvider != filter.ProviderNftables { + return "", fmt.Errorf("%w: %s synchronization only supports iptables and nftables targets", filter.ErrInvalidRule, subsystem) + } + return request.TargetProvider, nil +} + +func forwardingSyncPreview(target filter.Provider, candidates []forwardingRuleSyncCandidate, actual []forwarding.Rule) dto.FirewallRuleSyncPreview { + desired := make([]firewallSyncDesired[forwarding.Rule], 0, len(candidates)) + for _, candidate := range candidates { + desired = append(desired, firewallSyncDesired[forwarding.Rule]{ + Value: candidate.rule, + Payload: dto.FirewallRuleSyncItem{ + SourceUUID: candidate.rule.Identity(), ForwardRule: &dto.ForwardRule{Family: candidate.rule.Family, Protocol: candidate.rule.Protocol, Port: candidate.rule.Port, TargetIP: candidate.rule.TargetIP, TargetPort: candidate.rule.TargetPort, Interface: candidate.rule.Interface}, + }, + Err: candidate.err, + }) + } + return firewallDiffPreview( + "forwarding", target, desired, actual, + func(rule forwarding.Rule) string { return rule.Identity() }, + func(rule forwarding.Rule) dto.FirewallRuleSyncItem { + return dto.FirewallRuleSyncItem{SourceUUID: rule.Identity(), ForwardRule: &dto.ForwardRule{Family: rule.Family, Protocol: rule.Protocol, Port: rule.Port, TargetIP: rule.TargetIP, TargetPort: rule.TargetPort, Interface: rule.Interface}} + }, + ) +} + +func firewallDiffPreview[T any](subsystem string, target filter.Provider, desired []firewallSyncDesired[T], actual []T, key func(T) string, actualItem func(T) dto.FirewallRuleSyncItem) dto.FirewallRuleSyncPreview { + preview := dto.FirewallRuleSyncPreview{Subsystem: subsystem, TargetProvider: target, 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)) + for _, candidate := range desired { + item := candidate.Payload + item.Status, item.ReasonCode, item.Reason = "", "", "" + switch { + case candidate.Err != nil: + item.Status, item.ReasonCode, item.Reason = firewallsync.StatusBlocked, firewallsync.ReasonInvalidPolicy, candidate.Err.Error() + default: + match := -1 + for _, index := range actualByKey[key(candidate.Value)] { + if !matched[index] { + match = index + break + } + } + if match >= 0 { + matched[match] = true + item.Status, item.ReasonCode = firewallsync.StatusExisting, firewallsync.ReasonAlreadyExists + item.Reason = firewallSyncReasonMessage(item.ReasonCode) + } else { + item.Status = firewallsync.StatusReady + } + } + preview.Add(item) + } + for index, value := range actual { + if matched[index] { + continue + } + item := actualItem(value) + item.Status, item.ReasonCode = firewallsync.StatusRemove, firewallsync.ReasonOnlyExistsInTarget + item.Reason = firewallSyncReasonMessage(item.ReasonCode) + preview.Add(item) + } + return preview +} + +func firewallSyncResult(preview dto.FirewallRuleSyncPreview, cause error, executed bool) dto.FirewallRuleSyncResult { + result := dto.FirewallRuleSyncResult{Subsystem: preview.Subsystem, TargetProvider: preview.TargetProvider} + for _, item := range preview.Items { + if item.Status != firewallsync.StatusRemove { + result.Total++ + } + switch item.Status { + case firewallsync.StatusExisting: + result.Skipped++ + case firewallsync.StatusBlocked: + appendDatabaseSyncFailure(&result, item, errors.New(item.Reason)) + case firewallsync.StatusReady: + if executed { + if cause != nil { + appendDatabaseSyncFailure(&result, item, cause) + } else { + result.Succeeded++ + } + } + case firewallsync.StatusRemove: + if executed && cause == nil { + result.Removed++ + } + } + } + return result +} + +func firewallRuleStatesEqual[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 (s *DockerPortGuardService) loadRuleSyncCandidates(ctx context.Context, request dto.FirewallRuleSyncRequest) (string, []model.DockerPortGuardPolicy, dockerfirewall.Runtime, error) { + targetProvider, err := databaseRuleSyncTarget(request, "Docker") + if err != nil { + return "", nil, nil, err + } + target := string(targetProvider) + selected, err := s.selectedRuleSyncBackend(ctx) + if err != nil { + return "", nil, nil, err + } + if target != selected { + return "", nil, nil, fmt.Errorf( + "%w: selected Docker firewall backend is %s, requested target is %s", + filter.ErrProviderUnavailable, selected, target, + ) + } + policies, err := s.policies.ListManaged(ctx) + if err != nil { + return "", nil, nil, err + } + return target, policies, s.guardRuntime(target), nil +} + +func (s *DockerPortGuardService) selectedRuleSyncBackend(ctx context.Context) (string, error) { + selected, _ := settingRepo.GetValueByKey(constant.FirewallDockerBackendKey) + selected = strings.ToLower(strings.TrimSpace(selected)) + if selected == constant.FirewallProviderIptables || selected == constant.FirewallProviderNftables { + return selected, nil + } + if s.client == nil { + return "", buserr.New("ErrDockerFailed") + } + cli, err := s.client() + if err != nil { + return "", buserr.WithDetail("ErrDockerFailed", err.Error(), err) + } + defer cli.Close() + info, err := cli.Info(ctx) + if err != nil { + return "", buserr.WithDetail("ErrDockerFailed", err.Error(), err) + } + return selectedDockerFirewallBackend(dockerFirewallBackend(info)), nil +} + +func dockerSyncPreview(target filter.Provider, policies []model.DockerPortGuardPolicy, inventory dockerfirewall.PolicyInventory) dto.FirewallRuleSyncPreview { + desired := make([]firewallSyncDesired[dockerfirewall.Policy], 0, len(policies)) + for _, policy := range policies { + sources := []string{} + _ = json.Unmarshal([]byte(policy.Sources), &sources) + desired = append(desired, firewallSyncDesired[dockerfirewall.Policy]{ + Value: dockerfirewall.Policy{ + UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, + Protocol: policy.Protocol, Mode: policy.Mode, Sources: sources, + }, + Payload: dto.FirewallRuleSyncItem{SourceUUID: policy.UUID, DockerRule: &dto.DockerPortGuardEndpoint{ + Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, + PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: sources, Description: policy.Description, + TrafficPath: dockerTrafficPathUnknown, ManagementTarget: dockerManagementNeedsDiagnosis, + ManagementReason: dockerReasonNoMatchingPath, + }}, + }) + } + preview := firewallDiffPreview( + "docker", target, desired, inventory.Policies, dockerFirewallPolicyKey, + func(policy dockerfirewall.Policy) dto.FirewallRuleSyncItem { + return dto.FirewallRuleSyncItem{SourceUUID: policy.UUID, DockerRule: &dto.DockerPortGuardEndpoint{ + Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, + PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: append([]string(nil), policy.Sources...), + TrafficPath: dockerTrafficPathUnknown, ManagementTarget: dockerManagementNeedsDiagnosis, + ManagementReason: dockerReasonNoMatchingPath, + }} + }, + ) + for _, policy := range inventory.ReadOnly { + preview.Add(dto.FirewallRuleSyncItem{ + SourceUUID: dockerGuardReadOnlyPolicyUUID(policy), + DockerRule: &dto.DockerPortGuardEndpoint{ + Family: policy.Policy.Family, HostIP: policy.Policy.HostIP, HostPort: policy.Policy.HostPort, + Protocol: policy.Policy.Protocol, PolicyUUID: dockerGuardReadOnlyPolicyUUID(policy), Sources: append([]string(nil), policy.Policy.Sources...), + NativeAction: policy.Action, ReadOnly: true, TrafficPath: dockerTrafficPathUnknown, + ManagementTarget: dockerManagementNeedsDiagnosis, ManagementReason: dockerReasonNoMatchingPath, + }, + Status: firewallsync.StatusBlocked, + ReasonCode: firewallsync.ReasonReadOnlyRule, + Reason: firewallSyncReasonMessage(firewallsync.ReasonReadOnlyRule), + }) + } + return preview +} + +func discoverDockerEndpoints(ctx context.Context, cli *client.Client, all bool) ([]dto.DockerPortGuardEndpoint, error) { + containers, err := cli.ContainerList(ctx, containertypes.ListOptions{All: all}) + if err != nil { + return nil, err + } + endpoints := make([]dto.DockerPortGuardEndpoint, 0) + for _, item := range containers { + name := "" + if len(item.Names) > 0 { + name = strings.TrimPrefix(item.Names[0], "/") + } + compose := item.Labels[dockerGuardComposeProjectLabel] + application := "" + if created, ok := item.Labels[dockerGuardComposeCreatedBy]; ok && created == "Apps" { + application = compose + } + for _, port := range item.Ports { + if port.PublicPort == 0 || (port.Type != "tcp" && port.Type != "udp") { + continue + } + family := dockerfirewall.FamilyIPv4 + hostIP := port.IP + if addr, err := netip.ParseAddr(hostIP); err == nil && addr.Is6() { + family = dockerfirewall.FamilyIPv6 + } else if hostIP == "" { + hostIP = "0.0.0.0" + } + endpoints = append(endpoints, dto.DockerPortGuardEndpoint{Family: family, HostIP: hostIP, HostPort: port.PublicPort, Protocol: port.Type, ContainerID: item.ID, ContainerName: name, ContainerState: item.State, ContainerPort: port.PrivatePort, Compose: compose, Application: application, Sources: []string{}}) + } + } + return endpoints, nil +} + +func annotateDockerEndpointManagement(endpoints []dto.DockerPortGuardEndpoint, backend string) { + rules := map[string]dockerfirewall.DNATRules{ + constant.FirewallFamilyIPv4: dockerfirewall.ReadDNATRules(backend, constant.FirewallFamilyIPv4), + constant.FirewallFamilyIPv6: dockerfirewall.ReadDNATRules(backend, constant.FirewallFamilyIPv6), + } + proxies := dockerfirewall.ReadProxyEndpoints() + inspections := make(map[string]dockerfirewall.EndpointInspection, len(rules)) + for family, familyRules := range rules { + inspections[family] = dockerfirewall.InspectEndpoints(backend, family, familyRules, proxies) + } + for i := range endpoints { + endpoints[i].TrafficPath, endpoints[i].ManagementTarget, endpoints[i].ManagementReason = + dockerEndpointManagement(inspections[endpoints[i].Family], endpoints[i]) + } +} + +func dockerEndpointManagement(inspection dockerfirewall.EndpointInspection, endpoint dto.DockerPortGuardEndpoint) (string, string, string) { + if !inspection.DNATInspected { + return dockerTrafficPathUnknown, dockerManagementNeedsDiagnosis, dockerReasonNATInspectFailed + } + dnatMatched := inspection.DNATMatches(endpoint.HostIP, endpoint.HostPort, endpoint.Protocol) + if dnatMatched && inspection.IngressReachable { + return dockerTrafficPathForward, dockerManagementContainerGuard, "" + } + if !inspection.ProxyInspected { + return dockerTrafficPathUnknown, dockerManagementNeedsDiagnosis, dockerReasonProxyInspectFailed + } + if inspection.ProxyMatches(endpoint.HostIP, endpoint.HostPort, endpoint.Protocol) { + return dockerTrafficPathInput, dockerManagementHostFirewall, "" + } + if dnatMatched { + return dockerTrafficPathUnknown, dockerManagementNeedsDiagnosis, dockerReasonNATChainUnreachable + } + return dockerTrafficPathUnknown, dockerManagementNeedsDiagnosis, dockerReasonNoMatchingPath +} + +func groupDockerGuardContainers(endpoints []dto.DockerPortGuardEndpoint) []dto.DockerPortGuardContainer { + containers := make(map[string]*dto.DockerPortGuardContainer) + order := make([]string, 0) + for _, endpoint := range endpoints { + key := endpoint.ContainerID + if key == "" { + key = "__orphan__" + } + container, ok := containers[key] + if !ok { + container = &dto.DockerPortGuardContainer{ + Key: key, Name: endpoint.ContainerName, Compose: endpoint.Compose, + Application: endpoint.Application, Endpoints: []dto.DockerPortGuardEndpoint{}, + } + containers[key] = container + order = append(order, key) + } + container.Endpoints = append(container.Endpoints, endpoint) + } + sort.Slice(order, func(i, j int) bool { + return containers[order[i]].Name < containers[order[j]].Name + }) + + result := make([]dto.DockerPortGuardContainer, 0, len(order)) + for _, key := range order { + container := containers[key] + items := make([]docker.PortRangeItem, 0, len(container.Endpoints)) + for i, endpoint := range container.Endpoints { + sources := append([]string(nil), endpoint.Sources...) + sort.Strings(sources) + policyKey := fmt.Sprintf("%t|%s|%s|%t|%s|%s|%s", endpoint.PolicyUUID != "", endpoint.Mode, strings.Join(sources, ","), endpoint.Effective, endpoint.Description, endpoint.ManagementTarget, endpoint.ManagementReason) + items = append(items, docker.PortRangeItem{ + Key: endpoint.Family + "|" + endpoint.HostIP + "|" + endpoint.Protocol + "|" + policyKey, + PublicPort: endpoint.HostPort, PrivatePort: endpoint.ContainerPort, + HasPrivatePort: endpoint.ContainerPort != 0, Position: i, + }) + } + container.PortGroups = make([]dto.DockerPortGuardPortGroup, 0, len(items)) + for _, portRange := range docker.MergePortRanges(items) { + start := container.Endpoints[portRange.Start.Position] + address := start.HostIP + if strings.Contains(address, ":") { + address = "[" + address + "]" + } + ports := fmt.Sprintf("%d", portRange.Start.PublicPort) + if portRange.Start.PublicPort != portRange.End.PublicPort { + ports = fmt.Sprintf("%d-%d", portRange.Start.PublicPort, portRange.End.PublicPort) + } + container.PortGroups = append(container.PortGroups, dto.DockerPortGuardPortGroup{ + Key: fmt.Sprintf("%s|%d-%d", portRange.Start.Key, portRange.Start.PublicPort, portRange.End.PublicPort), + Label: fmt.Sprintf("%s:%s/%s", address, ports, start.Protocol), Endpoint: start, + Endpoints: func() []dto.DockerPortGuardEndpoint { + members := make([]dto.DockerPortGuardEndpoint, 0, len(portRange.Items)) + for _, item := range portRange.Items { + members = append(members, container.Endpoints[item.Position]) + } + return members + }(), + }) + } + result = append(result, *container) + } + return result +} + +func normalizeDockerFirewallUUIDs(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, buserr.WithDetail("ErrInvalidParams", "policy UUID cannot be empty", nil) + } + if _, exists := seen[policyUUID]; exists { + continue + } + seen[policyUUID] = struct{}{} + uuids = append(uuids, policyUUID) + } + if len(uuids) == 0 { + return nil, buserr.WithDetail("ErrInvalidParams", "policy UUIDs cannot be empty", nil) + } + return uuids, nil +} + +func queueFirewallRuleTask(subsystem, operation string, labels []string, apply func(context.Context) error) (dto.FilterChainOperationResponse, error) { + taskItem, err := task.NewTask(firewallTaskName(operation, subsystem, ""), operation, task.TaskScopeFirewall, "", 0) + if err != nil { + return dto.FilterChainOperationResponse{}, err + } + taskItem.AddSubTaskWithOps(taskItem.Name, func(t *task.Task) error { + t.Logf("rules=%d", len(labels)) + err := t.TaskCtx.Err() + if err == nil { + err = apply(t.TaskCtx) + } + succeeded, failed := 0, 0 + for _, label := range labels { + if err != nil { + failed++ + t.LogFailedWithErr(label, err) + } else { + succeeded++ + t.LogSuccess(label) + } + } + t.Log(i18n.GetMsgWithMap("FirewallRuleOperationResult", map[string]interface{}{ + "succeeded": succeeded, "failed": failed, + })) + return err + }, nil, 0, 0) + if err := repo.NewITaskRepo().Save(context.Background(), taskItem.Task); err != nil { + taskItem.LogFailedWithErr(taskItem.Name, err) + closeUnstartedFirewallTask(taskItem) + return dto.FilterChainOperationResponse{}, fmt.Errorf("save firewall rule task: %w", err) + } + go func() { _ = taskItem.Execute() }() + return dto.FilterChainOperationResponse{TaskID: taskItem.TaskID, Queued: true}, nil +} + +func normalizeDockerFirewallPolicy(policy dockerfirewall.Policy) (dockerfirewall.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 != dockerfirewall.FamilyIPv4 && policy.Family != dockerfirewall.FamilyIPv6) || + (policy.Mode != dockerfirewall.ModeAll && policy.Mode != dockerfirewall.ModeSources && policy.Mode != dockerfirewall.ModeAllow) { + return dockerfirewall.Policy{}, buserr.WithDetail("ErrInvalidParams", "invalid policy fields", nil) + } + address, err := netip.ParseAddr(policy.HostIP) + if err != nil || (policy.Family == dockerfirewall.FamilyIPv4) != address.Is4() { + return dockerfirewall.Policy{}, buserr.WithDetail("ErrInvalidParams", "host IP does not match address family", nil) + } + 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 dockerfirewall.Policy{}, buserr.WithDetail("ErrInvalidParams", fmt.Sprintf("invalid source address %q", source), nil) + } + } + if (policy.Family == dockerfirewall.FamilyIPv4) != prefix.Addr().Is4() { + return dockerfirewall.Policy{}, buserr.WithDetail("ErrInvalidParams", fmt.Sprintf("source %q does not match address family", source), nil) + } + canonical := prefix.Masked().String() + if _, exists := seen[canonical]; !exists { + seen[canonical] = struct{}{} + normalizedSources = append(normalizedSources, canonical) + } + } + if policy.Mode != dockerfirewall.ModeAll && len(normalizedSources) == 0 { + return dockerfirewall.Policy{}, buserr.WithDetail("ErrInvalidParams", "source-based modes require at least one source", nil) + } + if policy.Mode == dockerfirewall.ModeAll { + normalizedSources = []string{} + } + sort.Strings(normalizedSources) + policy.Sources = normalizedSources + return policy, nil +} + +func (i forwardingInventoryItem) SyncStatus() string { + switch { + case i.IsDesired && i.IsRuntime: + return forwardingSyncConverged + case i.IsDesired: + return forwardingSyncMissing + default: + return forwardingSyncRuntimeOnly + } +} + +func forwardingOperationsOnlyRemove(operations []dto.ForwardRuleOperation) bool { + if len(operations) == 0 { + return false + } + for _, operation := range operations { + if operation.Operation != string(forwarding.OperationRemove) { + return false + } + } + return true +} + +func (e *firewallDockerRestartError) Error() string { + return fmt.Sprintf("failed to restart Docker: %v", e.Err) +} + +func (e *firewallDockerRestartError) Unwrap() error { + return e.Err +} + +func (e *firewallCompletedOperationError) Error() string { + return fmt.Sprintf("firewall %s completed with recovery errors: %v", e.Operation, e.Err) +} + +func (e *firewallCompletedOperationError) Unwrap() error { + return e.Err +} + +func InitializeFirewallWhitelistPorts(entries []firewall.PortWhitelist) ([]firewall.PortWhitelist, error) { + entries = slices.Clone(entries) + var sshPort string + for i := range entries { + entry := &entries[i] + if entry.Type == "" || entry.Port != "" { + continue + } + switch entry.Type { + case firewall.PortWhitelistTypePanel: + entry.Port = LoadPanelPort() + case firewall.PortWhitelistTypeSSH: + if sshPort == "" { + var err error + sshPort, err = loadSSHWhitelistPortFrom(sshPath) + if err != nil { + return nil, err + } + } + entry.Port = sshPort + } + } + return firewall.ValidatePortWhitelist(entries) +} diff --git a/agent/app/service/forward.go b/agent/app/service/forward.go index 8c40d34f71d3..492bf7216ee1 100644 --- a/agent/app/service/forward.go +++ b/agent/app/service/forward.go @@ -7,6 +7,7 @@ import ( "strconv" "strings" "sync" + "time" "github.com/1Panel-dev/1Panel/agent/app/dto" "github.com/1Panel-dev/1Panel/agent/app/model" @@ -16,11 +17,19 @@ import ( "github.com/1Panel-dev/1Panel/agent/constant" "github.com/1Panel-dev/1Panel/agent/global" "github.com/1Panel-dev/1Panel/agent/i18n" + "github.com/1Panel-dev/1Panel/agent/utils/cmd" "github.com/1Panel-dev/1Panel/agent/utils/firewall" + "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding" "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" ) +const ( + forwardingSyncConverged = "converged" + forwardingSyncMissing = "missing" + forwardingSyncRuntimeOnly = "runtime_only" +) + type IForwardingService interface { LoadBaseInfo() (dto.FirewallSubsystemStatus, error) SearchRules(request dto.ForwardRuleSearch) (int64, []dto.ForwardRule, error) @@ -31,7 +40,7 @@ type IForwardingService interface { } type ForwardingService struct { - managerFactory func() (*forwarding.Manager, error) + clientFactory func() (forwarding.Adapter, error) rules repo.IForwardingRuleRepo enabled func() (bool, error) persistBackend func(string) error @@ -39,43 +48,27 @@ type ForwardingService struct { } var errForwardingBackendUnavailable = errors.New("no supported forwarding backend detected") -var forwardingMutationMu sync.Mutex -const ( - forwardingSyncConverged = "converged" - forwardingSyncMissing = "missing" - forwardingSyncRuntimeOnly = "runtime_only" -) +var forwardingMutationMu sync.Mutex var ( forwardingSyncStateMu sync.RWMutex forwardingLastSyncErr error ) -func NewIForwardingService() IForwardingService { - return newForwardingService() -} - -func newForwardingService() *ForwardingService { - return &ForwardingService{ - managerFactory: newForwardingManager, - rules: repo.NewIForwardingRuleRepo(), - enabled: forwardingPersistedEnabled, - markEnabled: func() error { - return settingRepo.UpdateOrCreate(constant.FirewallForwardingInitializedKey, constant.StatusEnable) - }, - persistBackend: func(backend string) error { - return settingRepo.UpdateOrCreate(constant.FirewallForwardingBackendKey, backend) - }, - } -} - func (s *ForwardingService) LoadBaseInfo() (dto.FirewallSubsystemStatus, error) { - selected := configuredForwardingBackend() + selected, _ := settingRepo.GetValueByKey(constant.FirewallForwardingBackendKey) + selected = strings.TrimSpace(selected) + if selected == "" { + selected = constant.FirewallProviderIptables + } baseInfo := dto.FirewallSubsystemStatus{ - Version: "-", Name: forwardingDisplayName(selected), Backend: selected, SyncError: lastForwardingSyncError(), + Version: "-", Name: selected, Backend: selected, SyncError: lastForwardingSyncError(), + } + if selected == constant.FirewallProviderIptables || selected == constant.FirewallProviderNftables { + baseInfo.Name += "-forward" } - manager, err := s.managerFactory() + manager, err := s.clientFactory() if err != nil { if errors.Is(err, errForwardingBackendUnavailable) { baseInfo.Reason = constant.FirewallBackendNotInstalled @@ -83,37 +76,39 @@ func (s *ForwardingService) LoadBaseInfo() (dto.FirewallSubsystemStatus, error) } return baseInfo, err } - status, err := manager.Status() + client, err := lifecycle.NewClient(manager.Name()) if err != nil { return baseInfo, err } + version, versionErr := client.Version() + status, statusErr := loadForwardingFirewallOverview(manager) + if err := errors.Join(versionErr, statusErr); err != nil { + return baseInfo, err + } baseInfo.IsExist = true - baseInfo.Name, baseInfo.Backend = forwardingDisplayName(status.Name), status.Name - baseInfo.Version = status.Version + baseInfo.Name, baseInfo.Backend = manager.Name(), manager.Name() + if baseInfo.Backend == constant.FirewallProviderIptables || baseInfo.Backend == constant.FirewallProviderNftables { + baseInfo.Name += "-forward" + } + baseInfo.Version = version baseInfo.PingStatus = firewall.LoadPingStatus() baseInfo.IsInit, baseInfo.IsBind = status.IsInit, status.IsBind - baseInfo.IPv4 = loadForwardingFamilyInfo(manager, status.Name, constant.FirewallFamilyIPv4) - baseInfo.IPv6 = loadForwardingFamilyInfo(manager, status.Name, constant.FirewallFamilyIPv6) - return baseInfo, nil -} - -func loadForwardingFamilyInfo(manager *forwarding.Manager, backend, family string) dto.FirewallBackendFamilyStatus { - initialized, bound, err := manager.FamilyStatus(family) - available := err == nil - if backend == constant.FirewallProviderIptables && family == constant.FirewallFamilyIPv6 { - commands, commandErr := lifecycle.ResolveIptablesCommands() - available = available && commandErr == nil && commands.IPv6Available() - } - return dto.FirewallBackendFamilyStatus{Available: available, Initialized: initialized, Bound: bound} -} - -func forwardingDisplayName(backend string) string { - switch backend { - case constant.FirewallProviderIptables, constant.FirewallProviderNftables: - return backend + "-forward" - default: - return backend + baseInfo.IPv4, baseInfo.IPv6 = status.IPv4, status.IPv6 + for _, family := range []struct { + command string + status *dto.FirewallBackendFamilyStatus + }{ + {"iptables", &baseInfo.IPv4}, + {"ip6tables", &baseInfo.IPv6}, + } { + policy, err := loadForwardPolicy(family.command) + if err != nil { + global.LOG.Warnf("inspect %s FORWARD policy: %v", family.command, err) + continue + } + family.status.ForwardPolicy = policy } + return baseInfo, nil } func (s *ForwardingService) SearchRules(request dto.ForwardRuleSearch) (int64, []dto.ForwardRule, error) { @@ -124,11 +119,11 @@ func (s *ForwardingService) SearchRules(request dto.ForwardRuleSearch) (int64, [ if err != nil { return 0, nil, err } - manager, err := s.managerFactory() + manager, err := s.clientFactory() if err != nil { return 0, nil, err } - runtime, err := manager.List("", "") + runtime, err := manager.List() if err != nil { return 0, nil, err } @@ -178,24 +173,18 @@ func (s *ForwardingService) SearchRules(request dto.ForwardRuleSearch) (int64, [ return int64(total), items, nil } -func forwardingRuleMatchesKeyword(item forwardingInventoryItem, keyword string) bool { - values := []string{ - item.Rule.Family, item.Rule.Protocol, item.Rule.Port, item.Rule.TargetIP, - item.Rule.TargetPort, item.Rule.Interface, item.SyncStatus(), - } - for _, value := range values { - if strings.Contains(strings.ToLower(value), keyword) { - return true +func (s *ForwardingService) OperateRules(request dto.ForwardRuleOperate) (dto.FilterChainOperationResponse, error) { + count := 0 + for _, rule := range request.Rules { + if rule.Operation == "add" { + count += strings.Count(rule.Protocol, "/") + 1 + } + if count > filter.MaxAtomicExpansion { + return dto.FilterChainOperationResponse{}, fmt.Errorf("create or import at most %d rules per batch (after expansion)", filter.MaxAtomicExpansion) } } - return false -} - -func (s *ForwardingService) OperateRules(request dto.ForwardRuleOperate) (dto.FilterChainOperationResponse, error) { - labels := make([]string, len(request.Rules)) operation := task.TaskCreate - for i, rule := range request.Rules { - labels[i] = fmt.Sprintf("[%d/%d] %s %s %s %s -> %s:%s", i+1, len(request.Rules), rule.Operation, rule.Family, rule.Protocol, rule.Port, rule.TargetIP, rule.TargetPort) + for _, rule := range request.Rules { if rule.Operation != "add" { operation = task.TaskUpdate } @@ -203,48 +192,26 @@ func (s *ForwardingService) OperateRules(request dto.ForwardRuleOperate) (dto.Fi if forwardingOperationsOnlyRemove(request.Rules) { operation = task.TaskDelete } - return queueFirewallRuleTask(firewallTaskForwarding, operation, labels, func(ctx context.Context) error { - return s.operateRules(ctx, request) - }) -} - -func (s *ForwardingService) operateRules(ctx context.Context, request dto.ForwardRuleOperate) error { - forwardingMutationMu.Lock() - defer forwardingMutationMu.Unlock() - if err := ctx.Err(); err != nil { - return err - } - stored, err := s.rules.List(ctx) + taskItem, err := task.NewTask(firewallTaskName(operation, firewallTaskForwarding, ""), operation, task.TaskScopeFirewall, "", 0) if err != nil { - return err - } - desired, err := applyForwardingOperations(forwardingRulesFromModels(stored), request.Rules) - if errors.Is(err, forwarding.ErrRuleExists) { - return buserr.New("ErrRecordExist") - } else if err != nil { - return err - } - if err := s.rules.ReplaceAll(ctx, forwardingRuleModels(desired)); err != nil { - return err + return dto.FilterChainOperationResponse{}, err } - if err := s.reconcile(desired); err != nil { - recordForwardingSyncError(err) - if request.ForceDelete && forwardingOperationsOnlyRemove(request.Rules) { - if global.LOG != nil { - global.LOG.Error(err) - } - return nil - } - return err + taskItem.AddSubTaskWithOps(taskItem.Name, func(t *task.Task) error { + return s.operateRules(t.TaskCtx, request, t) + }, nil, 0, 0) + if err := taskRepo.Save(context.Background(), taskItem.Task); err != nil { + taskItem.LogFailedWithErr(taskItem.Name, err) + closeUnstartedFirewallTask(taskItem) + return dto.FilterChainOperationResponse{}, err } - recordForwardingSyncError(nil) - return nil + go func() { _ = taskItem.Execute() }() + return dto.FilterChainOperationResponse{TaskID: taskItem.TaskID, Queued: true}, nil } func (s *ForwardingService) Enable() error { forwardingMutationMu.Lock() defer forwardingMutationMu.Unlock() - manager, err := s.managerFactory() + manager, err := s.clientFactory() if err != nil { recordForwardingSyncError(err) return err @@ -253,7 +220,7 @@ func (s *ForwardingService) Enable() error { recordForwardingSyncError(err) return err } - if err := s.activateManager(manager); err != nil { + if err := s.initializeForwarding(manager); err != nil { recordForwardingSyncError(err) return err } @@ -262,14 +229,12 @@ func (s *ForwardingService) Enable() error { recordForwardingSyncError(err) return err } - err = manager.Reconcile(forwardingRulesFromModels(rules)) + err = manager.ReplaceRules(forwardingRulesFromModels(rules)) recordForwardingSyncError(err) return err } -func (s *ForwardingService) QueueInitialization( - request dto.FirewallInitializationTask, -) (dto.FilterChainOperationResponse, error) { +func (s *ForwardingService) QueueInitialization(request dto.FirewallInitializationTask) (dto.FilterChainOperationResponse, error) { if err := task.CheckScopeTaskIsExecuting(task.TaskScopeFirewall, 0); err != nil { return dto.FilterChainOperationResponse{}, err } @@ -277,13 +242,13 @@ func (s *ForwardingService) QueueInitialization( if err != nil { return dto.FilterChainOperationResponse{}, fmt.Errorf("create forwarding initialization task: %w", err) } - var manager *forwarding.Manager + var manager forwarding.Adapter var backend string taskItem.AddSubTask(i18n.GetMsgByKey("FirewallEnableForwardingStep"), func(t *task.Task) error { forwardingMutationMu.Lock() defer forwardingMutationMu.Unlock() var err error - manager, err = s.managerFactory() + manager, err = s.clientFactory() if err != nil { recordForwardingSyncError(err) return err @@ -294,7 +259,7 @@ func (s *ForwardingService) QueueInitialization( recordForwardingSyncError(err) return err } - if err := s.activateManager(manager); err != nil { + if err := s.initializeForwarding(manager); err != nil { recordForwardingSyncError(err) return err } @@ -308,7 +273,7 @@ func (s *ForwardingService) QueueInitialization( recordForwardingSyncError(err) return err } - err = manager.Reconcile(forwardingRulesFromModels(rules)) + err = manager.ReplaceRules(forwardingRulesFromModels(rules)) recordForwardingSyncError(err) return err }, nil) @@ -319,131 +284,43 @@ func (s *ForwardingService) QueueInitialization( return dto.FilterChainOperationResponse{TaskID: taskItem.TaskID, Queued: true}, nil } -func (s *ForwardingService) Restore(ctx context.Context) error { - forwardingMutationMu.Lock() - defer forwardingMutationMu.Unlock() - enabled, err := s.forwardingEnabled() - if err != nil || !enabled { - if err != nil { - recordForwardingSyncError(err) - } - return err - } - manager, err := s.managerFactory() - if err != nil { - recordForwardingSyncError(err) - return err - } - stored, err := s.rules.List(ctx) - if err != nil { - recordForwardingSyncError(err) - return err - } - if err := s.activateManager(manager); err != nil { - recordForwardingSyncError(err) - return err - } - err = manager.Reconcile(forwardingRulesFromModels(stored)) - recordForwardingSyncError(err) - return err -} - -func (s *ForwardingService) reconcile(rules []forwarding.Rule) error { - manager, err := s.managerFactory() - if err != nil { - return err - } - return s.reconcileWithManager(manager, rules) -} - -func (s *ForwardingService) reconcileWithManager(manager *forwarding.Manager, rules []forwarding.Rule) error { - enabled, err := s.forwardingEnabled() - if err != nil || !enabled { - return err - } - if err := s.activateManager(manager); err != nil { - return err - } - return manager.Reconcile(rules) -} - -func (s *ForwardingService) activateManager(manager *forwarding.Manager) error { - if err := s.saveForwardingBackend(manager.Name()); err != nil { - return err - } - return manager.Enable() -} - -func (s *ForwardingService) forwardingEnabled() (bool, error) { - if s.enabled != nil { - return s.enabled() - } - return forwardingPersistedEnabled() -} - -func (s *ForwardingService) saveForwardingBackend(backend string) error { - if s.persistBackend != nil { - return s.persistBackend(backend) - } - return settingRepo.UpdateOrCreate(constant.FirewallForwardingBackendKey, backend) +func NewIForwardingService() IForwardingService { + return newForwardingService() } -func (s *ForwardingService) persistForwardingEnabled() error { - if s.markEnabled != nil { - return s.markEnabled() +func loadForwardPolicy(command string) (string, error) { + if !cmd.Which(command) { + command += "-nft" + if !cmd.Which(command) { + return "", nil + } } - return settingRepo.UpdateOrCreate(constant.FirewallForwardingInitializedKey, constant.StatusEnable) -} - -func forwardingPersistedEnabled() (bool, error) { - status, err := settingRepo.GetValueByKey(constant.FirewallForwardingInitializedKey) - return status == constant.StatusEnable, err -} - -func forwardingRulesFromModels(stored []model.ForwardingRule) []forwarding.Rule { - rules := make([]forwarding.Rule, 0, len(stored)) - for _, rule := range stored { - rules = append(rules, forwarding.Rule{ - Family: rule.Family, Protocol: rule.Protocol, Port: rule.Port, TargetIP: rule.TargetIP, - TargetPort: rule.TargetPort, Interface: rule.Interface, - }) + output, err := cmd.NewCommandMgr(cmd.WithTimeout(5*time.Second)).RunWithOptionalSudoAndStdout(command, "-t", "filter", "-w", "2", "-S", "FORWARD") + if err != nil { + return "", err } - return rules -} - -func forwardingRuleModels(rules []forwarding.Rule) []model.ForwardingRule { - stored := make([]model.ForwardingRule, 0, len(rules)) - for _, rule := range rules { - stored = append(stored, model.ForwardingRule{ - Family: rule.Family, Protocol: rule.Protocol, Port: rule.Port, TargetIP: rule.TargetIP, - TargetPort: rule.TargetPort, Interface: rule.Interface, - }) + for _, line := range strings.Split(output, "\n") { + fields := strings.Fields(line) + if len(fields) == 3 && fields[0] == "-P" && fields[1] == "FORWARD" { + if fields[2] != "ACCEPT" && fields[2] != "DROP" { + return "", fmt.Errorf("unexpected FORWARD policy: %s", fields[2]) + } + return fields[2], nil + } } - return stored + return "", errors.New("FORWARD default policy was not found") } -type forwardingInventoryItem struct { - ID uint - Rule forwarding.Rule - IsDesired bool - IsRuntime bool -} - -func (i forwardingInventoryItem) SyncStatus() string { - switch { - case i.IsDesired && i.IsRuntime: - return forwardingSyncConverged - case i.IsDesired: - return forwardingSyncMissing - default: - return forwardingSyncRuntimeOnly +func lastForwardingSyncError() string { + forwardingSyncStateMu.RLock() + defer forwardingSyncStateMu.RUnlock() + if forwardingLastSyncErr == nil { + return "" } + return forwardingLastSyncErr.Error() } -func mergeForwardingInventory( - stored []model.ForwardingRule, - runtime []forwarding.Rule, -) ([]forwardingInventoryItem, error) { +func mergeForwardingInventory(stored []model.ForwardingRule, runtime []forwarding.Rule) ([]forwardingInventoryItem, error) { items := make([]forwardingInventoryItem, 0, len(stored)+len(runtime)) byIdentity := make(map[string]int, len(stored)+len(runtime)) for _, record := range stored { @@ -474,104 +351,202 @@ func mergeForwardingInventory( return items, nil } -func recordForwardingSyncError(err error) { - forwardingSyncStateMu.Lock() - forwardingLastSyncErr = err - forwardingSyncStateMu.Unlock() -} - -func lastForwardingSyncError() string { - forwardingSyncStateMu.RLock() - defer forwardingSyncStateMu.RUnlock() - if forwardingLastSyncErr == nil { - return "" +func forwardingRuleMatchesKeyword(item forwardingInventoryItem, keyword string) bool { + values := []string{ + item.Rule.Family, item.Rule.Protocol, item.Rule.Port, item.Rule.TargetIP, + item.Rule.TargetPort, item.Rule.Interface, item.SyncStatus(), } - return forwardingLastSyncErr.Error() + for _, value := range values { + if strings.Contains(strings.ToLower(value), keyword) { + return true + } + } + return false } -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 := forwarding.NormalizeRule(rule) - if err != nil { - return nil, fmt.Errorf("normalize persisted forwarding rule: %w", err) - } - desired = append(desired, normalized) +func (s *ForwardingService) operateRules(ctx context.Context, request dto.ForwardRuleOperate, t *task.Task) (resultErr error) { + forwardingMutationMu.Lock() + defer forwardingMutationMu.Unlock() + if err := ctx.Err(); err != nil { + return err + } + type operationBatch struct { + operation forwarding.OperationType + rules []forwarding.Rule } - for _, operation := range requested { + groups := make([]operationBatch, 0) + for _, operation := range request.Rules { + kind := forwarding.OperationType(operation.Operation) + if kind != forwarding.OperationAdd && kind != forwarding.OperationRemove { + return fmt.Errorf("unsupported forwarding operation %q", operation.Operation) + } + if len(groups) == 0 || groups[len(groups)-1].operation != kind { + groups = append(groups, operationBatch{operation: kind}) + } for _, protocol := range strings.Split(operation.Protocol, "/") { rule, err := forwarding.NormalizeRule(forwarding.Rule{ - Family: operation.Family, Protocol: protocol, Port: operation.Port, TargetIP: operation.TargetIP, - TargetPort: operation.TargetPort, Interface: operation.Interface, + Family: operation.Family, Protocol: protocol, Port: operation.Port, + TargetIP: operation.TargetIP, TargetPort: operation.TargetPort, Interface: operation.Interface, }) if err != nil { - return nil, err - } - index := forwardingRuleIndex(desired, rule) - switch forwarding.OperationType(operation.Operation) { - case forwarding.OperationAdd: - if index >= 0 { - return nil, forwarding.ErrRuleExists - } - desired = append(desired, rule) - case forwarding.OperationRemove: - if index >= 0 { - desired = append(desired[:index], desired[index+1:]...) - } - default: - return nil, fmt.Errorf("unsupported forwarding operation %q", operation.Operation) + return err } + groups[len(groups)-1].rules = append(groups[len(groups)-1].rules, rule) } } - return desired, nil -} - -func forwardingRuleIndex(rules []forwarding.Rule, wanted forwarding.Rule) int { - wantedIdentity := wanted.Identity() - for index, rule := range rules { - if rule.Identity() == wantedIdentity { - return index - } - } - return -1 -} - -func forwardingOperationsOnlyRemove(operations []dto.ForwardRuleOperation) bool { - if len(operations) == 0 { - return false + stored, err := s.rules.List(ctx) + if err != nil { + return err } - for _, operation := range operations { - if operation.Operation != string(forwarding.OperationRemove) { - return false + byIdentity := make(map[string]model.ForwardingRule, len(stored)) + for index, rule := range forwardingRulesFromModels(stored) { + normalized, err := forwarding.NormalizeRule(rule) + if err != nil { + return err + } + byIdentity[normalized.Identity()] = stored[index] + } + succeeded, failed, skipped := 0, 0, 0 + var nativeFailure error + defer func() { + recordForwardingSyncError(errors.Join(resultErr, nativeFailure)) + if t != nil { + t.Log(i18n.GetMsgWithMap("FirewallRuleOperationResult", map[string]interface{}{"succeeded": succeeded, "failed": failed})) + if skipped > 0 { + t.Logf("%s: %d", i18n.GetMsgByKey("FirewallCreateRuleSkipped"), skipped) + } + } + }() + record := func(operation forwarding.OperationType, rule forwarding.Rule, status string, cause error) { + label := fmt.Sprintf("%s %s %s %s -> %s:%s", operation, rule.Family, rule.Protocol, rule.Port, rule.TargetIP, rule.TargetPort) + switch status { + case "skipped": + skipped++ + if t != nil { + t.Logf("%s %s: %v", label, i18n.GetMsgByKey("FirewallCreateRuleSkipped"), cause) + } + case "failed": + failed++ + if t != nil { + t.LogFailedWithErr(label, cause) + } + default: + succeeded++ + if t != nil { + t.LogSuccess(label) + } } } - return true -} - -func newForwardingManager() (*forwarding.Manager, error) { - return newForwardingManagerFor(configuredForwardingBackend()) -} - -func configuredForwardingBackend() string { - selected, _ := settingRepo.GetValueByKey(constant.FirewallForwardingBackendKey) - selected = strings.TrimSpace(selected) - if selected == "" { - return constant.FirewallProviderIptables - } - return selected -} - -func newForwardingManagerFor(backend string) (*forwarding.Manager, error) { - client, err := lifecycle.NewClientFor(backend) - if err != nil { - return nil, fmt.Errorf( - "%w: selected forwarding backend %s: %w", - errForwardingBackendUnavailable, backend, err, - ) + if len(request.Rules) == 2 && len(groups) == 2 && groups[0].operation == forwarding.OperationRemove && groups[1].operation == forwarding.OperationAdd { + old := make(map[string]bool, len(groups[0].rules)) + for _, rule := range groups[0].rules { + old[rule.Identity()] = true + } + unchanged := len(old) == len(groups[1].rules) + duplicate := false + for _, rule := range groups[1].rules { + key := rule.Identity() + unchanged = unchanged && old[key] + if _, exists := byIdentity[key]; exists && !old[key] { + duplicate = true + } + } + if unchanged || duplicate { + for _, group := range groups { + for _, rule := range group.rules { + record(group.operation, rule, "skipped", buserr.New("ErrRecordExist")) + } + } + return nil + } } - adapter, err := forwarding.New(client.Name()) - if err != nil { - return nil, err + var client forwarding.Adapter + var failures []error + for _, group := range groups { + byFamily := make(map[string][]forwarding.Rule, 2) + seen := make(map[string]bool, len(group.rules)) + for _, rule := range group.rules { + key := rule.Identity() + _, exists := byIdentity[key] + if seen[key] || (group.operation == forwarding.OperationAdd && exists) { + record(group.operation, rule, "skipped", buserr.New("ErrRecordExist")) + continue + } + seen[key] = true + byFamily[rule.Family] = append(byFamily[rule.Family], rule) + } + for _, family := range []string{forwarding.FamilyIPv4, forwarding.FamilyIPv6} { + rules := byFamily[family] + if len(rules) == 0 { + continue + } + err := ctx.Err() + if err == nil && client == nil { + var enabled bool + enabled, err = s.forwardingEnabled() + if err == nil && !enabled { + err = fmt.Errorf("%w: forwarding is not initialized", filter.ErrProviderUnavailable) + } + if err == nil { + client, err = s.clientFactory() + } + } + if err == nil { + if group.operation == forwarding.OperationAdd { + err = client.CreateRules(ctx, rules) + } else { + err = client.DeleteRules(ctx, rules) + } + } + if err != nil { + nativeFailure = errors.Join(nativeFailure, err) + if !request.ForceDelete || !forwardingOperationsOnlyRemove(request.Rules) || ctx.Err() != nil { + failures = append(failures, err) + for _, rule := range rules { + record(group.operation, rule, "failed", err) + } + continue + } + if t != nil { + t.Logf("force delete database records: %v", err) + } + } + for start := 0; start < len(rules); start += 500 { + batch := rules[start:min(start+500, len(rules))] + records := make([]model.ForwardingRule, 0, len(batch)) + ids := make([]uint, 0, len(batch)) + for _, rule := range batch { + if group.operation == forwarding.OperationAdd { + records = append(records, model.ForwardingRule{Family: rule.Family, Protocol: rule.Protocol, Port: rule.Port, TargetIP: rule.TargetIP, TargetPort: rule.TargetPort, Interface: rule.Interface}) + } else if stored, exists := byIdentity[rule.Identity()]; exists { + ids = append(ids, stored.ID) + } + } + if group.operation == forwarding.OperationAdd { + err = s.rules.CreateBatch(context.WithoutCancel(ctx), records) + } else { + err = s.rules.DeleteBatch(context.WithoutCancel(ctx), ids) + } + if err != nil { + failures = append(failures, err) + } + for index, rule := range batch { + if err != nil { + record(group.operation, rule, "failed", err) + continue + } + if group.operation == forwarding.OperationAdd { + byIdentity[rule.Identity()] = records[index] + } else { + delete(byIdentity, rule.Identity()) + } + record(group.operation, rule, "succeeded", nil) + } + } + } + if group.operation == forwarding.OperationRemove && len(failures) > 0 { + return errors.Join(failures...) + } } - return forwarding.NewManager(adapter, client), nil + return errors.Join(failures...) } diff --git a/agent/app/service/website_domain.go b/agent/app/service/website_domain.go index afea4d1d8b4e..89af8838d287 100644 --- a/agent/app/service/website_domain.go +++ b/agent/app/service/website_domain.go @@ -8,6 +8,7 @@ import ( "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" "github.com/1Panel-dev/1Panel/agent/utils/files" "path" "strconv" @@ -32,7 +33,9 @@ func (w WebsiteService) CreateWebsiteDomain(create request.WebsiteDomainCreate) return nil, err } go func() { - _ = ensureFirewallPorts(addPorts) + if err := ensureFirewallPorts(addPorts); err != nil { + global.LOG.Errorf("allow website firewall ports failed: %v", err) + } }() nginxInstall, err := getAppInstallByKey(constant.AppOpenresty) diff --git a/agent/i18n/lang/en.yaml b/agent/i18n/lang/en.yaml index c566d0828412..0c73f9116a25 100644 --- a/agent/i18n/lang/en.yaml +++ b/agent/i18n/lang/en.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'Persist Docker port guard status' ErrFirewallRuleScopeChange: "The current firewall does not support changing a rule's scope (such as its IPv4/IPv6 address family). Please create a new rule." FirewallWhitelistReleased: "{{ .name }}: whitelist protection released; allow rule retained. To close the port, delete the rule manually from the rule list" FirewallWhitelistRequired: "{{ .name }}: protected by mandatory system port rules" + +ErrFirewallBackendCleanupRequired: "The current backend {{ .current }} still contains 1Panel rules. Clean it up before switching to {{ .target }}." +ErrDockerIPv4ForwardingDisabled: "IPv4 forwarding is disabled. Set net.ipv4.ip_forward=1 before using Docker's firewall backend." +ErrFirewallRuleSavedApplyFailed: "The new rule configuration was saved, but applying it to the firewall failed. Retry by synchronizing: {{ .detail }}" diff --git a/agent/i18n/lang/es-ES.yaml b/agent/i18n/lang/es-ES.yaml index 4f1896958925..7283f97e27a4 100644 --- a/agent/i18n/lang/es-ES.yaml +++ b/agent/i18n/lang/es-ES.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'Guardar el estado de protección de puertos de ErrFirewallRuleScopeChange: "El cortafuegos actual no permite cambiar el ámbito de una regla (como su familia de direcciones IPv4/IPv6). Cree una regla nueva." FirewallWhitelistReleased: "{{ .name }}: protección de la lista de permitidos retirada; se conserva la regla de permiso. Para cerrar el puerto, elimine la regla manualmente de la lista" FirewallWhitelistRequired: "{{ .name }}: protegido por las reglas de puertos obligatorios del sistema" + +ErrFirewallBackendCleanupRequired: "El backend actual {{ .current }} aún contiene reglas de 1Panel. Elimínelas antes de cambiar a {{ .target }}." +ErrDockerIPv4ForwardingDisabled: "El reenvío IPv4 está desactivado. Configure net.ipv4.ip_forward=1 antes de usar el backend de cortafuegos de Docker." +ErrFirewallRuleSavedApplyFailed: "Se guardó la configuración de la nueva regla, pero no se pudo aplicar al cortafuegos. Vuelva a intentarlo mediante la sincronización: {{ .detail }}" diff --git a/agent/i18n/lang/fa.yaml b/agent/i18n/lang/fa.yaml index ff98e50b78ae..5346371427d5 100644 --- a/agent/i18n/lang/fa.yaml +++ b/agent/i18n/lang/fa.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'ذخیره وضعیت محافظت پورت Doc ErrFirewallRuleScopeChange: "فایروال فعلی از تغییر محدودهٔ قانون (مانند خانوادهٔ آدرس IPv4/IPv6) پشتیبانی نمی‌کند. لطفاً یک قانون جدید ایجاد کنید." FirewallWhitelistReleased: "{{ .name }}: حفاظت فهرست مجاز برداشته شد؛ قانون اجازه حفظ می‌شود. برای بستن پورت، قانون را به‌صورت دستی از فهرست قوانین حذف کنید" FirewallWhitelistRequired: "{{ .name }}: توسط قوانین پورت‌های ضروری سیستم محافظت می‌شود" + +ErrFirewallBackendCleanupRequired: "بک‌اند فعلی {{ .current }} هنوز شامل قوانین 1Panel است. پیش از تغییر به {{ .target }} آن‌ها را پاک کنید." +ErrDockerIPv4ForwardingDisabled: "ارسال IPv4 غیرفعال است. پیش از استفاده از بک‌اند فایروال Docker، مقدار net.ipv4.ip_forward=1 را تنظیم کنید." +ErrFirewallRuleSavedApplyFailed: "پیکربندی قانون جدید ذخیره شد، اما اعمال آن در دیوار آتش ناموفق بود. با همگام‌سازی دوباره تلاش کنید: {{ .detail }}" diff --git a/agent/i18n/lang/ja.yaml b/agent/i18n/lang/ja.yaml index ef4a768a276f..badb08d986ec 100644 --- a/agent/i18n/lang/ja.yaml +++ b/agent/i18n/lang/ja.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'Docker ポート保護状態を保存' ErrFirewallRuleScopeChange: "現在のファイアウォールでは、ルールの適用範囲(IPv4/IPv6 アドレスファミリーなど)を変更できません。新しいルールを作成してください。" FirewallWhitelistReleased: "{{ .name }}:許可リストの保護を解除しました。許可ルールは保持されます。ポートを閉じるには、ルール一覧から手動で削除してください" FirewallWhitelistRequired: "{{ .name }}:システム必須ポートのルールで保護されています" + +ErrFirewallBackendCleanupRequired: "現在のバックエンド {{ .current }} に 1Panel ルールが残っています。{{ .target }} に切り替える前に削除してください。" +ErrDockerIPv4ForwardingDisabled: "IPv4 転送が無効です。Docker のファイアウォールバックエンドを使用する前に net.ipv4.ip_forward=1 を設定してください。" +ErrFirewallRuleSavedApplyFailed: "新しいルール設定は保存されましたが、ファイアウォールへの適用に失敗しました。同期で再試行してください:{{ .detail }}" diff --git a/agent/i18n/lang/ko.yaml b/agent/i18n/lang/ko.yaml index 161920e7b5dd..fa0484d22c1f 100644 --- a/agent/i18n/lang/ko.yaml +++ b/agent/i18n/lang/ko.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'Docker 포트 보호 상태 저장' ErrFirewallRuleScopeChange: "현재 방화벽에서는 규칙의 적용 범위(예: IPv4/IPv6 주소 패밀리)를 변경할 수 없습니다. 새 규칙을 생성하세요." FirewallWhitelistReleased: "{{ .name }}: 허용 목록 보호가 해제되었으며 허용 규칙은 유지됩니다. 포트를 닫으려면 규칙 목록에서 수동으로 삭제하세요" FirewallWhitelistRequired: "{{ .name }}: 시스템 필수 포트 규칙으로 보호됩니다" + +ErrFirewallBackendCleanupRequired: "현재 백엔드 {{ .current }}에 1Panel 규칙이 남아 있습니다. {{ .target }}로 전환하기 전에 정리하세요." +ErrDockerIPv4ForwardingDisabled: "IPv4 전달이 비활성화되어 있습니다. Docker 방화벽 백엔드를 사용하기 전에 net.ipv4.ip_forward=1을 설정하세요." +ErrFirewallRuleSavedApplyFailed: "새 규칙 설정이 저장되었지만 방화벽에 적용하지 못했습니다. 동기화하여 다시 시도하세요: {{ .detail }}" diff --git a/agent/i18n/lang/lo.yaml b/agent/i18n/lang/lo.yaml index 5d27c049c81e..9b22c24e2e2b 100644 --- a/agent/i18n/lang/lo.yaml +++ b/agent/i18n/lang/lo.yaml @@ -710,3 +710,7 @@ FirewallPersistDockerGuardStep: 'ບັນທຶກສະຖານະປ້ອ ErrFirewallRuleScopeChange: "ໄຟວໍປັດຈຸບັນບໍ່ຮອງຮັບການປ່ຽນຂອບເຂດຂອງກົດ (ເຊັ່ນ ຕະກູນທີ່ຢູ່ IPv4/IPv6). ກະລຸນາສ້າງກົດໃໝ່." FirewallWhitelistReleased: "{{ .name }}: ຍົກເລີກການປ້ອງກັນລາຍຊື່ທີ່ອະນຸຍາດແລ້ວ; ຍັງຄົງກົດອະນຸຍາດໄວ້. ຫາກຕ້ອງການປິດພອດ ໃຫ້ລຶບກົດດ້ວຍຕົນເອງຈາກລາຍການກົດ" FirewallWhitelistRequired: "{{ .name }}: ປ້ອງກັນໂດຍກົດພອດທີ່ຈຳເປັນຂອງລະບົບ" + +ErrFirewallBackendCleanupRequired: "ແບັກເອນປັດຈຸບັນ {{ .current }} ຍັງມີກົດຂອງ 1Panel. ກະລຸນາລຶບອອກກ່ອນປ່ຽນໄປ {{ .target }}." +ErrDockerIPv4ForwardingDisabled: "ການສົ່ງຕໍ່ IPv4 ຖືກປິດ. ກະລຸນາຕັ້ງ net.ipv4.ip_forward=1 ກ່ອນໃຊ້ແບັກເອນໄຟວໍຂອງ Docker." +ErrFirewallRuleSavedApplyFailed: "ບັນທຶກການຕັ້ງຄ່າກົດໃໝ່ແລ້ວ ແຕ່ນຳໃຊ້ກັບໄຟວໍບໍ່ສຳເລັດ. ລອງອີກຄັ້ງດ້ວຍການຊິງຂໍ້ມູນ: {{ .detail }}" diff --git a/agent/i18n/lang/ms.yaml b/agent/i18n/lang/ms.yaml index 8adaf1beebeb..c94d7671e10d 100644 --- a/agent/i18n/lang/ms.yaml +++ b/agent/i18n/lang/ms.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'Simpan status perlindungan port Docker' ErrFirewallRuleScopeChange: "Tembok api semasa tidak menyokong perubahan skop peraturan (seperti keluarga alamat IPv4/IPv6). Sila cipta peraturan baharu." FirewallWhitelistReleased: "{{ .name }}: perlindungan senarai dibenarkan telah dilepaskan; peraturan izin dikekalkan. Untuk menutup port, padamkan peraturan secara manual daripada senarai peraturan" FirewallWhitelistRequired: "{{ .name }}: dilindungi oleh peraturan port wajib sistem" + +ErrFirewallBackendCleanupRequired: "Bahagian belakang semasa {{ .current }} masih mengandungi peraturan 1Panel. Buangkannya sebelum beralih kepada {{ .target }}." +ErrDockerIPv4ForwardingDisabled: "Pemajuan IPv4 dilumpuhkan. Tetapkan net.ipv4.ip_forward=1 sebelum menggunakan bahagian belakang tembok api Docker." +ErrFirewallRuleSavedApplyFailed: "Konfigurasi peraturan baharu telah disimpan, tetapi gagal digunakan pada tembok api. Cuba lagi melalui penyegerakan: {{ .detail }}" diff --git a/agent/i18n/lang/pt-BR.yaml b/agent/i18n/lang/pt-BR.yaml index 2674a1781866..b8534f7c47f8 100644 --- a/agent/i18n/lang/pt-BR.yaml +++ b/agent/i18n/lang/pt-BR.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'Salvar o status da proteção de portas do Dock ErrFirewallRuleScopeChange: "O firewall atual não permite alterar o escopo de uma regra (como a família de endereços IPv4/IPv6). Crie uma nova regra." FirewallWhitelistReleased: "{{ .name }}: proteção da lista de permissões removida; regra de permissão mantida. Para fechar a porta, exclua a regra manualmente da lista" FirewallWhitelistRequired: "{{ .name }}: protegido pelas regras de portas obrigatórias do sistema" + +ErrFirewallBackendCleanupRequired: "O backend atual {{ .current }} ainda contém regras do 1Panel. Remova-as antes de mudar para {{ .target }}." +ErrDockerIPv4ForwardingDisabled: "O encaminhamento IPv4 está desativado. Defina net.ipv4.ip_forward=1 antes de usar o backend de firewall do Docker." +ErrFirewallRuleSavedApplyFailed: "A configuração da nova regra foi salva, mas não pôde ser aplicada ao firewall. Tente novamente por meio da sincronização: {{ .detail }}" diff --git a/agent/i18n/lang/ru.yaml b/agent/i18n/lang/ru.yaml index a00d4134ecbe..6a5cd320c65b 100644 --- a/agent/i18n/lang/ru.yaml +++ b/agent/i18n/lang/ru.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'Сохранить состояние защи ErrFirewallRuleScopeChange: "Текущий межсетевой экран не поддерживает изменение области действия правила (например, семейства адресов IPv4/IPv6). Создайте новое правило." FirewallWhitelistReleased: "{{ .name }}: защита списка разрешённых портов снята; разрешающее правило сохранено. Чтобы закрыть порт, удалите правило вручную из списка правил" FirewallWhitelistRequired: "{{ .name }}: защищён обязательными правилами системных портов" + +ErrFirewallBackendCleanupRequired: "В текущем бэкенде {{ .current }} остались правила 1Panel. Удалите их перед переключением на {{ .target }}." +ErrDockerIPv4ForwardingDisabled: "Пересылка IPv4 отключена. Перед использованием бэкенда межсетевого экрана Docker установите net.ipv4.ip_forward=1." +ErrFirewallRuleSavedApplyFailed: "Настройки нового правила сохранены, но применить их к межсетевому экрану не удалось. Повторите попытку с помощью синхронизации: {{ .detail }}" diff --git a/agent/i18n/lang/tr.yaml b/agent/i18n/lang/tr.yaml index 10972e8b67e6..af870cccba24 100644 --- a/agent/i18n/lang/tr.yaml +++ b/agent/i18n/lang/tr.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: 'Docker bağlantı noktası koruma durumunu kayd ErrFirewallRuleScopeChange: "Mevcut güvenlik duvarı, kuralın kapsamını (IPv4/IPv6 adres ailesi gibi) değiştirmeyi desteklemiyor. Lütfen yeni bir kural oluşturun." FirewallWhitelistReleased: "{{ .name }}: izin listesi koruması kaldırıldı; izin kuralı korundu. Portu kapatmak için kuralı kural listesinden elle silin" FirewallWhitelistRequired: "{{ .name }}: zorunlu sistem portu kuralları tarafından korunuyor" + +ErrFirewallBackendCleanupRequired: "Mevcut {{ .current }} arka ucunda hâlâ 1Panel kuralları var. {{ .target }} arka ucuna geçmeden önce bunları temizleyin." +ErrDockerIPv4ForwardingDisabled: "IPv4 yönlendirmesi devre dışı. Docker güvenlik duvarı arka ucunu kullanmadan önce net.ipv4.ip_forward=1 ayarını yapın." +ErrFirewallRuleSavedApplyFailed: "Yeni kural yapılandırması kaydedildi, ancak güvenlik duvarına uygulanamadı. Eşitleme yaparak yeniden deneyin: {{ .detail }}" diff --git a/agent/i18n/lang/zh-Hant.yaml b/agent/i18n/lang/zh-Hant.yaml index 97fbe6c692ec..23908761c198 100644 --- a/agent/i18n/lang/zh-Hant.yaml +++ b/agent/i18n/lang/zh-Hant.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: '儲存 Docker 連接埠防護狀態' ErrFirewallRuleScopeChange: "目前的防火牆不支援修改規則的作用範圍(如 IPv4/IPv6 位址族),請建立新規則。" FirewallWhitelistReleased: "{{ .name }}:已解除白名單保護,放行規則保留;如需關閉連接埠,請在規則清單手動刪除" FirewallWhitelistRequired: "{{ .name }}:由系統必要連接埠規則保護" + +ErrFirewallBackendCleanupRequired: "目前後端 {{ .current }} 中仍有 1Panel 規則,請先清理後再切換至 {{ .target }}。" +ErrDockerIPv4ForwardingDisabled: "IPv4 轉送尚未啟用,請先設定 net.ipv4.ip_forward=1,再使用 Docker 防火牆後端。" +ErrFirewallRuleSavedApplyFailed: "新規則設定已儲存,但套用至防火牆失敗。可透過同步重試:{{ .detail }}" diff --git a/agent/i18n/lang/zh.yaml b/agent/i18n/lang/zh.yaml index eaaa9bfcc746..ff15f8b8833b 100644 --- a/agent/i18n/lang/zh.yaml +++ b/agent/i18n/lang/zh.yaml @@ -719,3 +719,7 @@ FirewallPersistDockerGuardStep: "保存 Docker 端口防护状态" ErrFirewallRuleScopeChange: "当前防火墙不支持修改规则的作用范围(如 IPv4/IPv6 地址族),请新建规则。" FirewallWhitelistReleased: "{{ .name }}:已解除白名单保护,放行规则保留;如需关闭端口,请在规则列表手动删除" FirewallWhitelistRequired: "{{ .name }}:由系统必需端口规则保护" + +ErrFirewallBackendCleanupRequired: "当前后端 {{ .current }} 中仍有 1Panel 规则,请先清理后再切换到 {{ .target }}。" +ErrDockerIPv4ForwardingDisabled: "IPv4 转发未开启,请先设置 net.ipv4.ip_forward=1,再使用 Docker 防火墙后端。" +ErrFirewallRuleSavedApplyFailed: "新规则配置已保存,但应用到防火墙失败。可通过同步重试:{{ .detail }}" diff --git a/agent/init/firewall/firewall.go b/agent/init/firewall/firewall.go index f27b356f4fb5..b88c22078a0a 100644 --- a/agent/init/firewall/firewall.go +++ b/agent/init/firewall/firewall.go @@ -103,10 +103,12 @@ func repairIptablesBaseChains(clientName string) { if status != constant.StatusEnable { return } - manager := iptables_helper.Manager{ - LoadRequiredPorts: service.LoadRequiredFirewallPortWhiteList, + ports, err := service.LoadRequiredFirewallPortWhiteList() + if err != nil { + global.LOG.Warnf("load required firewall ports for base chain repair failed, err: %v", err) + return } - if err := manager.RepairBaseChains(); err != nil { + if err := iptables_helper.RepairBaseChains(ports); err != nil { global.LOG.Warnf("repair iptables base chains failed, err: %v", err) } } diff --git a/agent/init/migration/migrations/firewall_whitelist.go b/agent/init/migration/migrations/firewall_whitelist.go index 4dff4eb7aaa0..93783955e9f3 100644 --- a/agent/init/migration/migrations/firewall_whitelist.go +++ b/agent/init/migration/migrations/firewall_whitelist.go @@ -162,6 +162,9 @@ func migrateFirewallPortWhitelist(value string) ([]firewall.PortWhitelist, error for _, rule := range defaults { index, found := indexes[key(rule)] if !found { + if rule.Type == "" { + continue + } index = len(rules) indexes[key(rule)] = index rules = append(rules, rule) diff --git a/agent/init/migration/migrations/utils/host_firewall_transfer.go b/agent/init/migration/migrations/utils/host_firewall_transfer.go index 4a7603844be8..5914dd9e854c 100644 --- a/agent/init/migration/migrations/utils/host_firewall_transfer.go +++ b/agent/init/migration/migrations/utils/host_firewall_transfer.go @@ -2,6 +2,9 @@ package utils import ( "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" "errors" "fmt" "net/netip" @@ -142,7 +145,7 @@ func convertLegacyHostFirewallRecords(records []legacyHostFirewallRecord, provid } continue } - identity := item.PolicyKey() + identity := hostFirewallPolicyKey(item) if index, exists := byIdentity[identity]; exists { if item.Description != "" { converted[index].Description = item.Description @@ -336,10 +339,30 @@ func legacyIPOrPrefix(value string) bool { } func hostFirewallRuleModel(rule filter.FirewallRule) (model.FirewallRule, error) { - record, err := model.FirewallRuleFromDomain(rule) + normalized, err := filter.NormalizeRule(rule) if err != nil { return model.FirewallRule{}, err } + switch normalized.NativeKind { + case "", filter.NativeKindRule, filter.NativeKindZonePort, filter.NativeKindRichRule, filter.NativeKindUFWRule: + default: + return model.FirewallRule{}, fmt.Errorf("%w: native rule %q cannot be stored as a provider-neutral policy", filter.ErrUnsupportedScope, normalized.NativeKind) + } + record := model.FirewallRule{ + Family: string(normalized.Scope.Family), + Protocol: normalized.Protocol, + SourceAddress: normalized.SourceAddress, + SourcePort: normalized.SourcePort, + DestinationAddress: normalized.DestinationAddress, + DestinationPort: normalized.DestinationPort, + Interface: normalized.Interface, + ConnectionStates: strings.Join(normalized.ConnectionStates, ","), + Action: string(normalized.Action), + Description: normalized.Description, + } + if normalized.Scope.Provider == filter.ProviderFirewalld { + record.Priority = normalized.Priority + } record.UUID = uuid.NewString() record.Origin = constant.FirewallRuleOriginAdopted record.Owner = constant.FirewallRuleSourceUser @@ -347,6 +370,27 @@ func hostFirewallRuleModel(rule filter.FirewallRule) (model.FirewallRule, error) return record, nil } +func hostFirewallPolicyKey(rule model.FirewallRule) string { + payload, _ := json.Marshal(struct { + Family string `json:"family"` + Protocol string `json:"protocol"` + SourceAddress string `json:"sourceAddress,omitempty"` + SourcePort string `json:"sourcePort,omitempty"` + DestinationAddress string `json:"destinationAddress,omitempty"` + DestinationPort string `json:"destinationPort,omitempty"` + Interface string `json:"interface,omitempty"` + ConnectionStates string `json:"connectionStates,omitempty"` + Action string `json:"action"` + }{ + Family: rule.Family, Protocol: rule.Protocol, + SourceAddress: rule.SourceAddress, SourcePort: rule.SourcePort, + DestinationAddress: rule.DestinationAddress, DestinationPort: rule.DestinationPort, + Interface: rule.Interface, ConnectionStates: rule.ConnectionStates, Action: rule.Action, + }) + sum := sha256.Sum256(payload) + return hex.EncodeToString(sum[:]) +} + func importLegacyHostFirewallRules(tx *gorm.DB, rules []model.FirewallRule) error { var existing []model.FirewallRule if err := tx.Find(&existing).Error; err != nil { @@ -354,10 +398,10 @@ func importLegacyHostFirewallRules(tx *gorm.DB, rules []model.FirewallRule) erro } byIdentity := make(map[string]model.FirewallRule, len(existing)) for _, item := range existing { - byIdentity[item.PolicyKey()] = item + byIdentity[hostFirewallPolicyKey(item)] = item } for _, item := range rules { - identity := item.PolicyKey() + identity := hostFirewallPolicyKey(item) if current, exists := byIdentity[identity]; exists { if current.Description == "" && item.Description != "" { if err := tx.Model(&model.FirewallRule{}).Where("uuid = ?", current.UUID). diff --git a/agent/utils/cmd/cmdx.go b/agent/utils/cmd/cmdx.go index 76162016106e..3b85d8c53f19 100644 --- a/agent/utils/cmd/cmdx.go +++ b/agent/utils/cmd/cmdx.go @@ -27,6 +27,7 @@ type CommandHelper struct { outputFile string scriptPath string stdin io.Reader + stderr io.Writer env []string timeout time.Duration taskItem *task.Task @@ -360,6 +361,9 @@ func (c *CommandHelper) run(name string, arg ...string) (string, error) { cmd.Stdout = &stdout cmd.Stderr = &stderr } + if c.stderr != nil { + cmd.Stderr = io.MultiWriter(cmd.Stderr, c.stderr) + } env := os.Environ() env = append(env, c.env...) cmd.Env = env @@ -481,6 +485,11 @@ func WithStdin(stdin io.Reader) Option { s.stdin = stdin } } +func WithStderr(stderr io.Writer) Option { + return func(s *CommandHelper) { + s.stderr = stderr + } +} func WithEnv(env ...string) Option { return func(s *CommandHelper) { s.env = append(s.env, env...) diff --git a/agent/utils/firewall/docker_guard/inspect.go b/agent/utils/firewall/docker_guard/inspect.go new file mode 100644 index 000000000000..027a580ed206 --- /dev/null +++ b/agent/utils/firewall/docker_guard/inspect.go @@ -0,0 +1,247 @@ +package docker_guard + +import ( + "net/netip" + "path/filepath" + "slices" + "strconv" + "strings" + "time" + + "github.com/1Panel-dev/1Panel/agent/constant" + "github.com/1Panel-dev/1Panel/agent/utils/cmd" + "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" +) + +func ReadDNATRules(backend, family string) DNATRules { + manager := cmd.NewCommandMgr(cmd.WithTimeout(10*time.Second), cmd.WithEnv("LC_ALL=C")) + if backend == constant.FirewallProviderNftables { + tableFamily := "ip" + if family == constant.FirewallFamilyIPv6 { + tableFamily = "ip6" + } + tables, err := manager.RunWithOptionalSudoAndStdout("nft", "list", "tables") + if err != nil { + return DNATRules{} + } + if !strings.Contains(tables, "table "+tableFamily+" docker-bridges") { + return DNATRules{Inspected: true} + } + output, err := manager.RunWithOptionalSudoAndStdout("nft", "list", "table", tableFamily, "docker-bridges") + return DNATRules{Output: output, Inspected: err == nil} + } + commands, err := lifecycle.ResolveIptablesCommands() + if err != nil { + return DNATRules{} + } + executable := commands.IPv4 + if family == constant.FirewallFamilyIPv6 { + executable = commands.IPv6 + } + if executable == "" { + return DNATRules{} + } + output, err := manager.RunWithOptionalSudoAndStdout(executable, "-w", "-t", "nat", "-S") + return DNATRules{Output: output, Inspected: err == nil} +} + +func ReadProxyEndpoints() ProxyEndpoints { + manager := cmd.NewCommandMgr(cmd.WithTimeout(10*time.Second), cmd.WithEnv("LC_ALL=C")) + output, err := manager.RunWithStdout("ps", "-ww", "-eo", "args=") + if err != nil { + return ProxyEndpoints{} + } + return ProxyEndpoints{Items: parseDockerProxyEndpoints(output), Inspected: true} +} + +func parseDockerProxyEndpoints(output string) []ProxyEndpoint { + result := make([]ProxyEndpoint, 0) + for _, line := range strings.Split(output, "\n") { + fields := strings.Fields(line) + isProxy := slices.ContainsFunc(fields, func(field string) bool { return filepath.Base(field) == "docker-proxy" }) + if !isProxy { + continue + } + protocol := commandFlagValue(fields, "-proto") + hostIP := commandFlagValue(fields, "-host-ip") + hostPortValue := commandFlagValue(fields, "-host-port") + hostPort, err := strconv.ParseUint(hostPortValue, 10, 16) + if err != nil || (protocol != "tcp" && protocol != "udp") || hostIP == "" { + continue + } + hostIP = strings.TrimSpace(hostIP) + if address, err := netip.ParseAddr(hostIP); err == nil { + hostIP = address.String() + } + result = append(result, ProxyEndpoint{Protocol: protocol, HostIP: hostIP, HostPort: uint16(hostPort)}) + } + return result +} + +func commandFlagValue(fields []string, name string) string { + for i := 0; i < len(fields); i++ { + if fields[i] == name && i+1 < len(fields) { + return fields[i+1] + } + if strings.HasPrefix(fields[i], name+"=") { + return strings.TrimPrefix(fields[i], name+"=") + } + } + return "" +} + +func ProxyEndpointMatches(proxies []ProxyEndpoint, family, hostIP string, hostPort uint16, protocol string) bool { + for _, proxy := range proxies { + if proxy.Protocol == protocol && proxy.HostPort == hostPort && hostAddressMatches(proxy.HostIP, hostIP, family) { + return true + } + } + return false +} + +func DNATRuleMatches(backend, output string, family, hostIP string, hostPort uint16, protocol string) bool { + return InspectEndpoints(backend, family, DNATRules{Output: output}, ProxyEndpoints{}).DNATMatches(hostIP, hostPort, protocol) +} + +func DNATIngressReachable(backend, output string) bool { + if backend == constant.FirewallProviderNftables { + return strings.Contains(output, "hook prerouting") + } + for _, line := range strings.Split(output, "\n") { + fields := strings.Fields(line) + if len(fields) >= 4 && fields[0] == "-A" && fields[1] == "PREROUTING" && commandFlagValue(fields, "-j") == "DOCKER" { + return true + } + } + return false +} + +type EndpointInspection struct { + DNATInspected, ProxyInspected, IngressReachable bool + family string + dnat, proxies map[ProxyEndpoint]bool +} + +func InspectEndpoints(backend, family string, rules DNATRules, proxies ProxyEndpoints) EndpointInspection { + inspection := EndpointInspection{ + DNATInspected: rules.Inspected, ProxyInspected: proxies.Inspected, + IngressReachable: DNATIngressReachable(backend, rules.Output), family: family, + dnat: make(map[ProxyEndpoint]bool), proxies: make(map[ProxyEndpoint]bool), + } + for _, proxy := range proxies.Items { + if isWildcardHostAddress(proxy.HostIP, family) { + wildcard := proxy + wildcard.HostIP = "" + inspection.proxies[wildcard] = true + } + proxy.HostIP = normalizedEndpointAddress(proxy.HostIP) + inspection.proxies[proxy] = true + } + replacer := strings.NewReplacer("{", " ", "}", " ", ",", " ", ";", " ") + addressToken := "ip" + if family == constant.FirewallFamilyIPv6 { + addressToken = "ip6" + } + for line := range strings.SplitSeq(rules.Output, "\n") { + if backend != constant.FirewallProviderNftables { + fields := strings.Fields(line) + if commandFlagValue(fields, "-j") != "DNAT" { + continue + } + port, err := strconv.ParseUint(commandFlagValue(fields, "--dport"), 10, 16) + if err != nil || strconv.FormatUint(port, 10) != commandFlagValue(fields, "--dport") { + continue + } + address, _, _ := strings.Cut(commandFlagValue(fields, "-d"), "/") + inspection.dnat[ProxyEndpoint{Protocol: commandFlagValue(fields, "-p"), HostIP: normalizedEndpointAddress(address), HostPort: uint16(port)}] = true + continue + } + fields := strings.Fields(replacer.Replace(line)) + if !slices.Contains(fields, "dnat") { + continue + } + destination := "" + protocols := make([]string, 0, 1) + for i := 0; i+2 < len(fields); i++ { + if fields[i] == "meta" && fields[i+1] == "l4proto" { + protocols = append(protocols, fields[i+2]) + } + if destination == "" && fields[i] == addressToken && fields[i+1] == "daddr" { + destination = fields[i+2] + } + } + destination, _, _ = strings.Cut(destination, "/") + destination = normalizedEndpointAddress(destination) + for i := 0; i+2 < len(fields); i++ { + if fields[i+1] != "dport" { + continue + } + port, err := strconv.ParseUint(fields[i+2], 10, 16) + if err != nil || strconv.FormatUint(port, 10) != fields[i+2] { + continue + } + matches := []string{fields[i]} + if fields[i] == "th" { + matches = protocols + } + for _, protocol := range matches { + inspection.dnat[ProxyEndpoint{Protocol: protocol, HostIP: destination, HostPort: uint16(port)}] = true + } + } + } + return inspection +} + +func normalizedEndpointAddress(address string) string { + address = strings.TrimSpace(address) + if parsed, err := netip.ParseAddr(address); err == nil { + return parsed.String() + } + return address +} + +func (inspection EndpointInspection) DNATMatches(hostIP string, hostPort uint16, protocol string) bool { + key := ProxyEndpoint{Protocol: protocol, HostPort: hostPort} + if inspection.dnat[key] { + return true + } + if isWildcardHostAddress(hostIP, inspection.family) { + return false + } + key.HostIP = normalizedEndpointAddress(hostIP) + return inspection.dnat[key] +} + +func (inspection EndpointInspection) ProxyMatches(hostIP string, hostPort uint16, protocol string) bool { + key := ProxyEndpoint{Protocol: protocol, HostPort: hostPort, HostIP: normalizedEndpointAddress(hostIP)} + if inspection.proxies[key] { + return true + } + if !isWildcardHostAddress(hostIP, inspection.family) { + return false + } + key.HostIP = "" + return inspection.proxies[key] +} + +func hostAddressMatches(left, right, family string) bool { + if isWildcardHostAddress(left, family) && isWildcardHostAddress(right, family) { + return true + } + left, right = strings.TrimSpace(left), strings.TrimSpace(right) + if address, err := netip.ParseAddr(left); err == nil { + left = address.String() + } + if address, err := netip.ParseAddr(right); err == nil { + right = address.String() + } + return left == right +} + +func isWildcardHostAddress(value, family string) bool { + value = strings.TrimSpace(value) + if family == constant.FirewallFamilyIPv6 { + return value == "" || value == "::" + } + return value == "" || value == "0.0.0.0" +} diff --git a/agent/utils/firewall/docker_guard/manager.go b/agent/utils/firewall/docker_guard/manager.go index 0b1a8f546e06..d17c3a382bd9 100644 --- a/agent/utils/firewall/docker_guard/manager.go +++ b/agent/utils/firewall/docker_guard/manager.go @@ -3,73 +3,19 @@ package docker_guard import ( "errors" "fmt" + "github.com/1Panel-dev/1Panel/agent/buserr" "sort" "strconv" "strings" "sync" "time" - "github.com/1Panel-dev/1Panel/agent/constant" "github.com/1Panel-dev/1Panel/agent/utils/cmd" firewallutil "github.com/1Panel-dev/1Panel/agent/utils/firewall" "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" "github.com/1Panel-dev/1Panel/agent/utils/firewall/nftables_helper" ) -const ( - Chain = "1PANEL_DOCKER" - DockerChain = "DOCKER-USER" - FamilyIPv4 = constant.FirewallFamilyIPv4 - FamilyIPv6 = constant.FirewallFamilyIPv6 - ModeSources = "deny_sources" - ModeAllow = "allow_sources" - ModeAll = "deny_all" - - StatusEffective = "effective" - StatusDisabled = "disabled" - StatusNotEffective = "not_effective" - - ReasonCommandMissing = "command_missing" - ReasonDockerChainMissing = "docker_chain_missing" - ReasonGuardChainMissing = "guard_chain_missing" - ReasonJumpMissing = "jump_missing" - ReasonJumpNotFirst = "jump_not_first" - ReasonJumpDuplicate = "jump_duplicate" - ReasonInspectFailed = "inspect_failed" -) - -var ( - ErrDockerChainUnavailable = errors.New("Docker DOCKER-USER chain is unavailable") - ErrDockerIptablesChainUnavailable = fmt.Errorf("%w for iptables", ErrDockerChainUnavailable) - ErrDockerNftablesChainUnavailable = fmt.Errorf("%w for nftables", ErrDockerChainUnavailable) -) - -type FamilyError struct { - Family string - Err error -} - -func (e *FamilyError) Error() string { return fmt.Sprintf("%s Docker port guard: %v", e.Family, e.Err) } -func (e *FamilyError) Unwrap() error { return e.Err } - -type Policy struct { - UUID string - Family string - HostIP string - HostPort uint16 - Protocol string - Mode string - Sources []string -} - -type FamilyStatus struct { - State string - Reason string - Initialized bool - Bound bool - Effective bool -} - type Runner interface { Run(executable string, args ...string) (string, error) RunInput(executable, input string, args ...string) (string, error) @@ -122,24 +68,20 @@ func dockerGuardExecutable(logical string) string { } } -type Manager struct { +type Iptables struct { runner Runner } var mutationMu sync.Mutex -func NewManager() *Manager { return &Manager{runner: commandRunner{}} } +func NewIptables() *Iptables { return &Iptables{runner: commandRunner{}} } -func (m *Manager) Initialize(policies []Policy) error { +func (m *Iptables) Initialize(policies []Policy, inventory PolicyInventory) error { mutationMu.Lock() defer mutationMu.Unlock() if err := CheckIPv4Forwarding(); err != nil { return err } - inventory, err := m.ListPolicies() - if err != nil { - return err - } if !m.runner.Exists("iptables-restore") { return errors.New("iptables-restore is not installed") } @@ -163,7 +105,7 @@ func (m *Manager) Initialize(policies []Policy) error { return m.rebuildLocked(policies, inventory) } -func (m *Manager) Bind() error { +func (m *Iptables) Bind() error { mutationMu.Lock() defer mutationMu.Unlock() if err := m.bindExistingFamily("iptables", true); err != nil { @@ -177,17 +119,13 @@ func (m *Manager) Bind() error { return nil } -func (m *Manager) Reconcile(policies []Policy) error { +func (m *Iptables) ReplacePolicies(policies []Policy, inventory PolicyInventory) error { mutationMu.Lock() defer mutationMu.Unlock() - inventory, err := m.ListPolicies() - if err != nil { - return err - } return m.rebuildLocked(policies, inventory) } -func (m *Manager) ListPolicies() (PolicyInventory, error) { +func (m *Iptables) ListPolicies() (PolicyInventory, error) { inventory := PolicyInventory{Policies: make([]Policy, 0), ManagedRuleOrders: make(map[string][]int64)} for _, family := range []string{FamilyIPv4, FamilyIPv6} { executable := executableForFamily(family) @@ -218,7 +156,7 @@ func (m *Manager) ListPolicies() (PolicyInventory, error) { return inventory, nil } -func (m *Manager) Unbind() error { +func (m *Iptables) Unbind() error { mutationMu.Lock() defer mutationMu.Unlock() if err := m.unbindFamily("iptables"); err != nil { @@ -232,7 +170,7 @@ func (m *Manager) Unbind() error { return nil } -func (m *Manager) Cleanup() error { +func (m *Iptables) Cleanup() error { mutationMu.Lock() defer mutationMu.Unlock() for _, executable := range []string{"iptables", "ip6tables"} { @@ -254,7 +192,7 @@ func (m *Manager) Cleanup() error { return nil } -func (m *Manager) Initialized(family string) (bool, error) { +func (m *Iptables) Initialized(family string) (bool, error) { executable := executableForFamily(family) if executable == "" || !m.runner.Exists(executable) { return false, nil @@ -262,7 +200,7 @@ func (m *Manager) Initialized(family string) (bool, error) { return m.chainExists(executable, Chain) } -func (m *Manager) Status(family string) FamilyStatus { +func (m *Iptables) Status(family string) FamilyStatus { executable := executableForFamily(family) if executable == "" || !m.runner.Exists(executable) { return FamilyStatus{State: StatusDisabled, Reason: ReasonCommandMissing} @@ -278,12 +216,7 @@ func (m *Manager) Status(family string) FamilyStatus { return FamilyStatus{State: StatusDisabled, Reason: ReasonGuardChainMissing} } status := FamilyStatus{State: StatusNotEffective, Initialized: true} - rules, err := m.run(executable, "-S", DockerChain) - if err != nil { - status.Reason = ReasonInspectFailed - return status - } - jumps := countJumps(rules) + jumps := countJumps(chains) if jumps == 0 { status.Reason = ReasonJumpMissing return status @@ -292,7 +225,7 @@ func (m *Manager) Status(family string) FamilyStatus { status.Reason = ReasonJumpDuplicate return status } - if !hasFirstUniqueJump(rules) { + if !hasFirstUniqueJump(chains) { status.Reason = ReasonJumpNotFirst return status } @@ -302,7 +235,7 @@ func (m *Manager) Status(family string) FamilyStatus { return status } -func (m *Manager) bindExistingFamily(executable string, required bool) error { +func (m *Iptables) bindExistingFamily(executable string, required bool) error { if !m.runner.Exists(executable) { if required { return fmt.Errorf("%s is not installed", executable) @@ -315,7 +248,7 @@ func (m *Manager) bindExistingFamily(executable string, required bool) error { } if !chainDeclared(output, DockerChain) { if required { - return ErrDockerIptablesChainUnavailable + return buserr.New("ErrDockerIptablesChainUnavailable") } return nil } @@ -328,7 +261,7 @@ func (m *Manager) bindExistingFamily(executable string, required bool) error { return m.restoreLifecycle(executable, dockerGuardLifecycleRules(output, true, false)) } -func (m *Manager) ensureFamily(executable string, required bool) error { +func (m *Iptables) ensureFamily(executable string, required bool) error { if !m.runner.Exists(executable) { if required { return fmt.Errorf("%s is not installed", executable) @@ -341,7 +274,7 @@ func (m *Manager) ensureFamily(executable string, required bool) error { } if !chainDeclared(output, DockerChain) { if required { - return ErrDockerIptablesChainUnavailable + return buserr.New("ErrDockerIptablesChainUnavailable") } return nil } @@ -364,7 +297,7 @@ func dockerGuardLifecycleRules(output string, bind, createOwned bool) [][]string return rules } -func (m *Manager) restoreLifecycle(executable string, rules [][]string) error { +func (m *Iptables) restoreLifecycle(executable string, rules [][]string) error { if len(rules) == 0 { return nil } @@ -382,7 +315,7 @@ func (m *Manager) restoreLifecycle(executable string, rules [][]string) error { return nil } -func (m *Manager) rebuildLocked(policies []Policy, inventory PolicyInventory) error { +func (m *Iptables) rebuildLocked(policies []Policy, inventory PolicyInventory) error { for _, family := range []string{FamilyIPv4, FamilyIPv6} { executable := executableForFamily(family) if executable == "" || !m.runner.Exists(executable) { @@ -447,7 +380,7 @@ func orderedIPTablesRules(family string, policies []Policy, inventory PolicyInve continue } compiled := compilePolicy(policy) - orders := inventory.ManagedRuleOrders[managedOrderKey(policy.Family, policy.UUID)] + orders := inventory.ManagedRuleOrders[policy.Family+"\x00"+policy.UUID] for ruleIndex, rule := range compiled { order := int64(0) if ruleIndex < len(orders) { @@ -524,7 +457,7 @@ func compilePolicy(policy Policy) [][]string { return rules } -func (m *Manager) unbindFamily(executable string) error { +func (m *Iptables) unbindFamily(executable string) error { if !m.runner.Exists(executable) { return nil } @@ -538,7 +471,7 @@ func (m *Manager) unbindFamily(executable string) error { return nil } -func (m *Manager) chainExists(executable, chain string) (bool, error) { +func (m *Iptables) chainExists(executable, chain string) (bool, error) { output, err := m.run(executable, "-S") if err != nil { return false, err @@ -556,7 +489,7 @@ func chainDeclared(output, chain string) bool { return false } -func (m *Manager) run(executable string, args ...string) (string, error) { +func (m *Iptables) run(executable string, args ...string) (string, error) { commandArgs := append([]string{"-w", "-t", "filter"}, args...) return m.runner.Run(executable, commandArgs...) } diff --git a/agent/utils/firewall/docker_guard/nftables.go b/agent/utils/firewall/docker_guard/nftables.go index b459a4171eff..d92418b52033 100644 --- a/agent/utils/firewall/docker_guard/nftables.go +++ b/agent/utils/firewall/docker_guard/nftables.go @@ -3,6 +3,7 @@ package docker_guard import ( "errors" "fmt" + "github.com/1Panel-dev/1Panel/agent/buserr" "sort" "strconv" "strings" @@ -18,13 +19,13 @@ const ( dockerNftTable = "docker-bridges" ) -type NftablesManager struct { +type Nftables struct { runner Runner } -func NewNftablesManager() *NftablesManager { return &NftablesManager{runner: commandRunner{}} } +func NewNftables() *Nftables { return &Nftables{runner: commandRunner{}} } -func (m *NftablesManager) Initialize(policies []Policy) error { +func (m *Nftables) Initialize(policies []Policy, inventory PolicyInventory) error { mutationMu.Lock() defer mutationMu.Unlock() if !m.runner.Exists("nft") { @@ -36,10 +37,6 @@ func (m *NftablesManager) Initialize(policies []Policy) error { if err := m.checkForwardPolicy(); err != nil { return err } - inventory, err := m.ListPolicies() - if err != nil { - return err - } if err := m.ensureFamily(FamilyIPv4, true); err != nil { return err } @@ -49,7 +46,7 @@ func (m *NftablesManager) Initialize(policies []Policy) error { return m.rebuildLocked(policies, inventory) } -func (m *NftablesManager) Bind() error { +func (m *Nftables) Bind() error { mutationMu.Lock() defer mutationMu.Unlock() if err := m.bindExistingFamily(FamilyIPv4, true); err != nil { @@ -61,17 +58,13 @@ func (m *NftablesManager) Bind() error { return nil } -func (m *NftablesManager) Reconcile(policies []Policy) error { +func (m *Nftables) ReplacePolicies(policies []Policy, inventory PolicyInventory) error { mutationMu.Lock() defer mutationMu.Unlock() - inventory, err := m.ListPolicies() - if err != nil { - return err - } return m.rebuildLocked(policies, inventory) } -func (m *NftablesManager) ListPolicies() (PolicyInventory, error) { +func (m *Nftables) ListPolicies() (PolicyInventory, error) { if !m.runner.Exists("nft") { return PolicyInventory{}, nil } @@ -98,7 +91,7 @@ func (m *NftablesManager) ListPolicies() (PolicyInventory, error) { return inventory, nil } -func (m *NftablesManager) Unbind() error { +func (m *Nftables) Unbind() error { mutationMu.Lock() defer mutationMu.Unlock() for _, family := range []string{FamilyIPv4, FamilyIPv6} { @@ -109,7 +102,7 @@ func (m *NftablesManager) Unbind() error { return nil } -func (m *NftablesManager) Cleanup() error { +func (m *Nftables) Cleanup() error { mutationMu.Lock() defer mutationMu.Unlock() if !m.runner.Exists("nft") { @@ -126,7 +119,7 @@ func (m *NftablesManager) Cleanup() error { return m.runBatch(commands) } -func (m *NftablesManager) Initialized(family string) (bool, error) { +func (m *Nftables) Initialized(family string) (bool, error) { if nftTableFamily(family) == "" || !m.runner.Exists("nft") { return false, nil } @@ -137,7 +130,7 @@ func (m *NftablesManager) Initialized(family string) (bool, error) { return m.objectExists("chain", tableFamily, NftTable, NftChain), nil } -func (m *NftablesManager) Status(family string) FamilyStatus { +func (m *Nftables) Status(family string) FamilyStatus { tableFamily := nftTableFamily(family) if tableFamily == "" || !m.runner.Exists("nft") { return FamilyStatus{State: StatusDisabled, Reason: ReasonCommandMissing} @@ -145,16 +138,32 @@ func (m *NftablesManager) Status(family string) FamilyStatus { if !m.objectExists("table", tableFamily, dockerNftTable) { return FamilyStatus{State: StatusDisabled, Reason: ReasonDockerChainMissing} } - if !m.objectExists("chain", tableFamily, NftTable, NftBaseChain) || - !m.objectExists("chain", tableFamily, NftTable, NftChain) { + output, err := m.run("-a", "list", "table", tableFamily, NftTable) + if err != nil { return FamilyStatus{State: StatusDisabled, Reason: ReasonGuardChainMissing} } - status := FamilyStatus{State: StatusNotEffective, Initialized: true} - rules, err := m.run("-a", "list", "chain", tableFamily, NftTable, NftBaseChain) - if err != nil { - status.Reason = ReasonInspectFailed - return status + baseExists, guardExists := false, false + currentChain := "" + var baseRules strings.Builder + for _, line := range strings.Split(output, "\n") { + fields := strings.Fields(line) + if len(fields) >= 3 && fields[0] == "chain" && fields[2] == "{" { + currentChain = fields[1] + baseExists = baseExists || currentChain == NftBaseChain + guardExists = guardExists || currentChain == NftChain + } else if strings.TrimSpace(line) == "}" { + currentChain = "" + } + if currentChain == NftBaseChain { + baseRules.WriteString(line) + baseRules.WriteByte('\n') + } + } + if !baseExists || !guardExists { + return FamilyStatus{State: StatusDisabled, Reason: ReasonGuardChainMissing} } + status := FamilyStatus{State: StatusNotEffective, Initialized: true} + rules := baseRules.String() jumps := nftJumpHandles(rules) if len(jumps) == 0 { status.Reason = ReasonJumpMissing @@ -174,14 +183,14 @@ func (m *NftablesManager) Status(family string) FamilyStatus { return status } -func (m *NftablesManager) ensureFamily(family string, required bool) error { +func (m *Nftables) ensureFamily(family string, required bool) error { tableFamily := nftTableFamily(family) if tableFamily == "" { return fmt.Errorf("unsupported address family %q", family) } if !m.objectExists("table", tableFamily, dockerNftTable) { if required { - return fmt.Errorf("%w %s", ErrDockerNftablesChainUnavailable, family) + return fmt.Errorf("%w %s", buserr.New("ErrDockerNftablesChainUnavailable"), family) } return nil } @@ -213,7 +222,7 @@ func (m *NftablesManager) ensureFamily(family string, required bool) error { return m.runBatch(commands) } -func (m *NftablesManager) bindExistingFamily(family string, required bool) error { +func (m *Nftables) bindExistingFamily(family string, required bool) error { tableFamily := nftTableFamily(family) if !m.runner.Exists("nft") { if required { @@ -223,7 +232,7 @@ func (m *NftablesManager) bindExistingFamily(family string, required bool) error } if !m.objectExists("table", tableFamily, dockerNftTable) { if required { - return fmt.Errorf("%w %s", ErrDockerNftablesChainUnavailable, family) + return fmt.Errorf("%w %s", buserr.New("ErrDockerNftablesChainUnavailable"), family) } return nil } @@ -237,7 +246,7 @@ func (m *NftablesManager) bindExistingFamily(family string, required bool) error return m.ensureJump(family) } -func (m *NftablesManager) ensureJump(family string) error { +func (m *Nftables) ensureJump(family string) error { tableFamily := nftTableFamily(family) output, err := m.run("-a", "list", "chain", tableFamily, NftTable, NftBaseChain) if err != nil { @@ -251,7 +260,7 @@ func (m *NftablesManager) ensureJump(family string) error { return m.runBatch(commands) } -func (m *NftablesManager) rebuildLocked(policies []Policy, inventory PolicyInventory) error { +func (m *Nftables) rebuildLocked(policies []Policy, inventory PolicyInventory) error { if !m.runner.Exists("nft") { return nil } @@ -312,7 +321,7 @@ func orderedNftRules(family string, policies []Policy, inventory PolicyInventory continue } compiled := compileNftPolicy(policy) - orders := inventory.ManagedRuleOrders[managedOrderKey(policy.Family, policy.UUID)] + orders := inventory.ManagedRuleOrders[policy.Family+"\x00"+policy.UUID] for ruleIndex, rule := range compiled { order := int64(0) if ruleIndex < len(orders) { @@ -408,7 +417,7 @@ func validNftToken(token string) bool { return !strings.ContainsAny(token, " \t\\\"'") } -func (m *NftablesManager) unbindFamily(family string) error { +func (m *Nftables) unbindFamily(family string) error { if !m.runner.Exists("nft") { return nil } @@ -427,7 +436,7 @@ func (m *NftablesManager) unbindFamily(family string) error { return m.runBatch(commands) } -func (m *NftablesManager) runBatch(commands [][]string) error { +func (m *Nftables) runBatch(commands [][]string) error { if len(commands) == 0 { return nil } @@ -441,13 +450,13 @@ func (m *NftablesManager) runBatch(commands [][]string) error { return nil } -func (m *NftablesManager) objectExists(kind string, args ...string) bool { +func (m *Nftables) objectExists(kind string, args ...string) bool { command := append([]string{"list", kind}, args...) _, err := m.run(command...) return err == nil } -func (m *NftablesManager) run(args ...string) (string, error) { +func (m *Nftables) run(args ...string) (string, error) { return m.runner.Run("nft", args...) } @@ -497,9 +506,7 @@ func nftHasFirstUniqueJump(output string) bool { return false } -var ErrDockerForwardPolicyDrop = errors.New("iptables FORWARD default policy is DROP") - -func (m *NftablesManager) checkForwardPolicy() error { +func (m *Nftables) checkForwardPolicy() error { for _, family := range []struct{ command, name string }{ {"iptables", FamilyIPv4}, {"ip6tables", FamilyIPv6}, @@ -519,7 +526,11 @@ func (m *NftablesManager) checkForwardPolicy() error { } found = true if fields[2] == "DROP" { - return &FamilyError{Family: family.name, Err: ErrDockerForwardPolicyDrop} + label := "IPv4" + if family.name == FamilyIPv6 { + label = "IPv6" + } + return &FamilyError{Family: family.name, Err: buserr.WithMap("ErrDockerForwardPolicyDrop", map[string]interface{}{"family": label}, nil)} } if fields[2] != "ACCEPT" { return &FamilyError{Family: family.name, Err: fmt.Errorf("unexpected iptables FORWARD policy: %s", fields[2])} diff --git a/agent/utils/firewall/docker_guard/policy.go b/agent/utils/firewall/docker_guard/policy.go index 03aa77884547..287d24667e8c 100644 --- a/agent/utils/firewall/docker_guard/policy.go +++ b/agent/utils/firewall/docker_guard/policy.go @@ -1,8 +1,6 @@ package docker_guard import ( - "encoding/json" - "errors" "fmt" "net/netip" "sort" @@ -12,141 +10,6 @@ import ( "github.com/mattn/go-shellwords" ) -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 != ModeAll && len(normalizedSources) == 0 { - return Policy{}, fmt.Errorf("%w: source-based modes require 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 := make([]string, 0, len(policy.Sources)) - for _, source := range policy.Sources { - sources = append(sources, canonicalPolicySource(source)) - } - 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 canonicalPolicySource(value string) string { - value = strings.TrimSpace(value) - if prefix, err := netip.ParsePrefix(value); err == nil { - return prefix.Masked().String() - } - if address, err := netip.ParseAddr(value); err == nil { - address = address.Unmap() - return netip.PrefixFrom(address, address.BitLen()).String() - } - return value -} - -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 -} - type observedPolicy struct { policy Policy sequence int64 @@ -235,7 +98,7 @@ func parseDockerGuardPolicies(output, family string) (PolicyInventory, error) { return PolicyInventory{}, fmt.Errorf("Docker guard policy %s has no effective rules", group.policy.UUID) } inventory.Policies = append(inventory.Policies, group.policy) - inventory.ManagedRuleOrders[managedOrderKey(group.policy.Family, group.policy.UUID)] = append([]int64(nil), group.managedOrders...) + inventory.ManagedRuleOrders[group.policy.Family+"\x00"+group.policy.UUID] = append([]int64(nil), group.managedOrders...) } return inventory, nil } @@ -260,12 +123,11 @@ func nativeRuleTokens(tokens []string) []string { return result } -func managedOrderKey(family, policyUUID string) string { - return family + "\x00" + policyUUID -} - func parseDockerGuardRuleTokens(tokens []string, family string) (Policy, string, string, error) { - policy := Policy{Family: family, HostIP: wildcardHost(family)} + policy := Policy{Family: family, HostIP: "0.0.0.0"} + if family == FamilyIPv6 { + policy.HostIP = "::" + } source, action := "", "" for index := 0; index < len(tokens); index++ { switch tokens[index] { @@ -317,7 +179,7 @@ func parseDockerGuardRuleTokens(tokens []string, family string) (Policy, string, policy.HostPort = parsePolicyPort(nextPolicyToken(tokens, index+1)) } case "accept", "drop", "return": - if isCommentValue(tokens, index) { + if index > 0 && (tokens[index-1] == "comment" || tokens[index-1] == "--comment") { continue } action = tokens[index] @@ -337,20 +199,13 @@ func hasAcceptAction(tokens []string) bool { if token == "-j" && strings.EqualFold(nextPolicyToken(tokens, index), "accept") { return true } - if strings.EqualFold(token, "accept") && !isCommentValue(tokens, index) { + if strings.EqualFold(token, "accept") && !(index > 0 && (tokens[index-1] == "comment" || tokens[index-1] == "--comment")) { return true } } return false } -func isCommentValue(tokens []string, index int) bool { - if index == 0 { - return false - } - return tokens[index-1] == "comment" || tokens[index-1] == "--comment" -} - func normalizeObservedHost(value string) string { if prefix, err := netip.ParsePrefix(value); err == nil && prefix.Bits() == prefix.Addr().BitLen() { return prefix.Addr().String() @@ -373,13 +228,6 @@ func parsePolicyPort(value string) uint16 { return uint16(port) } -func wildcardHost(family string) string { - if family == FamilyIPv6 { - return "::" - } - return "0.0.0.0" -} - func uniqueSortedStrings(values []string) []string { seen := make(map[string]struct{}, len(values)) result := make([]string, 0, len(values)) diff --git a/agent/utils/firewall/docker_guard/runtime.go b/agent/utils/firewall/docker_guard/runtime.go index 6531b248663c..54545bac662b 100644 --- a/agent/utils/firewall/docker_guard/runtime.go +++ b/agent/utils/firewall/docker_guard/runtime.go @@ -1,15 +1,77 @@ package docker_guard import ( - "errors" + "github.com/1Panel-dev/1Panel/agent/buserr" + "fmt" "os" - "slices" "strings" "github.com/1Panel-dev/1Panel/agent/constant" ) +const ( + Chain = "1PANEL_DOCKER" + DockerChain = "DOCKER-USER" + FamilyIPv4 = constant.FirewallFamilyIPv4 + FamilyIPv6 = constant.FirewallFamilyIPv6 + ModeSources = "deny_sources" + ModeAllow = "allow_sources" + ModeAll = "deny_all" + StatusEffective = "effective" + StatusDisabled = "disabled" + StatusNotEffective = "not_effective" + ReasonCommandMissing = "command_missing" + ReasonDockerChainMissing = "docker_chain_missing" + ReasonGuardChainMissing = "guard_chain_missing" + ReasonJumpMissing = "jump_missing" + ReasonJumpNotFirst = "jump_not_first" + ReasonJumpDuplicate = "jump_duplicate" + ReasonInspectFailed = "inspect_failed" +) + +type FamilyError struct { + Family string + Err error +} + +func (e *FamilyError) Error() string { return fmt.Sprintf("%s Docker port guard: %v", e.Family, e.Err) } +func (e *FamilyError) Unwrap() error { return e.Err } + +type ProxyEndpoint struct { + Protocol string + HostIP string + HostPort uint16 +} + +type ProxyEndpoints struct { + Items []ProxyEndpoint + Inspected bool +} + +type DNATRules struct { + Output string + Inspected bool +} + +type Policy struct { + UUID string + Family string + HostIP string + HostPort uint16 + Protocol string + Mode string + Sources []string +} + +type FamilyStatus struct { + State string + Reason string + Initialized bool + Bound bool + Effective bool +} + type NativeRule struct { Family string `json:"family"` Order int64 `json:"order"` @@ -30,9 +92,9 @@ type PolicyInventory struct { } type Runtime interface { - Initialize([]Policy) error + Initialize([]Policy, PolicyInventory) error Bind() error - Reconcile([]Policy) error + ReplacePolicies([]Policy, PolicyInventory) error Unbind() error Cleanup() error Initialized(string) (bool, error) @@ -40,118 +102,8 @@ type Runtime interface { ListPolicies() (PolicyInventory, error) } -func NewRuntime(provider string) Runtime { - if provider == constant.FirewallProviderNftables { - return NewNftablesManager() - } - return NewManager() -} - -func Verify(runtime Runtime, desired []Policy, preserved []ReadOnlyPolicy) error { - inventory, err := runtime.ListPolicies() - if err != nil { - return fmt.Errorf("verify synchronized Docker firewall policies: %w", err) - } - if !PolicyStatesEqual(inventory.Policies, desired) { - return fmt.Errorf("verify synchronized Docker firewall policies: target policies do not match the database") - } - if !readOnlyStatesEqual(inventory.ReadOnly, preserved) { - return fmt.Errorf("verify synchronized Docker firewall policies: read-only runtime rules changed") - } - return nil -} - -func readOnlyStatesEqual(left, right []ReadOnlyPolicy) bool { - if len(left) != len(right) { - return false - } - leftRules := flattenNativeRules(left) - rightRules := flattenNativeRules(right) - if len(leftRules) != len(rightRules) { - return false - } - for index := range leftRules { - if leftRules[index].Family != rightRules[index].Family || !slices.Equal(leftRules[index].Tokens, rightRules[index].Tokens) { - return false - } - } - return true -} - -func flattenNativeRules(policies []ReadOnlyPolicy) []NativeRule { - rules := make([]NativeRule, 0) - for _, policy := range policies { - rules = append(rules, policy.NativeRules...) - } - slices.SortStableFunc(rules, func(left, right NativeRule) int { - if left.Family < right.Family { - return -1 - } - if left.Family > right.Family { - return 1 - } - if left.Order < right.Order { - return -1 - } - if left.Order > right.Order { - return 1 - } - return 0 - }) - return rules -} - -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 -} - const ipv4ForwardingPath = "/proc/sys/net/ipv4/ip_forward" -var ErrIPv4ForwardingDisabled = errors.New("IPv4 forwarding is disabled; set net.ipv4.ip_forward=1 before using Docker's firewall backend") - func CheckIPv4Forwarding() error { return checkIPv4Forwarding(os.ReadFile) } @@ -162,7 +114,7 @@ func checkIPv4Forwarding(readFile func(string) ([]byte, error)) error { return fmt.Errorf("inspect IPv4 forwarding: %w", err) } if strings.TrimSpace(string(value)) != "1" { - return ErrIPv4ForwardingDisabled + return buserr.New("ErrDockerIPv4ForwardingDisabled") } return nil } diff --git a/agent/utils/firewall/filter/adapter.go b/agent/utils/firewall/filter/adapter.go index edb0807d7676..b40a927314be 100644 --- a/agent/utils/firewall/filter/adapter.go +++ b/agent/utils/firewall/filter/adapter.go @@ -21,7 +21,7 @@ const ( ChangeReorder ChangeOperation = "reorder" ) -type DesiredChange struct { +type RuleChange struct { CommandOnly bool `json:"-"` UnmarkedAdopted bool `json:"-"` Operation ChangeOperation `json:"operation"` @@ -39,7 +39,7 @@ type NativeCommand struct { Stdin string `json:"stdin,omitempty"` } -type NativeRulePlan struct { +type RuleCommands struct { RuleUUID string `json:"ruleUUID"` Operation ChangeOperation `json:"operation"` Commands []NativeCommand `json:"commands"` @@ -48,15 +48,14 @@ type NativeRulePlan struct { Expected ObservedRule `json:"expected"` } -type BackendPlan struct { - CommandOnly bool `json:"-"` - Provider Provider `json:"provider"` - Scope Scope `json:"scope"` - SnapshotRevision string `json:"snapshotRevision"` - Rules []NativeRulePlan `json:"rules"` +type CommandBatch struct { + CommandOnly bool `json:"-"` + Provider Provider `json:"provider"` + Scope Scope `json:"scope"` + Rules []RuleCommands `json:"rules"` } -func (p BackendPlan) CreatesOnly() bool { +func (p CommandBatch) CreatesOnly() bool { if len(p.Rules) == 0 { return false } @@ -68,40 +67,17 @@ func (p BackendPlan) CreatesOnly() bool { return true } -type ApplyResult struct { - Applied []ObservedRule `json:"applied"` - Verification *VerifyResult `json:"verification,omitempty"` -} - -type VerifyResult struct { - Snapshot Snapshot `json:"snapshot"` - Matched bool `json:"matched"` -} - type Adapter interface { Provider() Provider Capabilities(context.Context) (Capabilities, error) - Observe(context.Context, Scope) (Snapshot, error) - Compile(Snapshot, []DesiredChange) (BackendPlan, error) - Apply(context.Context, BackendPlan) (ApplyResult, error) - Verify(context.Context, BackendPlan) (VerifyResult, error) -} - -type MultiScopeObserver interface { - ObserveScopes(context.Context, []Scope) ([]Snapshot, error) -} - -type ObservationSessionFactory interface { - NewObservationSession() Adapter -} - -type CreatePlanner interface { - Compile(DesiredChange) (BackendPlan, error) - Applied(ObservedRule) + ListRules(context.Context, Scope) (RuleSet, error) + BuildCommands(RuleSet, []RuleChange) (CommandBatch, error) + RunCommands(context.Context, CommandBatch) error + Rollback(context.Context, CommandBatch) error } -type CreatePlannerFactory interface { - NewCreatePlanner(Snapshot) CreatePlanner +type MultiScopeReader interface { + ListRuleScopes(context.Context, []Scope) ([]RuleSet, error) } type RulePreparer interface { @@ -112,7 +88,8 @@ type RuleChecker interface { CheckRule(context.Context, FirewallRule) error } -type UnverifiedRuleAppender interface { +type ExternalRuleAdapter interface { + ListRulesByComment(context.Context, []Scope, string) ([]ObservedRule, error) AppendUnverified(context.Context, FirewallRule, string) error } @@ -120,6 +97,6 @@ type NativeDetailReader interface { NativeDetail(context.Context, string, bool) (string, error) } -type PlanRollbacker interface { - Rollback(context.Context, BackendPlan) error +type RuleSaver interface { + SaveRules(context.Context, Scope) error } diff --git a/agent/utils/firewall/filter/external.go b/agent/utils/firewall/filter/external.go new file mode 100644 index 000000000000..9a0185238efa --- /dev/null +++ b/agent/utils/firewall/filter/external.go @@ -0,0 +1,20 @@ +package filter + +import ( + "context" + "time" + + "github.com/1Panel-dev/1Panel/agent/utils/cmd" +) + +type CommentRuleReader interface { + ReadRulesByComment(context.Context, Scope, string) (string, error) +} + +func ReadRulesByComment(ctx context.Context, executable string, args []string, comment string) (string, error) { + name, args := cmd.WrapWithOptionalSudo(executable, args...) + return cmd.NewCommandMgr(cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second), cmd.WithEnv("LC_ALL=C", "LANGUAGE=en_US:en")).RunPipe( + cmd.PipeCommand{Name: name, Args: args}, + cmd.PipeCommand{Name: "sh", Args: []string{"-c", `grep -F -- "$1"; result=$?; if [ "$result" -eq 1 ]; then exit 0; fi; exit "$result"`, "sh", comment}}, + ) +} diff --git a/agent/utils/firewall/filter/identity.go b/agent/utils/firewall/filter/identity.go index 690033a11654..a854e5d2dd9b 100644 --- a/agent/utils/firewall/filter/identity.go +++ b/agent/utils/firewall/filter/identity.go @@ -5,7 +5,6 @@ import ( "encoding/hex" "encoding/json" "fmt" - "sort" "strings" ) @@ -76,93 +75,6 @@ func SameRuleContent(before, after FirewallRule) (bool, error) { return err == nil && previous == requested && before.Action == after.Action, err } -type RuleCollisionIndex map[string][]Action - -func (index RuleCollisionIndex) Add(rule FirewallRule) error { - key, err := RuleMatchKey(rule) - if err != nil { - return err - } - index[key] = append(index[key], rule.Action) - return nil -} - -func (index RuleCollisionIndex) CheckDuplicate(rule FirewallRule) error { - key, err := RuleMatchKey(rule) - if err != nil { - return err - } - for _, action := range index[key] { - if action == rule.Action { - return checkCollisionActions(rule.Action, action) - } - } - return nil -} - -func (index RuleCollisionIndex) Check(rule FirewallRule) error { - key, err := RuleMatchKey(rule) - if err != nil { - return err - } - for _, action := range index[key] { - if err := checkCollisionActions(rule.Action, action); err != nil { - return err - } - } - return nil -} - -func CheckRuleCollision(requested, existing FirewallRule) error { - wanted, err := RuleMatchKey(requested) - if err != nil { - return err - } - actual, err := RuleMatchKey(existing) - if err != nil { - return err - } - if wanted != actual { - return nil - } - return checkCollisionActions(requested.Action, existing.Action) -} - -func checkCollisionActions(requested, existing Action) error { - if requested == existing { - return fmt.Errorf("%w: equivalent rule already exists", ErrRuleOperation) - } - if OppositeActions(requested, existing) { - return ErrRuleConflict - } - return nil -} - -func CheckObservedRuleCollisions(snapshot Snapshot, requested FirewallRule, excluded *Locator) error { - for _, observed := range snapshot.Rules { - if observed.ParseStatus != ParseStatusSupported || excluded != nil && SameLocator(observed.Locator, *excluded) { - continue - } - if err := CheckRuleCollision(requested, observed.Rule); err != nil { - return err - } - } - return nil -} - -func ObservedRuleCollisionIndex(snapshot Snapshot) (RuleCollisionIndex, error) { - index := make(RuleCollisionIndex, len(snapshot.Rules)) - for _, observed := range snapshot.Rules { - if observed.ParseStatus != ParseStatusSupported { - continue - } - if err := index.Add(observed.Rule); err != nil { - return nil, err - } - } - return index, nil -} - func normalizedRuleKey(normalized FirewallRule) (string, error) { identity := ruleIdentity{ Scope: normalized.Scope.Key(), @@ -222,49 +134,25 @@ func opaqueInstanceKey(rule ObservedRule, locator Locator) (string, error) { }{Raw: strings.TrimSpace(rule.Raw), Locator: locator, Persistence: rule.Persistence}) } -func SnapshotRevision(scope Scope, rules []ObservedRule) (string, error) { +func NewRuleSet(scope Scope, rules []ObservedRule) (RuleSet, error) { scope = scope.Normalize() if err := scope.ValidateMVP(); err != nil { - return "", err + return RuleSet{}, err } - - identities := make([]string, 0, len(rules)) for _, rule := range rules { if rule.Rule.Scope.Normalize().Key() != scope.Key() { - return "", fmt.Errorf("%w: observed rule scope %q does not match snapshot scope %q", ErrInvalidRule, rule.Rule.Scope.Key(), scope.Key()) + return RuleSet{}, fmt.Errorf("%w: observed rule scope does not match read scope", ErrInvalidRule) } - if rule.ParseStatus != ParseStatusSupported { - locator, err := validatedLocator(rule.Locator, scope) - if err != nil { - return "", err - } - identity, err := opaqueInstanceKey(rule, locator) - if err != nil { - return "", err + if rule.ParseStatus == ParseStatusSupported { + if _, err := NormalizeRule(rule.Rule); err != nil { + return RuleSet{}, err } - identities = append(identities, identity) - continue } - - identity, err := InstanceKey(rule) - if err != nil { - return "", err + if _, err := validatedLocator(rule.Locator, scope); err != nil { + return RuleSet{}, err } - identities = append(identities, identity) - } - sort.Strings(identities) - return hashJSON(struct { - Scope string `json:"scope"` - Rules []string `json:"rules"` - }{Scope: scope.Key(), Rules: identities}) -} - -func NewSnapshot(scope Scope, rules []ObservedRule) (Snapshot, error) { - revision, err := SnapshotRevision(scope, rules) - if err != nil { - return Snapshot{}, err } - return Snapshot{Scope: scope.Normalize(), Revision: revision, Rules: rules}, nil + return RuleSet{Scope: scope, Rules: rules}, nil } func normalizeLocator(locator Locator, scope Scope) Locator { diff --git a/agent/utils/firewall/filter/inventory.go b/agent/utils/firewall/filter/inventory.go index 3ad473ce0d1a..0df1529c5a57 100644 --- a/agent/utils/firewall/filter/inventory.go +++ b/agent/utils/firewall/filter/inventory.go @@ -1,9 +1,6 @@ package filter import ( - "fmt" - "strings" - "github.com/1Panel-dev/1Panel/agent/constant" ) @@ -71,273 +68,3 @@ type InventoryMergeInput struct { Desired []DesiredRule ProtectedObservedKeys map[string]struct{} } - -type observedInventoryCandidate struct { - rule ObservedRule - ruleKey string - instanceKey string - claimed bool -} - -func MergeInventory(input InventoryMergeInput) ([]InventoryItem, error) { - candidates := make([]observedInventoryCandidate, len(input.Observed)) - byRuleKey := make(map[string][]int) - byInstanceKey := make(map[string][]int) - byMarker := make(map[string][]int) - for index, observed := range input.Observed { - candidate := observedInventoryCandidate{rule: observed} - if marker := strings.TrimSpace(candidate.rule.Marker); marker != "" { - markerKey := candidate.rule.Rule.Scope.Key() + "\x00" + marker - byMarker[markerKey] = append(byMarker[markerKey], index) - } - if observed.ParseStatus == ParseStatusSupported { - normalized, err := NormalizeRule(observed.Rule) - if err != nil { - return nil, fmt.Errorf("normalize observed firewall rule %d: %w", index, err) - } - candidate.rule.Rule = normalized - candidate.ruleKey, err = normalizedRuleKey(normalized) - if err != nil { - return nil, err - } - byRuleKey[candidate.ruleKey] = append(byRuleKey[candidate.ruleKey], index) - if instanceKey, err := instanceKeyWithRuleKey(candidate.rule, candidate.ruleKey); err == nil { - candidate.instanceKey = instanceKey - candidate.rule.InstanceKey = instanceKey - byInstanceKey[instanceKey] = append(byInstanceKey[instanceKey], index) - } - } - candidates[index] = candidate - } - - normalizedDesired := make([]DesiredRule, 0, len(input.Desired)) - desiredMatches := make(map[int]int) - desiredMatchStates := make([]InventoryMatch, 0, len(input.Desired)) - for _, desired := range input.Desired { - normalized, err := NormalizeRule(desired.Rule) - if err != nil { - return nil, fmt.Errorf("normalize desired firewall rule %q: %w", desired.UUID, err) - } - desired.Rule = normalized - calculatedKey, err := normalizedRuleKey(normalized) - if err != nil { - return nil, err - } - if desired.RuleKey != "" && desired.RuleKey != calculatedKey { - return nil, fmt.Errorf("%w: desired rule %q key does not match its semantics", ErrInvalidRule, desired.UUID) - } - desired.RuleKey = calculatedKey - - match, matchState := findObservedInventoryMatch(desired, candidates, byRuleKey, byInstanceKey, byMarker) - normalizedIndex := len(normalizedDesired) - normalizedDesired = append(normalizedDesired, desired) - desiredMatchStates = append(desiredMatchStates, matchState) - if match >= 0 { - candidates[match].claimed = true - desiredMatches[match] = normalizedIndex - } - } - - items := make([]InventoryItem, 0, len(candidates)+len(normalizedDesired)) - matchedDesired := make(map[int]struct{}, len(desiredMatches)) - for index := range candidates { - candidate := &candidates[index] - if desiredIndex, exists := desiredMatches[index]; exists { - desired := normalizedDesired[desiredIndex] - observed := candidate.rule - match := desiredMatchStates[desiredIndex] - if match == InventoryMatchExact && observed.ParseStatus != ParseStatusSupported { - orderIndex := observed.Rule.OrderIndex - observed.Rule = desired.Rule - observed.Rule.OrderIndex = orderIndex - observed.ParseStatus = ParseStatusSupported - observed.UncertainFields = nil - } - displayRule := observed.Rule - displayRule.Description = desired.Rule.Description - state := inventoryStateForDesired(desired, match) - if observed.Protected { - state = InventoryStateProtected - } else if observed.Persistence != "" && observed.Persistence != PersistenceStatusConverged { - state = InventoryStateDrifted - } - items = append(items, InventoryItem{ - Rule: displayRule, - Observed: &observed, - Desired: &desired, - State: state, - Match: match, - }) - matchedDesired[desiredIndex] = struct{}{} - continue - } - observed := candidate.rule - state := InventoryStateExternal - if observed.Protected { - state = InventoryStateProtected - } else if _, protected := input.ProtectedObservedKeys[candidate.ruleKey]; protected { - state = InventoryStateProtected - } - match := InventoryMatchNone - if observed.ParseStatus != ParseStatusSupported { - match = InventoryMatchOpaque - } - items = append(items, InventoryItem{Rule: observed.Rule, Observed: &observed, State: state, Match: match}) - } - for index, desired := range normalizedDesired { - if _, matched := matchedDesired[index]; matched { - continue - } - desiredCopy := desired - match := desiredMatchStates[index] - items = append(items, InventoryItem{ - Rule: desired.Rule, - Desired: &desiredCopy, - State: inventoryStateForDesired(desired, match), - Match: match, - }) - } - return items, nil -} - -func findObservedInventoryMatch( - desired DesiredRule, - candidates []observedInventoryCandidate, - byRuleKey map[string][]int, - byInstanceKey map[string][]int, - byMarker map[string][]int, -) (int, InventoryMatch) { - if marker := strings.TrimSpace(desired.Marker); marker != "" { - markerKey := desired.Rule.Scope.Key() + "\x00" + marker - match, status := uniqueUnclaimedCandidate(byMarker[markerKey], candidates) - if match >= 0 && candidates[match].rule.ParseStatus != ParseStatusOpaque && - !ObservedRuleMatchesExpected(candidates[match].rule, desired.Rule) { - return match, InventoryMatchChanged - } - if status != InventoryMatchMissing { - return match, status - } - if desired.Origin == RuleOriginAdopted { - match, status = uniqueUnclaimedSemanticCandidate(desired.Rule, "", candidates) - if status != InventoryMatchMissing { - if match >= 0 { - return match, InventoryMatchChanged - } - return match, status - } - } - legacyMarker := "1panel-rule:" + strings.TrimSpace(desired.UUID) - if legacyMarker != "1panel-rule:" && legacyMarker != marker { - match, status = uniqueUnclaimedSemanticCandidate(desired.Rule, legacyMarker, candidates) - if status != InventoryMatchMissing { - if match >= 0 { - return match, InventoryMatchChanged - } - return match, status - } - } - return match, status - } - if desired.ObservedInstanceKey != "" { - return uniqueUnclaimedCandidate(byInstanceKey[desired.ObservedInstanceKey], candidates) - } - return uniqueUnclaimedCandidate(byRuleKey[desired.RuleKey], candidates) -} - -func uniqueUnclaimedSemanticCandidate( - expected FirewallRule, - marker string, - candidates []observedInventoryCandidate, -) (int, InventoryMatch) { - match := -1 - count := 0 - for index := range candidates { - candidate := candidates[index] - if candidate.claimed || strings.TrimSpace(candidate.rule.Marker) != marker || - !ObservedRuleMatchesExpected(candidate.rule, expected) { - continue - } - match = index - count++ - } - switch count { - case 0: - return -1, InventoryMatchMissing - case 1: - return match, InventoryMatchExact - default: - return -1, InventoryMatchAmbiguous - } -} - -func uniqueUnclaimedCandidate(indices []int, candidates []observedInventoryCandidate) (int, InventoryMatch) { - match := -1 - count := 0 - for _, index := range indices { - if candidates[index].claimed { - continue - } - match = index - count++ - } - switch count { - case 0: - return -1, InventoryMatchMissing - case 1: - return match, InventoryMatchExact - default: - return -1, InventoryMatchAmbiguous - } -} - -func inventoryStateForDesired(desired DesiredRule, match InventoryMatch) InventoryState { - if match != InventoryMatchExact { - return InventoryStateDrifted - } - if desired.Protected { - return InventoryStateProtected - } - switch desired.Origin { - case RuleOriginAdopted: - return InventoryStateAdopted - default: - return InventoryStateManaged - } -} - -func InventoryPositionRanges(provider Provider, items []InventoryItem) (ipv4, ipv6 PositionRange) { - if provider == ProviderFirewalld { - return PositionRange{Min: -32768, Max: 32767}, PositionRange{Min: -32768, Max: 32767} - } - for _, item := range items { - if item.Observed == nil || item.Observed.Locator.Position == nil { - continue - } - scope := item.Observed.Rule.Scope - if scope.Provider != provider || scope.Direction != DirectionInput { - continue - } - if (provider == ProviderIptables || provider == ProviderNftables) && - (scope.Table != "filter" || scope.Chain != IptablesInputChain) { - continue - } - bounds := &ipv4 - if scope.Family == FamilyIPv6 { - bounds = &ipv6 - } else if scope.Family != FamilyIPv4 { - continue - } - position := *item.Observed.Locator.Position - if position < 1 { - continue - } - if bounds.Min == 0 || position < bounds.Min { - bounds.Min = position - } - bounds.Max = max(bounds.Max, position) - if provider != ProviderUFW { - bounds.Min = 1 - } - } - return -} diff --git a/agent/utils/firewall/filter/model.go b/agent/utils/firewall/filter/model.go index b80375c4e63c..1a8033c6477f 100644 --- a/agent/utils/firewall/filter/model.go +++ b/agent/utils/firewall/filter/model.go @@ -110,13 +110,12 @@ type ScopeNotice struct { } var ( - ErrInvalidScope = errors.New("invalid firewall scope") - ErrUnsupportedScope = errors.New("unsupported firewall scope") - ErrManagedScopeChange = fmt.Errorf("%w: managed rule scope cannot be changed", ErrUnsupportedScope) - ErrInvalidRule = errors.New("invalid firewall rule") - ErrProtectedRule = errors.New("protected firewall rule cannot be modified") - ErrCompositeRule = errors.New("firewall rule must be atomic") - ErrExpansionLimit = errors.New("firewall rule expansion limit exceeded") + ErrInvalidScope = errors.New("invalid firewall scope") + ErrUnsupportedScope = errors.New("unsupported firewall scope") + ErrInvalidRule = errors.New("invalid firewall rule") + ErrProtectedRule = errors.New("protected firewall rule cannot be modified") + ErrCompositeRule = errors.New("firewall rule must be atomic") + ErrExpansionLimit = errors.New("firewall rule expansion limit exceeded") ) type Scope struct { @@ -298,11 +297,11 @@ type ObservedRule struct { Persistence PersistenceStatus `json:"persistence,omitempty"` } -type Snapshot struct { - Scope Scope `json:"scope"` - Revision string `json:"revision"` - Rules []ObservedRule `json:"rules"` - Notices []ScopeNotice `json:"notices,omitempty"` +type RuleSet struct { + LastPosition int `json:"-"` + Scope Scope `json:"scope"` + Rules []ObservedRule `json:"rules"` + Notices []ScopeNotice `json:"notices,omitempty"` } type Capabilities struct { diff --git a/agent/utils/firewall/filter/normalize.go b/agent/utils/firewall/filter/normalize.go index c47933008c9d..b6bc3a9b3809 100644 --- a/agent/utils/firewall/filter/normalize.go +++ b/agent/utils/firewall/filter/normalize.go @@ -8,7 +8,7 @@ import ( "strings" ) -const MaxAtomicExpansion = 256 +const MaxAtomicExpansion = 500 func NormalizeRule(rule FirewallRule) (FirewallRule, error) { rule.Scope = rule.Scope.Normalize() @@ -16,9 +16,9 @@ func NormalizeRule(rule FirewallRule) (FirewallRule, error) { return FirewallRule{}, err } - if hasCompositeValue(rule.SourceAddress) || hasCompositeValue(rule.DestinationAddress) || - hasCompositeValue(rule.SourcePort) || - (hasCompositeValue(rule.DestinationPort) && !supportsNativeDestinationPortSet(rule.Scope.Provider)) || + if strings.Contains(rule.SourceAddress, ",") || strings.Contains(rule.DestinationAddress, ",") || + strings.Contains(rule.SourcePort, ",") || + (strings.Contains(rule.DestinationPort, ",") && rule.Scope.Provider != ProviderIptables && rule.Scope.Provider != ProviderUFW) || isCompositeProtocol(rule.Protocol) { return FirewallRule{}, fmt.Errorf("%w: expand addresses, ports and protocols before normalization", ErrCompositeRule) } @@ -37,11 +37,11 @@ func NormalizeRule(rule FirewallRule) (FirewallRule, error) { if err != nil { return FirewallRule{}, fmt.Errorf("%w: destination address: %v", ErrInvalidRule, err) } - rule.SourcePort, err = normalizePort(rule.SourcePort) + rule.SourcePort, err = normalizePortValue(rule.SourcePort, false) if err != nil { return FirewallRule{}, fmt.Errorf("%w: source port: %v", ErrInvalidRule, err) } - rule.DestinationPort, err = normalizePortValue(rule.DestinationPort, supportsNativeDestinationPortSet(rule.Scope.Provider)) + rule.DestinationPort, err = normalizePortValue(rule.DestinationPort, rule.Scope.Provider == ProviderIptables || rule.Scope.Provider == ProviderUFW) if err != nil { return FirewallRule{}, fmt.Errorf("%w: destination port: %v", ErrInvalidRule, err) } @@ -136,7 +136,7 @@ func ExpandAtomicRules(input FirewallRule) ([]FirewallRule, error) { destinationAddresses := splitValues(input.DestinationAddress) sourcePorts := splitValues(input.SourcePort) destinationPorts := splitValues(input.DestinationPort) - if supportsNativeDestinationPortSet(input.Scope.Provider) { + if input.Scope.Provider == ProviderIptables || input.Scope.Provider == ProviderUFW { destinationPorts = []string{input.DestinationPort} } @@ -272,10 +272,6 @@ func validateAddressFamily(address netip.Addr, family Family) error { return nil } -func normalizePort(value string) (string, error) { - return normalizePortValue(value, false) -} - func normalizePortValue(value string, allowSet bool) (string, error) { value = strings.TrimSpace(value) if value == "" || strings.EqualFold(value, "any") || strings.EqualFold(value, "anywhere") { @@ -348,10 +344,6 @@ func normalizePortValue(value string, allowSet bool) (string, error) { return fmt.Sprintf("%d-%d", start, end), nil } -func supportsNativeDestinationPortSet(provider Provider) bool { - return provider == ProviderIptables || provider == ProviderUFW -} - func parsePort(value string) (int, error) { port, err := strconv.Atoi(strings.TrimSpace(value)) if err != nil || port < 1 || port > 65535 { @@ -404,10 +396,6 @@ func normalizeConnectionStates(values []string, provider Provider) ([]string, er return states, nil } -func hasCompositeValue(value string) bool { - return strings.Contains(value, ",") -} - func isCompositeProtocol(value string) bool { value = strings.ToLower(strings.TrimSpace(value)) return value == "tcp/udp" || value == "udp/tcp" || strings.Contains(value, ",") diff --git a/agent/utils/firewall/filter/providers/firewalld/adapter.go b/agent/utils/firewall/filter/providers/firewalld/adapter.go index 2428b9ceb8fa..207a758fdb0b 100644 --- a/agent/utils/firewall/filter/providers/firewalld/adapter.go +++ b/agent/utils/firewall/filter/providers/firewalld/adapter.go @@ -15,6 +15,8 @@ import ( "github.com/mattn/go-shellwords" ) +var ErrAlreadyEnabled = errors.New("firewalld rule already exists") + type CommandReader interface { Read(context.Context, ...string) (string, error) } @@ -114,17 +116,17 @@ func parseFirewalldVersion(output string) (int, int, error) { return major, minor, nil } -func (a *Adapter) Observe(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) { +func (a *Adapter) ListRules(ctx context.Context, scope filter.Scope) (filter.RuleSet, error) { scope = scope.Normalize() if err := scope.ValidateMVP(); err != nil { - return filter.Snapshot{}, err + return filter.RuleSet{}, err } if scope.Provider != filter.ProviderFirewalld { - return filter.Snapshot{}, fmt.Errorf("%w: %s", filter.ErrUnsupportedScope, scope.Key()) + return filter.RuleSet{}, fmt.Errorf("%w: %s", filter.ErrUnsupportedScope, scope.Key()) } scope.Family = filter.FamilyInet if a.reader == nil { - return filter.Snapshot{}, errors.New("firewalld reader is required") + return filter.RuleSet{}, errors.New("firewalld reader is required") } var runtime, permanent zoneOutput @@ -141,15 +143,15 @@ func (a *Adapter) Observe(ctx context.Context, scope filter.Scope) (filter.Snaps }() reads.Wait() if err := errors.Join(runtimeErr, permanentErr); err != nil { - return filter.Snapshot{}, err + return filter.RuleSet{}, err } rules, err := mergeZoneObjects(scope, runtime, permanent) if err != nil { - return filter.Snapshot{}, err + return filter.RuleSet{}, err } - snapshot, err := filter.NewSnapshot(scope, rules) + snapshot, err := filter.NewRuleSet(scope, rules) if err != nil { - return filter.Snapshot{}, err + return filter.RuleSet{}, err } snapshot.Notices = publicZoneNotices(runtime, permanent) return snapshot, nil @@ -186,135 +188,144 @@ func (a *Adapter) PrepareRule(rule filter.FirewallRule) (filter.FirewallRule, er return normalized, nil } -func (a *Adapter) Compile(snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, error) { - if snapshot.Revision == "" { - return filter.BackendPlan{}, filter.ErrRuleStale - } +func (a *Adapter) BuildCommands(snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, error) { if err := validateFirewalldScope(snapshot.Scope); err != nil { - return filter.BackendPlan{}, err + return filter.CommandBatch{}, err } - if len(changes) != 1 { - return filter.BackendPlan{}, fmt.Errorf("%w: firewalld plans currently require exactly one change", filter.ErrInvalidRule) + if len(changes) == 0 || len(changes) > filter.MaxAtomicExpansion { + return filter.CommandBatch{}, fmt.Errorf("%w: invalid firewalld batch size", filter.ErrInvalidRule) } - rulePlan, err := a.compileChange(snapshot, changes[0]) - if err != nil { - return filter.BackendPlan{}, err + plan := filter.CommandBatch{Provider: filter.ProviderFirewalld, Scope: snapshot.Scope} + for _, change := range changes { + rulePlan, err := a.compileChange(snapshot, change) + if err != nil { + return filter.CommandBatch{}, err + } + plan.Rules = append(plan.Rules, rulePlan) } - return filter.BackendPlan{ - Provider: filter.ProviderFirewalld, Scope: snapshot.Scope, SnapshotRevision: snapshot.Revision, - Rules: []filter.NativeRulePlan{rulePlan}, - }, nil + return plan, nil } -func (a *Adapter) Apply(ctx context.Context, plan filter.BackendPlan) (filter.ApplyResult, error) { - if plan.Provider != filter.ProviderFirewalld || len(plan.Rules) != 1 { - return filter.ApplyResult{}, fmt.Errorf("%w: invalid firewalld backend plan", filter.ErrInvalidRule) - } - if err := validateFirewalldScope(plan.Scope); err != nil { - return filter.ApplyResult{}, err - } - rulePlan := plan.Rules[0] - if len(rulePlan.Commands) != 0 && a.writer == nil { - return filter.ApplyResult{}, errors.New("firewalld writer is required") +func (a *Adapter) RunCommands(ctx context.Context, plan filter.CommandBatch) error { + commands, err := batchCommands(plan) + if err != nil { + return err } - if len(rulePlan.Commands) != len(rulePlan.RollbackCommands) { - return filter.ApplyResult{}, fmt.Errorf("%w: incomplete firewalld rollback plan", filter.ErrInvalidRule) + if len(commands.Commands) != 0 && a.writer == nil { + return errors.New("firewalld writer is required") } - for _, command := range append(append([]filter.NativeCommand(nil), rulePlan.Commands...), rulePlan.RollbackCommands...) { - if err := validateScopeCommand(plan.Scope, command); err != nil { - return filter.ApplyResult{}, err + executed, alreadyEnabled := 0, 0 + for index, command := range commands.Commands { + err := a.writer.Run(ctx, command) + if errors.Is(err, ErrAlreadyEnabled) { + alreadyEnabled++ + err = nil } - } - executed := 0 - for index, command := range rulePlan.Commands { - if err := a.writer.Run(ctx, command); err != nil { + if err != nil { if plan.CommandOnly { - return filter.ApplyResult{}, err + return err } - return filter.ApplyResult{}, a.compensate(ctx, rulePlan, executed, err) + return a.compensate(ctx, commands, executed, err) } executed = index + 1 } - return filter.ApplyResult{Applied: []filter.ObservedRule{rulePlan.Expected}}, nil + if commands.Operation == filter.ChangeCreate && alreadyEnabled > 0 && alreadyEnabled == len(commands.Commands) { + return ErrAlreadyEnabled + } + return nil } -func (a *Adapter) Verify(ctx context.Context, plan filter.BackendPlan) (filter.VerifyResult, error) { - if plan.Provider != filter.ProviderFirewalld || len(plan.Rules) != 1 { - return filter.VerifyResult{}, fmt.Errorf("%w: invalid firewalld backend plan", filter.ErrInvalidRule) - } - snapshot, err := a.Observe(ctx, plan.Scope) +func (a *Adapter) Rollback(ctx context.Context, plan filter.CommandBatch) error { + commands, err := batchCommands(plan) if err != nil { - return filter.VerifyResult{}, err - } - rulePlan := plan.Rules[0] - if rulePlan.Operation == filter.ChangeDelete { - return filter.VerifyResult{Snapshot: snapshot, Matched: countCanonical(snapshot, rulePlan.Previous) == 0}, nil - } - if rulePlan.Operation == filter.ChangeUpdate && rulePlan.Previous != nil && - rulePlan.Previous.Locator.Canonical != rulePlan.Expected.Locator.Canonical && countCanonical(snapshot, rulePlan.Previous) != 0 { - return filter.VerifyResult{Snapshot: snapshot, Matched: false}, nil + return err } - matches := 0 - for _, observed := range snapshot.Rules { - if observed.Locator.Canonical != rulePlan.Expected.Locator.Canonical || observed.Persistence != filter.PersistenceStatusConverged { - continue - } - want, wantErr := filter.RuleKey(rulePlan.Expected.Rule) - got, gotErr := filter.RuleKey(observed.Rule) - if wantErr == nil && gotErr == nil && want == got { - matches++ - } + if a.writer == nil { + return errors.New("firewalld writer is required") } - return filter.VerifyResult{Snapshot: snapshot, Matched: matches == 1}, nil + return a.rollback(ctx, plan.Scope, commands, len(commands.RollbackCommands)) } -func (a *Adapter) Rollback(ctx context.Context, plan filter.BackendPlan) error { - if plan.Provider != filter.ProviderFirewalld { - return fmt.Errorf("%w: invalid firewalld backend plan", filter.ErrInvalidRule) +func batchCommands(plan filter.CommandBatch) (filter.RuleCommands, error) { + if plan.Provider != filter.ProviderFirewalld || len(plan.Rules) == 0 || len(plan.Rules) > filter.MaxAtomicExpansion { + return filter.RuleCommands{}, fmt.Errorf("%w: invalid firewalld backend plan", filter.ErrInvalidRule) } if err := validateFirewalldScope(plan.Scope); err != nil { - return err + return filter.RuleCommands{}, err } - if a.writer == nil { - return errors.New("firewalld writer is required") + result := filter.RuleCommands{Expected: filter.ObservedRule{Rule: filter.FirewallRule{Scope: plan.Scope}}} + if plan.CreatesOnly() { + result.Operation = filter.ChangeCreate } - for ruleIndex := len(plan.Rules) - 1; ruleIndex >= 0; ruleIndex-- { - rulePlan := plan.Rules[ruleIndex] - if err := a.rollback(ctx, plan.Scope, rulePlan, len(rulePlan.RollbackCommands)); err != nil { - return err + for _, rule := range plan.Rules { + if len(rule.Commands) != len(rule.RollbackCommands) { + return filter.RuleCommands{}, fmt.Errorf("%w: incomplete firewalld rollback plan", filter.ErrInvalidRule) + } + for _, commands := range [][]filter.NativeCommand{rule.Commands, rule.RollbackCommands} { + for _, command := range commands { + if len(command.Args) != 2 || command.Args[0] != "--zone="+filter.FirewalldInputZone { + return filter.RuleCommands{}, fmt.Errorf("%w: expected one firewalld rule option", filter.ErrInvalidRule) + } + if err := validateScopeCommand(plan.Scope, command); err != nil { + return filter.RuleCommands{}, err + } + } } } - return nil + for _, permanent := range []bool{false, true} { + previousOperation, commandBytes := "", 0 + for _, rule := range plan.Rules { + for index, command := range rule.Commands { + option := command.Args[len(command.Args)-1] + rollback := rule.RollbackCommands[index].Args[len(rule.RollbackCommands[index].Args)-1] + operation := strings.SplitN(strings.TrimPrefix(option, "--"), "-", 2)[0] + if operation != previousOperation || commandBytes+max(len(option), len(rollback))+1 > 64*1024 { + args := []string{"--zone=" + filter.FirewalldInputZone} + if permanent { + args = append(args, "--permanent") + } + result.Commands = append(result.Commands, filter.NativeCommand{Executable: "firewall-cmd", Args: append([]string(nil), args...)}) + result.RollbackCommands = append(result.RollbackCommands, filter.NativeCommand{Executable: "firewall-cmd", Args: append([]string(nil), args...)}) + previousOperation, commandBytes = operation, 0 + } + last := len(result.Commands) - 1 + result.Commands[last].Args = append(result.Commands[last].Args, option) + result.RollbackCommands[last].Args = append(result.RollbackCommands[last].Args, rollback) + commandBytes += max(len(option), len(rollback)) + 1 + } + } + } + return result, nil } -func (a *Adapter) compileChange(snapshot filter.Snapshot, change filter.DesiredChange) (filter.NativeRulePlan, error) { +func (a *Adapter) compileChange(snapshot filter.RuleSet, change filter.RuleChange) (filter.RuleCommands, error) { rule := change.After if change.Operation == filter.ChangeDelete { rule = change.Before } if rule == nil { - return filter.NativeRulePlan{}, fmt.Errorf("%w: %s rule is required", filter.ErrInvalidRule, change.Operation) + return filter.RuleCommands{}, fmt.Errorf("%w: %s rule is required", filter.ErrInvalidRule, change.Operation) } normalized, err := a.PrepareRule(*rule) if err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } if normalized.Scope.Key() != snapshot.Scope.Key() { - return filter.NativeRulePlan{}, fmt.Errorf("%w: change scope %s", filter.ErrUnsupportedScope, normalized.Scope.Key()) + return filter.RuleCommands{}, fmt.Errorf("%w: change scope %s", filter.ErrUnsupportedScope, normalized.Scope.Key()) } if normalized.UUID == "" { - return filter.NativeRulePlan{}, fmt.Errorf("%w: rule UUID is required", filter.ErrInvalidRule) + return filter.RuleCommands{}, fmt.Errorf("%w: rule UUID is required", filter.ErrInvalidRule) } expected := observedForRule(normalized) - plan := filter.NativeRulePlan{RuleUUID: normalized.UUID, Operation: change.Operation, Expected: expected} + plan := filter.RuleCommands{RuleUUID: normalized.UUID, Operation: change.Operation, Expected: expected} switch change.Operation { case filter.ChangeCreate: - plan.Commands, plan.RollbackCommands = missingRuleCommands(snapshot, normalized) + plan.Commands, plan.RollbackCommands = ruleCommands(nativeOption(normalized, "add"), nativeOption(normalized, "remove")) case filter.ChangeAdopt: target, targetErr := validateMutationTarget(snapshot, change, normalized, false) if targetErr != nil { - return filter.NativeRulePlan{}, targetErr + return filter.RuleCommands{}, targetErr } plan.Previous = &target plan.Expected = target @@ -322,35 +333,31 @@ func (a *Adapter) compileChange(snapshot filter.Snapshot, change filter.DesiredC case filter.ChangeUpdate: target, targetErr := validateMutationTarget(snapshot, change, normalized, true) if targetErr != nil { - return filter.NativeRulePlan{}, targetErr + return filter.RuleCommands{}, targetErr } plan.Previous = &target if target.Locator.Canonical == expected.Locator.Canonical { break } - removeCommands, restoreCommands := observedPairedCommands(target, "remove", "add") - addCommands, removeNewCommands := missingRuleCommands(snapshot, normalized) + removeCommands, restoreCommands := observedRuleCommands(target, "remove", "add") + addCommands, removeNewCommands := ruleCommands(nativeOption(normalized, "add"), nativeOption(normalized, "remove")) plan.Commands = append(removeCommands, addCommands...) plan.RollbackCommands = append(restoreCommands, removeNewCommands...) case filter.ChangeDelete: + if change.CommandOnly && change.Locator == nil { + plan.Commands, plan.RollbackCommands = ruleCommands(nativeOption(normalized, "remove"), nativeOption(normalized, "add")) + break + } target, targetErr := validateMutationTarget(snapshot, change, normalized, true) if targetErr != nil { - return filter.NativeRulePlan{}, targetErr + return filter.RuleCommands{}, targetErr } plan.Previous = &target plan.Expected = target plan.Expected.Rule.UUID = normalized.UUID - plan.Commands, plan.RollbackCommands = observedPairedCommands(target, "remove", "add") - if change.CommandOnly { - switch target.Persistence { - case filter.PersistenceStatusRuntimeOnly: - plan.Commands, plan.RollbackCommands = plan.Commands[:1], plan.RollbackCommands[:1] - case filter.PersistenceStatusPermanentOnly: - plan.Commands, plan.RollbackCommands = plan.Commands[1:], plan.RollbackCommands[1:] - } - } + plan.Commands, plan.RollbackCommands = observedRuleCommands(target, "remove", "add") default: - return filter.NativeRulePlan{}, fmt.Errorf("%w: unsupported operation %s", filter.ErrInvalidRule, change.Operation) + return filter.RuleCommands{}, fmt.Errorf("%w: unsupported operation %s", filter.ErrInvalidRule, change.Operation) } return plan, nil } @@ -407,52 +414,22 @@ func nativeCanonical(rule filter.FirewallRule) string { return "rich:" + canonicalRichRule(rule) } -func missingRuleCommands(snapshot filter.Snapshot, rule filter.FirewallRule) ([]filter.NativeCommand, []filter.NativeCommand) { - commands, rollback := pairedCommands(rule, "add", "remove") - canonical := nativeCanonical(rule) - var runtimeExists, permanentExists bool - for _, observed := range snapshot.Rules { - if observed.Locator.Canonical != canonical { - continue - } - runtimeExists = runtimeExists || observed.Persistence == filter.PersistenceStatusConverged || observed.Persistence == filter.PersistenceStatusRuntimeOnly - permanentExists = permanentExists || observed.Persistence == filter.PersistenceStatusConverged || observed.Persistence == filter.PersistenceStatusPermanentOnly - if runtimeExists && permanentExists { - break - } - } - var changes, inverses []filter.NativeCommand - for index, exists := range []bool{runtimeExists, permanentExists} { - if !exists { - changes = append(changes, commands[index]) - inverses = append(inverses, rollback[index]) - } - } - return changes, inverses -} - -func pairedCommands(rule filter.FirewallRule, operation, inverse string) ([]filter.NativeCommand, []filter.NativeCommand) { - return pairedNativeCommands(rule.Scope, nativeOption(rule, operation), nativeOption(rule, inverse)) -} - -func observedPairedCommands(observed filter.ObservedRule, operation, inverse string) ([]filter.NativeCommand, []filter.NativeCommand) { +func observedRuleCommands(observed filter.ObservedRule, operation, inverse string) ([]filter.NativeCommand, []filter.NativeCommand) { if observed.Rule.NativeKind != filter.NativeKindRichRule || observed.Raw == "" { - return pairedCommands(observed.Rule, operation, inverse) + return ruleCommands(nativeOption(observed.Rule, operation), nativeOption(observed.Rule, inverse)) } - return pairedNativeCommands(observed.Rule.Scope, + return ruleCommands( "--"+operation+"-rich-rule="+observed.Raw, "--"+inverse+"-rich-rule="+observed.Raw) } -func pairedNativeCommands(scope filter.Scope, option, rollback string) ([]filter.NativeCommand, []filter.NativeCommand) { - selector := scopeSelector(scope) +func ruleCommands(option, rollback string) ([]filter.NativeCommand, []filter.NativeCommand) { + selector := "--zone=" + filter.FirewalldInputZone commands := []filter.NativeCommand{ {Executable: "firewall-cmd", Args: []string{selector, option}}, - {Executable: "firewall-cmd", Args: []string{"--permanent", selector, option}}, } rollbackCommands := []filter.NativeCommand{ {Executable: "firewall-cmd", Args: []string{selector, rollback}}, - {Executable: "firewall-cmd", Args: []string{"--permanent", selector, rollback}}, } return commands, rollbackCommands } @@ -464,7 +441,7 @@ func nativeOption(rule filter.FirewallRule, operation string) string { return "--" + operation + "-rich-rule=" + canonicalRichRule(rule) } -func validateMutationTarget(snapshot filter.Snapshot, change filter.DesiredChange, normalized filter.FirewallRule, requireOwned bool) (filter.ObservedRule, error) { +func validateMutationTarget(snapshot filter.RuleSet, change filter.RuleChange, normalized filter.FirewallRule, requireOwned bool) (filter.ObservedRule, error) { if change.Locator == nil || change.Locator.Canonical == "" { return filter.ObservedRule{}, fmt.Errorf("%w: firewalld mutation requires canonical locator", filter.ErrInvalidRule) } @@ -487,8 +464,7 @@ func validateMutationTarget(snapshot filter.Snapshot, change filter.DesiredChang if target.Protected { return filter.ObservedRule{}, filter.ErrProtectedRule } - if target.ParseStatus != filter.ParseStatusSupported || - (target.Persistence != filter.PersistenceStatusConverged && !(change.CommandOnly && change.Operation == filter.ChangeDelete)) { + if target.ParseStatus != filter.ParseStatusSupported { return filter.ObservedRule{}, filter.ErrRuleStale } want := normalized @@ -511,30 +487,37 @@ func validateMutationTarget(snapshot filter.Snapshot, change filter.DesiredChang } func validateScopeCommand(scope filter.Scope, command filter.NativeCommand) error { - if command.Executable != "firewall-cmd" { - return fmt.Errorf("%w: unexpected firewalld executable %q", filter.ErrInvalidRule, command.Executable) + if err := validateFirewalldScope(scope); err != nil { + return err + } + if command.Executable != "firewall-cmd" || command.Stdin != "" { + return fmt.Errorf("%w: invalid firewalld command", filter.ErrInvalidRule) } - expected := scopeSelector(scope) - foundSelector := false + expected := "--zone=" + filter.FirewalldInputZone + foundSelector, options := false, 0 for _, arg := range command.Args { - if arg == expected { + switch { + case arg == expected: foundSelector = true - } - if (strings.HasPrefix(arg, "--zone=") || strings.HasPrefix(arg, "--policy=")) && arg != expected { - return fmt.Errorf("%w: firewalld command targets another scope", filter.ErrUnsupportedScope) + case arg == "--permanent": + case strings.HasPrefix(arg, "--add-port="), strings.HasPrefix(arg, "--remove-port="), + strings.HasPrefix(arg, "--add-rich-rule="), strings.HasPrefix(arg, "--remove-rich-rule="), + strings.HasPrefix(arg, "--add-service="), strings.HasPrefix(arg, "--remove-service="): + options++ + default: + return fmt.Errorf("%w: unexpected firewalld argument %q", filter.ErrInvalidRule, arg) } } if !foundSelector { return fmt.Errorf("%w: firewalld command must explicitly target %s", filter.ErrUnsupportedScope, expected) } + if options == 0 { + return fmt.Errorf("%w: firewalld command requires rule options", filter.ErrInvalidRule) + } return nil } -func scopeSelector(scope filter.Scope) string { - return "--zone=" + filter.FirewalldInputZone -} - -func (a *Adapter) compensate(ctx context.Context, plan filter.NativeRulePlan, executed int, cause error) error { +func (a *Adapter) compensate(ctx context.Context, plan filter.RuleCommands, executed int, cause error) error { if plan.Operation == filter.ChangeCreate { return cause } @@ -547,7 +530,7 @@ func (a *Adapter) compensate(ctx context.Context, plan filter.NativeRulePlan, ex return cause } -func (a *Adapter) rollback(ctx context.Context, scope filter.Scope, plan filter.NativeRulePlan, executed int) error { +func (a *Adapter) rollback(ctx context.Context, scope filter.Scope, plan filter.RuleCommands, executed int) error { var rollbackErr error for index := executed - 1; index >= 0; index-- { if index >= len(plan.RollbackCommands) { @@ -563,26 +546,13 @@ func (a *Adapter) rollback(ctx context.Context, scope filter.Scope, plan filter. } continue } - if err := a.writer.Run(ctx, command); err != nil && rollbackErr == nil { + if err := a.writer.Run(ctx, command); err != nil && !errors.Is(err, ErrAlreadyEnabled) && rollbackErr == nil { rollbackErr = err } } return rollbackErr } -func countCanonical(snapshot filter.Snapshot, expected *filter.ObservedRule) int { - if expected == nil { - return 0 - } - count := 0 - for _, observed := range snapshot.Rules { - if observed.Locator.Canonical == expected.Locator.Canonical { - count++ - } - } - return count -} - type zoneOutput struct { ports string rich string @@ -594,7 +564,7 @@ func (a *Adapter) readScope(ctx context.Context, scope filter.Scope, permanent b if permanent { args = append(args, "--permanent") } - args = append(args, scopeSelector(scope), "--list-all") + args = append(args, "--zone="+filter.FirewalldInputZone, "--list-all") output, err := a.reader.Read(ctx, args...) if err != nil { return zoneOutput{}, err @@ -1079,9 +1049,46 @@ func (systemBackend) Run(ctx context.Context, command filter.NativeCommand) erro if err := validateSystemCommand(command); err != nil { return err } - return cmd.NewCommandMgr( - cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second), cmd.WithEnv("LANGUAGE=en_US:en"), + options, removals, alreadyEnabled := 0, 0, 0 + for _, arg := range command.Args { + if strings.HasPrefix(arg, "--add-") || strings.HasPrefix(arg, "--remove-") { + options++ + } + if strings.HasPrefix(arg, "--remove-") { + removals++ + } + } + timeout := 60 * time.Second + if options > 1 { + timeout = 5 * time.Minute + } + var stderr strings.Builder + err := cmd.NewCommandMgr( + cmd.WithContext(ctx), cmd.WithTimeout(timeout), cmd.WithEnv("LC_ALL=C", "LANGUAGE=en_US:en"), cmd.WithStderr(&stderr), ).RunWithOptionalSudo(command.Executable, command.Args...) + if err != nil { + return err + } + for _, line := range strings.Split(stderr.String(), "\n") { + line = strings.TrimSpace(line) + if line == "" { + continue + } + switch { + case strings.HasPrefix(line, "Warning: ALREADY_ENABLED:"): + alreadyEnabled++ + case strings.HasPrefix(line, "Warning: NOT_ENABLED:"): + if removals != options { + return fmt.Errorf("%w: %s", filter.ErrRuleStale, stderr.String()) + } + default: + return fmt.Errorf("firewalld rule batch failed: %s", stderr.String()) + } + } + if alreadyEnabled > 0 && alreadyEnabled == options { + return ErrAlreadyEnabled + } + return nil } func validateSystemCommand(command filter.NativeCommand) error { @@ -1093,35 +1100,3 @@ func validateSystemCommand(command filter.NativeCommand) error { } return fmt.Errorf("%w: firewalld command must target the managed input zone", filter.ErrUnsupportedScope) } - -type createPlanner struct { - adapter *Adapter - snapshot filter.Snapshot - byCanonical map[string][]filter.ObservedRule -} - -func (a *Adapter) NewCreatePlanner(snapshot filter.Snapshot) filter.CreatePlanner { - byCanonical := make(map[string][]filter.ObservedRule, len(snapshot.Rules)) - for _, observed := range snapshot.Rules { - byCanonical[observed.Locator.Canonical] = append(byCanonical[observed.Locator.Canonical], observed) - } - snapshot.Rules = nil - return &createPlanner{adapter: a, snapshot: snapshot, byCanonical: byCanonical} -} - -func (p *createPlanner) Compile(change filter.DesiredChange) (filter.BackendPlan, error) { - if change.Operation != filter.ChangeCreate || change.After == nil { - return filter.BackendPlan{}, filter.ErrInvalidRule - } - rule, err := p.adapter.PrepareRule(*change.After) - if err != nil { - return filter.BackendPlan{}, err - } - snapshot := p.snapshot - snapshot.Rules = p.byCanonical[nativeCanonical(rule)] - return p.adapter.Compile(snapshot, []filter.DesiredChange{change}) -} - -func (p *createPlanner) Applied(rule filter.ObservedRule) { - p.byCanonical[rule.Locator.Canonical] = []filter.ObservedRule{rule} -} diff --git a/agent/utils/firewall/filter/providers/iptables/adapter.go b/agent/utils/firewall/filter/providers/iptables/adapter.go index c7895341f1ac..22f5811a3ab9 100644 --- a/agent/utils/firewall/filter/providers/iptables/adapter.go +++ b/agent/utils/firewall/filter/providers/iptables/adapter.go @@ -6,7 +6,6 @@ import ( "fmt" "strconv" "strings" - "sync" "time" "github.com/1Panel-dev/1Panel/agent/utils/cmd" @@ -17,140 +16,187 @@ import ( ) type RuleReader interface { + filter.CommentRuleReader ListChain(context.Context, filter.Scope) (string, error) } +type TableReader interface { + ListTable(context.Context, filter.Scope) (string, error) +} + type RuleWriter interface { Run(context.Context, filter.NativeCommand) error Save(context.Context, filter.Scope) error } -type MultiportChecker interface { - CheckMultiport(context.Context, filter.Family) error -} - type Adapter struct { - reader RuleReader - writer RuleWriter - checker MultiportChecker - multiportMu sync.Mutex - multiportOK map[filter.Family]bool + reader RuleReader + writer RuleWriter } func NewAdapter() *Adapter { backend := systemBackend{} - return &Adapter{reader: backend, writer: backend, checker: backend} + return &Adapter{reader: backend, writer: backend} } func NewAdapterWithReader(reader RuleReader) *Adapter { - adapter := &Adapter{reader: reader} - adapter.checker, _ = reader.(MultiportChecker) - return adapter + return &Adapter{reader: reader} } func NewAdapterWithBackend(reader RuleReader, writer RuleWriter) *Adapter { - adapter := &Adapter{reader: reader, writer: writer} - if checker, ok := reader.(MultiportChecker); ok { - adapter.checker = checker - } else if checker, ok := writer.(MultiportChecker); ok { - adapter.checker = checker - } - return adapter + return &Adapter{reader: reader, writer: writer} } func (a *Adapter) Provider() filter.Provider { return filter.ProviderIptables } -func (a *Adapter) CheckRule(ctx context.Context, rule filter.FirewallRule) error { - if !strings.Contains(rule.DestinationPort, ",") && !strings.Contains(rule.SourcePort, ",") { - return nil - } - if rule.Protocol != "tcp" && rule.Protocol != "udp" { - return fmt.Errorf("%w: iptables multiport requires tcp or udp", filter.ErrInvalidRule) - } - if a.checker == nil { - return nil - } - a.multiportMu.Lock() - defer a.multiportMu.Unlock() - if a.multiportOK[rule.Scope.Family] { - return nil - } - if err := a.checker.CheckMultiport(ctx, rule.Scope.Family); err != nil { - return fmt.Errorf("inspect iptables multiport for %s: %w", rule.Scope.Family, err) - } - if a.multiportOK == nil { - a.multiportOK = make(map[filter.Family]bool, 2) - } - a.multiportOK[rule.Scope.Family] = true - return nil -} - func (a *Adapter) Capabilities(context.Context) (filter.Capabilities, error) { return filter.Capabilities{ Marker: true, OwnedChains: true, ExplicitPosition: true, }, nil } -func (a *Adapter) Observe(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) { - scope = scope.Normalize() - if err := scope.ValidateMVP(); err != nil { - return filter.Snapshot{}, err +func (a *Adapter) AppendUnverified(ctx context.Context, rule filter.FirewallRule, comment string) error { + rule, err := filter.NormalizeRule(rule) + if err != nil { + return err } - if scope.Provider != filter.ProviderIptables { - return filter.Snapshot{}, fmt.Errorf("%w: %s", filter.ErrUnsupportedScope, scope.Key()) + if err := validateAdapterScope(rule.Scope); err != nil { + return err } - if a.reader == nil { - return filter.Snapshot{}, fmt.Errorf("iptables reader is required") + args := []string{"-w", "-t", rule.Scope.Table, "-A", rule.Scope.Chain} + args = append(args, compileRuleArgs(rule, comment)...) + return a.writer.Run(ctx, filter.NativeCommand{Executable: executableForFamily(rule.Scope.Family), Args: args}) +} + +func (a *Adapter) ListRulesByComment(ctx context.Context, scopes []filter.Scope, comment string) ([]filter.ObservedRule, error) { + var rules []filter.ObservedRule + for _, scope := range scopes { + scope = scope.Normalize() + if err := validateAdapterScope(scope); err != nil { + return nil, err + } + output, err := a.reader.ReadRulesByComment(ctx, scope, comment) + if err != nil { + return nil, err + } + rules = append(rules, parseChainRules(scope, output)...) } - output, err := a.reader.ListChain(ctx, scope) + return rules, nil +} + +func (a *Adapter) ListRules(ctx context.Context, scope filter.Scope) (filter.RuleSet, error) { + snapshots, err := a.ListRuleScopes(ctx, []filter.Scope{scope}) if err != nil { - if errors.Is(err, filter.ErrProviderUnavailable) { - snapshot, snapshotErr := filter.NewSnapshot(scope, nil) - if snapshotErr != nil { - return filter.Snapshot{}, snapshotErr + return filter.RuleSet{}, err + } + return snapshots[0], nil +} + +func (a *Adapter) ListRuleScopes(ctx context.Context, scopes []filter.Scope) ([]filter.RuleSet, error) { + normalized := make([]filter.Scope, len(scopes)) + for index, scope := range scopes { + scope = scope.Normalize() + if err := validateAdapterScope(scope); err != nil { + return nil, err + } + normalized[index] = scope + } + if a.reader == nil { + return nil, fmt.Errorf("iptables reader is required") + } + tableReader, readsTable := a.reader.(TableReader) + snapshots := make([]filter.RuleSet, len(scopes)) + for index, scope := range normalized { + if snapshots[index].Scope.Provider != "" { + continue + } + var output string + var err error + if readsTable { + output, err = tableReader.ListTable(ctx, scope) + } else { + output, err = a.reader.ListChain(ctx, scope) + } + if err != nil && (readsTable || !errors.Is(err, filter.ErrProviderUnavailable)) { + return nil, err + } + for target := index; target < len(normalized); target++ { + current := normalized[target] + if readsTable { + if current.Family != scope.Family || current.Table != scope.Table { + continue + } + } else if target != index { + continue } - snapshot.Notices = []filter.ScopeNotice{{ - Code: filter.ScopeNoticeManagedScopeMissing, Values: []string{string(scope.Family), scope.Chain}, - }} - return snapshot, nil + missing := errors.Is(err, filter.ErrProviderUnavailable) || readsTable && !containsChainDeclaration(output, current.Chain) + var rules []filter.ObservedRule + if !missing { + rules = parseChainRules(current, output) + } + snapshot, buildErr := filter.NewRuleSet(current, rules) + if buildErr != nil { + return nil, buildErr + } + if missing { + snapshot.Notices = []filter.ScopeNotice{{Code: filter.ScopeNoticeManagedScopeMissing, Values: []string{string(current.Family), current.Chain}}} + } + snapshots[target] = snapshot } - return filter.Snapshot{}, err } - rules := parseChainRules(scope, output) - return filter.NewSnapshot(scope, rules) + return snapshots, nil } -func (a *Adapter) Compile(snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, error) { - if snapshot.Revision == "" { - return filter.BackendPlan{}, filter.ErrRuleStale - } +func (a *Adapter) BuildCommands(snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, error) { snapshot.Scope = snapshot.Scope.Normalize() if err := validateAdapterScope(snapshot.Scope); err != nil { - return filter.BackendPlan{}, err + return filter.CommandBatch{}, err } if len(changes) == 0 { - return filter.BackendPlan{}, fmt.Errorf("%w: iptables plan requires at least one change", filter.ErrInvalidRule) + return filter.CommandBatch{}, fmt.Errorf("%w: iptables plan requires at least one change", filter.ErrInvalidRule) + } + createOnly, deleteOnly := true, true + for _, change := range changes { + createOnly = createOnly && change.Operation == filter.ChangeCreate && change.CommandOnly + deleteOnly = deleteOnly && change.Operation == filter.ChangeDelete && change.CommandOnly + } + if createOnly { + return compileCreateBatch(snapshot, changes) + } + externalDelete := len(changes) == 1 && changes[0].Locator == nil && (changes[0].UnmarkedAdopted || changes[0].PreviousMarker != "") + if deleteOnly && !externalDelete { + return compileDeleteBatch(snapshot, changes) + } + if len(changes) != 1 { + return filter.CommandBatch{}, fmt.Errorf("%w: iptables mutation requires exactly one change", filter.ErrInvalidRule) + } + rulePlan, err := compileChange(snapshot, changes[0]) + if err != nil { + return filter.CommandBatch{}, err } - return compileBatch(snapshot, changes) + return filter.CommandBatch{ + Provider: filter.ProviderIptables, Scope: snapshot.Scope, CommandOnly: changes[0].CommandOnly, + Rules: []filter.RuleCommands{rulePlan}, + }, nil } -func (a *Adapter) Apply(ctx context.Context, plan filter.BackendPlan) (filter.ApplyResult, error) { +func (a *Adapter) RunCommands(ctx context.Context, plan filter.CommandBatch) error { if a.writer == nil { - return filter.ApplyResult{}, fmt.Errorf("iptables writer is required") + return fmt.Errorf("iptables writer is required") } if plan.Provider != filter.ProviderIptables { - return filter.ApplyResult{}, fmt.Errorf("%w: backend plan provider %q", filter.ErrUnsupportedScope, plan.Provider) + return fmt.Errorf("%w: backend plan provider %q", filter.ErrUnsupportedScope, plan.Provider) } if err := validateAdapterScope(plan.Scope); err != nil { - return filter.ApplyResult{}, err + return err } if len(plan.Rules) == 0 { - return filter.ApplyResult{}, fmt.Errorf("%w: iptables plan requires at least one rule", filter.ErrInvalidRule) + return fmt.Errorf("%w: iptables plan requires at least one rule", filter.ErrInvalidRule) } for _, rulePlan := range plan.Rules { for _, command := range rulePlan.Commands { if err := validateNativeCommand(plan.Scope, command); err != nil { - return filter.ApplyResult{}, err + return err } } } @@ -159,142 +205,69 @@ func (a *Adapter) Apply(ctx context.Context, plan filter.BackendPlan) (filter.Ap for _, command := range rulePlan.Commands { if err := a.writer.Run(ctx, command); err != nil { if plan.CommandOnly { - return filter.ApplyResult{}, err + return err } - return filter.ApplyResult{}, a.compensate(ctx, plan, ruleIndex, executed, err) + return a.compensate(ctx, plan, ruleIndex, executed, err) } executed++ } } - if err := a.writer.Save(ctx, plan.Scope); err != nil { - if plan.CommandOnly { - return filter.ApplyResult{}, err - } - return filter.ApplyResult{}, a.compensate(ctx, plan, len(plan.Rules)-1, -1, err) - } - applied := make([]filter.ObservedRule, 0, len(plan.Rules)) - for _, rulePlan := range plan.Rules { - applied = append(applied, rulePlan.Expected) - } - return filter.ApplyResult{Applied: applied}, nil + + return nil } -func compileBatch(snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, error) { - plan := filter.BackendPlan{ - Provider: filter.ProviderIptables, Scope: snapshot.Scope, SnapshotRevision: snapshot.Revision, - Rules: make([]filter.NativeRulePlan, 0, len(changes)), - } - current := snapshot - current.Rules = make([]filter.ObservedRule, len(snapshot.Rules), len(snapshot.Rules)+len(changes)) - copy(current.Rules, snapshot.Rules) +func compileDeleteBatch(snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, error) { + plan := filter.CommandBatch{Provider: filter.ProviderIptables, Scope: snapshot.Scope, CommandOnly: true} + var script strings.Builder + fmt.Fprintf(&script, "*%s\n", snapshot.Scope.Table) for _, change := range changes { - rulePlan, err := compileChange(current, change) + rulePlan, err := compileChange(snapshot, change) if err != nil { - return filter.BackendPlan{}, err + return filter.CommandBatch{}, err } - rulePlan.Commands = nil - rulePlan.RollbackCommands = nil - plan.Rules = append(plan.Rules, rulePlan) - current, err = applyRestoreRulePlan(current, rulePlan) + line, err := restoreRuleLine(snapshot.Scope, *rulePlan.Previous) if err != nil { - return filter.BackendPlan{}, err + return filter.CommandBatch{}, err } + script.WriteString(strings.Replace(line, "-A ", "-D ", 1)) + script.WriteByte('\n') + rulePlan.Commands, rulePlan.RollbackCommands = nil, nil + plan.Rules = append(plan.Rules, rulePlan) } - - if _, err := filter.NewSnapshot(snapshot.Scope, current.Rules); err != nil { - return filter.BackendPlan{}, err - } - applyScript, err := buildRestoreScript(snapshot.Scope, current.Rules) - if err != nil { - return filter.BackendPlan{}, err - } - rollbackScript, err := buildRestoreScript(snapshot.Scope, snapshot.Rules) - if err != nil { - return filter.BackendPlan{}, err - } + script.WriteString("COMMIT\n") plan.Rules[0].Commands = []filter.NativeCommand{{ - Executable: restoreExecutableForFamily(snapshot.Scope.Family), Args: []string{"--noflush", "--wait"}, Stdin: applyScript, - }} - plan.Rules[0].RollbackCommands = []filter.NativeCommand{{ - Executable: restoreExecutableForFamily(snapshot.Scope.Family), Args: []string{"--noflush", "--wait"}, Stdin: rollbackScript, + Executable: restoreExecutableForFamily(snapshot.Scope.Family), Args: []string{"--noflush", "--wait"}, Stdin: script.String(), }} return plan, nil } -func applyRestoreRulePlan(snapshot filter.Snapshot, plan filter.NativeRulePlan) (filter.Snapshot, error) { - position := plan.Expected.Locator.Position - if plan.Operation == filter.ChangeDelete && plan.Previous != nil { - position = plan.Previous.Locator.Position - } - if position == nil { - return filter.Snapshot{}, fmt.Errorf("%w: batch rule has no target position", filter.ErrInvalidRule) - } - nativePosition := *position - rules := snapshot.Rules - firstChanged := nativePosition - 1 - switch plan.Operation { - case filter.ChangeCreate: - if nativePosition < 1 || nativePosition > len(rules)+1 { - return filter.Snapshot{}, fmt.Errorf("%w: batch target position %d is out of range", filter.ErrInvalidRule, nativePosition) - } - expected := plan.Expected - expected.Raw = "" - rules = append(rules, filter.ObservedRule{}) - copy(rules[nativePosition:], rules[nativePosition-1:]) - rules[nativePosition-1] = expected - case filter.ChangeDelete: - if nativePosition < 1 || nativePosition > len(rules) { - return filter.Snapshot{}, fmt.Errorf("%w: batch target position %d is out of range", filter.ErrRuleStale, nativePosition) - } - rules = append(rules[:nativePosition-1], rules[nativePosition:]...) - case filter.ChangeAdopt, filter.ChangeUpdate, filter.ChangeReorder: - if plan.Previous == nil || plan.Previous.Locator.Position == nil { - return filter.Snapshot{}, fmt.Errorf("%w: mutation has no previous position", filter.ErrInvalidRule) - } - previousPosition := *plan.Previous.Locator.Position - firstChanged = min(firstChanged, previousPosition-1) - if previousPosition < 1 || previousPosition > len(rules) || nativePosition < 1 || nativePosition > len(rules) { - return filter.Snapshot{}, fmt.Errorf("%w: mutation position is out of range", filter.ErrRuleStale) - } - expected := plan.Expected - expected.Raw = "" - if previousPosition == nativePosition { - rules[previousPosition-1] = expected - break - } - rules = append(rules[:previousPosition-1], rules[previousPosition:]...) - rules = append(rules, filter.ObservedRule{}) - copy(rules[nativePosition:], rules[nativePosition-1:]) - rules[nativePosition-1] = expected - default: - return filter.Snapshot{}, fmt.Errorf("%w: unsupported batch operation %s", filter.ErrInvalidRule, plan.Operation) - } - for index := firstChanged; index < len(rules); index++ { - position := index + 1 - rules[index].Locator.Position = &position - } - snapshot.Rules = rules - return snapshot, nil -} - -func buildRestoreScript(scope filter.Scope, rules []filter.ObservedRule) (string, error) { +func compileCreateBatch(snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, error) { + plan := filter.CommandBatch{Provider: filter.ProviderIptables, Scope: snapshot.Scope, CommandOnly: true} var script strings.Builder - script.WriteByte('*') - script.WriteString(scope.Table) - script.WriteByte('\n') - script.WriteString("-F ") - script.WriteString(scope.Chain) - script.WriteByte('\n') - for _, observed := range rules { - line, err := restoreRuleLine(scope, observed) + fmt.Fprintf(&script, "*%s\n", snapshot.Scope.Table) + for _, change := range changes { + rulePlan, err := compileChange(snapshot, change) + if err != nil { + return filter.CommandBatch{}, err + } + line, err := restoreRuleLine(snapshot.Scope, rulePlan.Expected) if err != nil { - return "", err + return filter.CommandBatch{}, err + } + if !change.Append { + line = strings.Replace(line, "-A "+snapshot.Scope.Chain+" ", fmt.Sprintf("-I %s %d ", snapshot.Scope.Chain, *rulePlan.Expected.Locator.Position), 1) } + rulePlan.Expected.Locator.Position = nil script.WriteString(line) script.WriteByte('\n') + rulePlan.Commands, rulePlan.RollbackCommands = nil, nil + plan.Rules = append(plan.Rules, rulePlan) } script.WriteString("COMMIT\n") - return script.String(), nil + plan.Rules[0].Commands = []filter.NativeCommand{{ + Executable: restoreExecutableForFamily(snapshot.Scope.Family), Args: []string{"--noflush", "--wait"}, Stdin: script.String(), + }} + return plan, nil } func restoreRuleLine(scope filter.Scope, observed filter.ObservedRule) (string, error) { @@ -317,56 +290,7 @@ func restoreRuleLine(scope filter.Scope, observed filter.ObservedRule) (string, return strings.Join(tokens, " "), nil } -func (a *Adapter) Verify(ctx context.Context, plan filter.BackendPlan) (filter.VerifyResult, error) { - if plan.Provider != filter.ProviderIptables { - return filter.VerifyResult{}, fmt.Errorf("%w: backend plan provider %q", filter.ErrUnsupportedScope, plan.Provider) - } - snapshot, err := a.Observe(ctx, plan.Scope) - if err != nil { - return filter.VerifyResult{}, err - } - byMarker := make(map[string][]int, len(snapshot.Rules)) - for index, observed := range snapshot.Rules { - byMarker[observed.Marker] = append(byMarker[observed.Marker], index) - } - for _, expected := range plan.Rules { - markerMatches := 0 - semanticMatches := 0 - for _, index := range byMarker[expected.Expected.Marker] { - observed := snapshot.Rules[index] - if observed.Marker != "" && observed.Marker == expected.Expected.Marker { - markerMatches++ - want, wantErr := filter.RuleKey(expected.Expected.Rule) - got, gotErr := filter.RuleKey(observed.Rule) - if wantErr == nil && gotErr == nil && want == got { - semanticMatches++ - } - } - } - requiresPositionMatch := expected.Operation == filter.ChangeReorder || - (expected.Operation == filter.ChangeUpdate && expected.Expected.Rule.OrderIndex != nil) - positionMatches := true - if requiresPositionMatch { - positionMatches = false - for _, index := range byMarker[expected.Expected.Marker] { - observed := snapshot.Rules[index] - if observed.Marker == expected.Expected.Marker && observed.Locator.Position != nil && - expected.Expected.Locator.Position != nil && *observed.Locator.Position == *expected.Expected.Locator.Position { - positionMatches = true - break - } - } - } - if (expected.Operation == filter.ChangeDelete && markerMatches != 0) || - (requiresPositionMatch && !positionMatches) || - (expected.Operation != filter.ChangeDelete && (markerMatches != 1 || semanticMatches != 1)) { - return filter.VerifyResult{Snapshot: snapshot, Matched: false}, nil - } - } - return filter.VerifyResult{Snapshot: snapshot, Matched: true}, nil -} - -func (a *Adapter) Rollback(ctx context.Context, plan filter.BackendPlan) error { +func (a *Adapter) Rollback(ctx context.Context, plan filter.CommandBatch) error { if a.writer == nil { return fmt.Errorf("iptables writer is required") } @@ -382,7 +306,7 @@ func (a *Adapter) Rollback(ctx context.Context, plan filter.BackendPlan) error { return a.rollback(ctx, plan, len(plan.Rules)-1, -1) } -func (a *Adapter) compensate(ctx context.Context, plan filter.BackendPlan, lastRule, lastCommandCount int, cause error) error { +func (a *Adapter) compensate(ctx context.Context, plan filter.CommandBatch, lastRule, lastCommandCount int, cause error) error { if plan.CreatesOnly() { return cause } @@ -395,7 +319,7 @@ func (a *Adapter) compensate(ctx context.Context, plan filter.BackendPlan, lastR return cause } -func (a *Adapter) rollback(ctx context.Context, plan filter.BackendPlan, lastRule, lastCommandCount int) error { +func (a *Adapter) rollback(ctx context.Context, plan filter.CommandBatch, lastRule, lastCommandCount int) error { var rollbackErr error for index := lastRule; index >= 0; index-- { commands := plan.Rules[index].RollbackCommands @@ -432,53 +356,65 @@ func validateAdapterScope(scope filter.Scope) error { return nil } -func compileChange(snapshot filter.Snapshot, change filter.DesiredChange) (filter.NativeRulePlan, error) { +func compileChange(snapshot filter.RuleSet, change filter.RuleChange) (filter.RuleCommands, error) { rule := change.After if change.Operation == filter.ChangeDelete { rule = change.Before } if rule == nil { - return filter.NativeRulePlan{}, fmt.Errorf("%w: %s rule is required", filter.ErrInvalidRule, change.Operation) + return filter.RuleCommands{}, fmt.Errorf("%w: %s rule is required", filter.ErrInvalidRule, change.Operation) } normalized, err := filter.NormalizeRule(*rule) if err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } if normalized.Scope.Key() != snapshot.Scope.Key() { - return filter.NativeRulePlan{}, fmt.Errorf("%w: change scope %s", filter.ErrUnsupportedScope, normalized.Scope.Key()) + return filter.RuleCommands{}, fmt.Errorf("%w: change scope %s", filter.ErrUnsupportedScope, normalized.Scope.Key()) } if normalized.UUID == "" { - return filter.NativeRulePlan{}, fmt.Errorf("%w: rule UUID is required", filter.ErrInvalidRule) + return filter.RuleCommands{}, fmt.Errorf("%w: rule UUID is required", filter.ErrInvalidRule) } if (normalized.Scope.Family == filter.FamilyIPv4 && normalized.Protocol == "icmpv6") || (normalized.Scope.Family == filter.FamilyIPv6 && normalized.Protocol == "icmp") { - return filter.NativeRulePlan{}, fmt.Errorf("%w: protocol %q does not match %s", filter.ErrInvalidRule, normalized.Protocol, normalized.Scope.Family) + return filter.RuleCommands{}, fmt.Errorf("%w: protocol %q does not match %s", filter.ErrInvalidRule, normalized.Protocol, normalized.Scope.Family) } marker := "1panel-rule:" + normalized.UUID + if change.Operation == filter.ChangeDelete && change.CommandOnly && change.Locator == nil { + previous := filter.ObservedRule{Rule: normalized, Marker: marker, ParseStatus: filter.ParseStatusSupported} + var commands []filter.NativeCommand + if change.UnmarkedAdopted || change.PreviousMarker != "" { + previous.Marker = change.PreviousMarker + args := []string{"-w", "-t", snapshot.Scope.Table, "-D", snapshot.Scope.Chain} + args = append(args, compileObservedRuleArgs(previous)...) + commands = []filter.NativeCommand{{Executable: executableForFamily(snapshot.Scope.Family), Args: args}} + } + return filter.RuleCommands{RuleUUID: normalized.UUID, Operation: change.Operation, Previous: &previous, Expected: previous, Commands: commands}, nil + } position := len(snapshot.Rules) + 1 verb := "-I" var target filter.ObservedRule switch change.Operation { case filter.ChangeCreate: - if normalized.OrderIndex != nil && (*normalized.OrderIndex < 1 || *normalized.OrderIndex > int64(len(snapshot.Rules)+1)) { - return filter.NativeRulePlan{}, fmt.Errorf("%w: create target is out of range", filter.ErrInvalidRule) + hasSnapshot := !change.CommandOnly || snapshot.Rules != nil + if normalized.OrderIndex != nil && (*normalized.OrderIndex < 1 || hasSnapshot && *normalized.OrderIndex > int64(len(snapshot.Rules)+1)) { + return filter.RuleCommands{}, fmt.Errorf("%w: create target is out of range", filter.ErrInvalidRule) } position = insertionPosition(snapshot, normalized) case filter.ChangeAdopt: position, target, err = validateMutationTarget(snapshot, change, normalized, marker) if err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } verb = "-R" case filter.ChangeUpdate: position, target, err = validateMutationTarget(snapshot, change, normalized, marker) if err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } targetPosition := position if normalized.OrderIndex != nil { if *normalized.OrderIndex < 1 || *normalized.OrderIndex > int64(len(snapshot.Rules)) { - return filter.NativeRulePlan{}, fmt.Errorf("%w: update target is out of range", filter.ErrInvalidRule) + return filter.RuleCommands{}, fmt.Errorf("%w: update target is out of range", filter.ErrInvalidRule) } targetPosition = int(*normalized.OrderIndex) } @@ -489,21 +425,21 @@ func compileChange(snapshot filter.Snapshot, change filter.DesiredChange) (filte case filter.ChangeDelete: position, target, err = validateMutationTarget(snapshot, change, normalized, marker) if err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } verb = "-D" case filter.ChangeReorder: position, target, err = validateMutationTarget(snapshot, change, normalized, marker) if err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } if normalized.OrderIndex == nil || *normalized.OrderIndex < 1 || *normalized.OrderIndex > int64(len(snapshot.Rules)) { - return filter.NativeRulePlan{}, fmt.Errorf("%w: reorder target is out of range", filter.ErrInvalidRule) + return filter.RuleCommands{}, fmt.Errorf("%w: reorder target is out of range", filter.ErrInvalidRule) } targetPosition := int(*normalized.OrderIndex) return positionalMutationPlan(snapshot, normalized, target, marker, position, targetPosition, change.Operation), nil default: - return filter.NativeRulePlan{}, fmt.Errorf("%w: unsupported operation %s", filter.ErrInvalidRule, change.Operation) + return filter.RuleCommands{}, fmt.Errorf("%w: unsupported operation %s", filter.ErrInvalidRule, change.Operation) } args := []string{"-w", "-t", snapshot.Scope.Table, verb, snapshot.Scope.Chain, strconv.Itoa(position)} if change.Operation != filter.ChangeDelete { @@ -525,36 +461,32 @@ func compileChange(snapshot filter.Snapshot, change filter.DesiredChange) (filte Rule: normalized, Marker: marker, ParseStatus: filter.ParseStatusSupported, Locator: filter.Locator{Provider: filter.ProviderIptables, ScopeKey: snapshot.Scope.Key(), Position: &position}, } - return filter.NativeRulePlan{ + var previous *filter.ObservedRule + if change.Operation != filter.ChangeCreate { + previous = &target + } + return filter.RuleCommands{ RuleUUID: normalized.UUID, Operation: change.Operation, Commands: []filter.NativeCommand{{Executable: executableForFamily(snapshot.Scope.Family), Args: args}}, RollbackCommands: []filter.NativeCommand{{Executable: executableForFamily(snapshot.Scope.Family), Args: rollbackArgs}}, - Previous: pointerToObserved(target, change.Operation != filter.ChangeCreate), + Previous: previous, Expected: expected, }, nil } -func positionalMutationPlan( - snapshot filter.Snapshot, - rule filter.FirewallRule, - previous filter.ObservedRule, - marker string, - position int, - targetPosition int, - operation filter.ChangeOperation, -) filter.NativeRulePlan { +func positionalMutationPlan(snapshot filter.RuleSet, rule filter.FirewallRule, previous filter.ObservedRule, marker string, position int, targetPosition int, operation filter.ChangeOperation) filter.RuleCommands { expected := filter.ObservedRule{ Rule: rule, Marker: marker, ParseStatus: filter.ParseStatusSupported, Locator: filter.Locator{Provider: filter.ProviderIptables, ScopeKey: snapshot.Scope.Key(), Position: &targetPosition}, } - plan := filter.NativeRulePlan{ + plan := filter.RuleCommands{ RuleUUID: rule.UUID, Operation: operation, Previous: &previous, Expected: expected, } if position == targetPosition { return plan } executable := executableForFamily(snapshot.Scope.Family) - deleteArgs := []string{"-w", "-t", snapshot.Scope.Table, "-D", snapshot.Scope.Chain, strconv.Itoa(position)} + deleteArgs := append([]string{"-w", "-t", snapshot.Scope.Table, "-D", snapshot.Scope.Chain}, compileObservedRuleArgs(previous)...) insertArgs := []string{"-w", "-t", snapshot.Scope.Table, "-I", snapshot.Scope.Chain, strconv.Itoa(targetPosition)} insertArgs = append(insertArgs, compileRuleArgs(rule, marker)...) restoreArgs := []string{"-w", "-t", snapshot.Scope.Table, "-I", snapshot.Scope.Chain, strconv.Itoa(position)} @@ -572,13 +504,6 @@ func positionalMutationPlan( return plan } -func pointerToObserved(rule filter.ObservedRule, include bool) *filter.ObservedRule { - if !include { - return nil - } - return &rule -} - func compileRuleArgs(rule filter.FirewallRule, marker string) []string { args := make([]string, 0, 24) if rule.Protocol != "all" { @@ -599,16 +524,16 @@ func compileRuleArgs(rule filter.FirewallRule, marker string) []string { } if rule.SourcePort != "" { if strings.Contains(rule.SourcePort, ",") { - args = append(args, "-m", "multiport", "--sports", nativePort(rule.SourcePort)) + args = append(args, "-m", "multiport", "--sports", strings.ReplaceAll(rule.SourcePort, "-", ":")) } else { - args = append(args, "--sport", nativePort(rule.SourcePort)) + args = append(args, "--sport", strings.ReplaceAll(rule.SourcePort, "-", ":")) } } if rule.DestinationPort != "" { if strings.Contains(rule.DestinationPort, ",") { - args = append(args, "-m", "multiport", "--dports", nativePort(rule.DestinationPort)) + args = append(args, "-m", "multiport", "--dports", strings.ReplaceAll(rule.DestinationPort, "-", ":")) } else { - args = append(args, "--dport", nativePort(rule.DestinationPort)) + args = append(args, "--dport", strings.ReplaceAll(rule.DestinationPort, "-", ":")) } } if len(rule.ConnectionStates) != 0 { @@ -629,11 +554,7 @@ func compileObservedRuleArgs(observed filter.ObservedRule) []string { return compileRuleArgs(observed.Rule, comment) } -func nativePort(port string) string { - return strings.ReplaceAll(port, "-", ":") -} - -func validateMutationTarget(snapshot filter.Snapshot, change filter.DesiredChange, after filter.FirewallRule, marker string) (int, filter.ObservedRule, error) { +func validateMutationTarget(snapshot filter.RuleSet, change filter.RuleChange, after filter.FirewallRule, marker string) (int, filter.ObservedRule, error) { if change.Locator == nil || change.Locator.Position == nil { return 0, filter.ObservedRule{}, fmt.Errorf("%w: mutation requires a position locator", filter.ErrInvalidRule) } @@ -673,7 +594,7 @@ func validateMutationTarget(snapshot filter.Snapshot, change filter.DesiredChang return position, observed, nil } -func insertionPosition(snapshot filter.Snapshot, rule filter.FirewallRule) int { +func insertionPosition(snapshot filter.RuleSet, rule filter.FirewallRule) int { if rule.OrderIndex != nil { return int(*rule.OrderIndex) } @@ -691,17 +612,27 @@ func insertionPosition(snapshot filter.Snapshot, rule filter.FirewallRule) int { type systemBackend struct{} -func (systemBackend) CheckMultiport(ctx context.Context, family filter.Family) error { - executable, err := runtimeExecutableForFamily(family) +func (systemBackend) ReadRulesByComment(ctx context.Context, scope filter.Scope, comment string) (string, error) { + executable, err := runtimeExecutable(executableForFamily(scope.Family)) if err != nil { - return err + return "", err } - return cmd.NewCommandMgr(cmd.WithContext(ctx), cmd.WithTimeout(20*time.Second)).RunWithOptionalSudo(executable, "-m", "multiport", "--help") + return filter.ReadRulesByComment(ctx, executable, []string{"-w", "-t", scope.Table, "-S", scope.Chain}, comment) } -func (systemBackend) ListChain(ctx context.Context, scope filter.Scope) (string, error) { - output, err := (systemBackend{}).ListTable(ctx, scope) - return chainOutput(scope, output, err) +func (systemBackend) ListTable(ctx context.Context, scope filter.Scope) (string, error) { + return native.ReadTable(ctx, scope.Table, scope.Family == filter.FamilyIPv6) +} + +func (b systemBackend) ListChain(ctx context.Context, scope filter.Scope) (string, error) { + output, err := b.ListTable(ctx, scope) + if err != nil { + return "", err + } + if !containsChainDeclaration(output, scope.Chain) { + return "", fmt.Errorf("%w: iptables %s chain %s is not initialized", filter.ErrProviderUnavailable, scope.Family, scope.Chain) + } + return output, nil } func containsChainDeclaration(output, chain string) bool { @@ -760,10 +691,6 @@ func restoreExecutableForFamily(family filter.Family) string { return "iptables-restore" } -func runtimeExecutableForFamily(family filter.Family) (string, error) { - return runtimeExecutable(executableForFamily(family)) -} - func runtimeExecutable(logical string) (string, error) { commands, err := lifecycle.ResolveIptablesCommands() if err != nil { @@ -937,54 +864,6 @@ func takeValue(args []string, index *int, target *string) bool { return true } -type tableReader interface { - ListTable(context.Context, filter.Scope) (string, error) -} - -type tableRead struct { - output string - err error -} - -type tableObservationReader struct { - reader tableReader - tables map[string]tableRead -} - -func (r *tableObservationReader) ListChain(ctx context.Context, scope filter.Scope) (string, error) { - if err := ctx.Err(); err != nil { - return "", err - } - key := string(scope.Family) + ":" + scope.Table - read, exists := r.tables[key] - if !exists { - read.output, read.err = r.reader.ListTable(ctx, scope) - r.tables[key] = read - } - return chainOutput(scope, read.output, read.err) -} - -func (a *Adapter) NewObservationSession() filter.Adapter { - reader, ok := a.reader.(tableReader) - if !ok { - return a - } - return &Adapter{ - reader: &tableObservationReader{reader: reader, tables: make(map[string]tableRead)}, - writer: a.writer, checker: a.checker, - } -} - -func (systemBackend) ListTable(ctx context.Context, scope filter.Scope) (string, error) { - return native.ReadTable(ctx, scope.Table, scope.Family == filter.FamilyIPv6) -} - -func chainOutput(scope filter.Scope, output string, err error) (string, error) { - if err != nil { - return "", err - } - if !containsChainDeclaration(output, scope.Chain) { - return "", fmt.Errorf("%w: iptables %s chain %s is not initialized", filter.ErrProviderUnavailable, scope.Family, scope.Chain) - } - return output, nil +func (a *Adapter) SaveRules(ctx context.Context, scope filter.Scope) error { + return a.writer.Save(ctx, scope) } diff --git a/agent/utils/firewall/filter/providers/nftables/adapter.go b/agent/utils/firewall/filter/providers/nftables/adapter.go index 0fbd525302d4..1a1f1ed61b96 100644 --- a/agent/utils/firewall/filter/providers/nftables/adapter.go +++ b/agent/utils/firewall/filter/providers/nftables/adapter.go @@ -15,11 +15,16 @@ import ( ) type Backend interface { + filter.CommentRuleReader ListChain(context.Context, filter.Scope) (string, error) Run(context.Context, filter.NativeCommand) error Save(context.Context) error } +type TableReader interface { + ListTable(context.Context, filter.Scope) (string, bool, error) +} + type Adapter struct{ backend Backend } func NewAdapter() *Adapter { return &Adapter{backend: systemBackend{}} } @@ -34,19 +39,47 @@ func (a *Adapter) Capabilities(context.Context) (filter.Capabilities, error) { }, nil } -func (a *Adapter) Observe(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) { +func (a *Adapter) AppendUnverified(ctx context.Context, rule filter.FirewallRule, comment string) error { + rule, err := filter.NormalizeRule(rule) + if err != nil { + return err + } + if err := validateScope(rule.Scope); err != nil { + return err + } + script := fmt.Sprintf("add rule %s %s %s %s\n", nftables_helper.TableFamily(rule.Scope.Family), nftables_helper.TableName, nativeChainName(rule.Scope), strings.Join(compileExpressionArgs(rule, comment), " ")) + return a.backend.Run(ctx, filter.NativeCommand{Executable: "nft", Stdin: script}) +} + +func (a *Adapter) ListRulesByComment(ctx context.Context, scopes []filter.Scope, comment string) ([]filter.ObservedRule, error) { + var rules []filter.ObservedRule + for _, scope := range scopes { + scope = scope.Normalize() + if err := validateScope(scope); err != nil { + return nil, err + } + output, err := a.backend.ReadRulesByComment(ctx, scope, comment) + if err != nil { + return nil, err + } + rules = append(rules, parseChain(scope, output)...) + } + return rules, nil +} + +func (a *Adapter) ListRules(ctx context.Context, scope filter.Scope) (filter.RuleSet, error) { scope = scope.Normalize() if err := validateScope(scope); err != nil { - return filter.Snapshot{}, err + return filter.RuleSet{}, err } if a.backend == nil { - return filter.Snapshot{}, fmt.Errorf("nftables backend is required") + return filter.RuleSet{}, fmt.Errorf("nftables backend is required") } output, err := a.backend.ListChain(ctx, scope) if errors.Is(err, nftables_helper.ErrChainNotFound) { - snapshot, snapshotErr := filter.NewSnapshot(scope, nil) + snapshot, snapshotErr := filter.NewRuleSet(scope, nil) if snapshotErr != nil { - return filter.Snapshot{}, snapshotErr + return filter.RuleSet{}, snapshotErr } snapshot.Notices = []filter.ScopeNotice{{ Code: filter.ScopeNoticeManagedScopeMissing, Values: []string{string(scope.Family), scope.Chain}, @@ -54,70 +87,189 @@ func (a *Adapter) Observe(ctx context.Context, scope filter.Scope) (filter.Snaps return snapshot, nil } if err != nil { - return filter.Snapshot{}, err + return filter.RuleSet{}, err } - return filter.NewSnapshot(scope, parseChain(scope, output)) + return filter.NewRuleSet(scope, parseChain(scope, output)) } -func (a *Adapter) Compile(snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, error) { - if snapshot.Revision == "" { - return filter.BackendPlan{}, filter.ErrRuleStale +func (a *Adapter) ListRuleScopes(ctx context.Context, scopes []filter.Scope) ([]filter.RuleSet, error) { + normalized := make([]filter.Scope, len(scopes)) + for index, scope := range scopes { + scope = scope.Normalize() + if err := validateScope(scope); err != nil { + return nil, err + } + normalized[index] = scope + } + if a.backend == nil { + return nil, fmt.Errorf("nftables backend is required") + } + snapshots := make([]filter.RuleSet, len(scopes)) + reader, readsTable := a.backend.(TableReader) + for index, scope := range normalized { + if snapshots[index].Scope.Provider != "" { + continue + } + if !readsTable { + snapshot, err := a.ListRules(ctx, scope) + if err != nil { + return nil, err + } + snapshots[index] = snapshot + continue + } + output, _, err := reader.ListTable(ctx, scope) + if err != nil { + return nil, err + } + chains := nftables_helper.ParseTableChains(output) + for target := index; target < len(normalized); target++ { + current := normalized[target] + if current.Family != scope.Family || current.Table != scope.Table { + continue + } + chain, exists := chains[nativeChainName(current)] + snapshot, err := filter.NewRuleSet(current, parseChain(current, chain)) + if err != nil { + return nil, err + } + if !exists { + snapshot.Notices = []filter.ScopeNotice{{Code: filter.ScopeNoticeManagedScopeMissing, Values: []string{string(current.Family), current.Chain}}} + } + snapshots[target] = snapshot + } } + return snapshots, nil +} + +func (a *Adapter) BuildCommands(snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, error) { if err := validateScope(snapshot.Scope); err != nil { - return filter.BackendPlan{}, err + return filter.CommandBatch{}, err } if len(changes) == 0 { - return filter.BackendPlan{}, fmt.Errorf("%w: nftables plan requires at least one change", filter.ErrInvalidRule) + return filter.CommandBatch{}, fmt.Errorf("%w: nftables plan requires at least one change", filter.ErrInvalidRule) + } + createOnly, deleteOnly := true, true + for _, change := range changes { + createOnly = createOnly && change.Operation == filter.ChangeCreate && change.CommandOnly + deleteOnly = deleteOnly && change.Operation == filter.ChangeDelete && change.CommandOnly + } + if createOnly { + return compileCreateBatch(snapshot, changes) + } + if deleteOnly { + return compileDeleteBatch(snapshot, changes) + } + if len(changes) != 1 { + return filter.CommandBatch{}, fmt.Errorf("%w: nftables mutation requires exactly one change", filter.ErrInvalidRule) + } + change := changes[0] + if change.Operation != filter.ChangeUpdate && change.Operation != filter.ChangeAdopt && change.Operation != filter.ChangeReorder { + return filter.CommandBatch{}, fmt.Errorf("%w: unsupported nftables mutation %s", filter.ErrInvalidRule, change.Operation) + } + expected, previous, err := compileChange(snapshot, change) + if err != nil { + return filter.CommandBatch{}, err + } + handle := previous.Locator.NativeID + if _, err := strconv.ParseUint(handle, 10, 64); err != nil || handle != change.Locator.NativeID { + return filter.CommandBatch{}, filter.ErrRuleStale + } + chain := strings.Join([]string{nftables_helper.TableFamily(snapshot.Scope.Family), nftables_helper.TableName, nativeChainName(snapshot.Scope)}, " ") + rulePlan := filter.RuleCommands{RuleUUID: ruleUUID(change), Operation: change.Operation, Previous: previous, Expected: expected} + target := *expected.Locator.Position + var script string + if target == *previous.Locator.Position { + if change.Operation != filter.ChangeReorder { + script = fmt.Sprintf("replace rule %s handle %s %s\n", chain, handle, expected.Raw) + if strings.ContainsAny(previous.Raw, "\r\n") || previous.Raw == "" { + return filter.CommandBatch{}, fmt.Errorf("%w: invalid native nftables rule", filter.ErrInvalidRule) + } + rulePlan.RollbackCommands = []filter.NativeCommand{{Executable: "nft", Stdin: fmt.Sprintf("replace rule %s handle %s %s\n", chain, handle, previous.Raw)}} + } + } else { + script = fmt.Sprintf("delete rule %s handle %s\n", chain, handle) + if target == len(snapshot.Rules) { + script += fmt.Sprintf("add rule %s %s\n", chain, expected.Raw) + } else { + anchorIndex := target - 1 + if target > *previous.Locator.Position { + anchorIndex++ + } + anchor := snapshot.Rules[anchorIndex].Locator.NativeID + if _, err := strconv.ParseUint(anchor, 10, 64); err != nil { + return filter.CommandBatch{}, filter.ErrRuleStale + } + script += fmt.Sprintf("insert rule %s position %s %s\n", chain, anchor, expected.Raw) + } } - if len(changes) > 1 && changes[0].Operation != filter.ChangeCreate && changes[0].Operation != filter.ChangeDelete { - return filter.BackendPlan{}, fmt.Errorf("%w: nftables batch plans only support create or delete operations", filter.ErrInvalidRule) + if script != "" { + rulePlan.Commands = []filter.NativeCommand{{Executable: "nft", Stdin: script}} } + return filter.CommandBatch{Provider: filter.ProviderNftables, Scope: snapshot.Scope, CommandOnly: change.CommandOnly, Rules: []filter.RuleCommands{rulePlan}}, nil +} - operation := changes[0].Operation - current := snapshot - current.Rules = make([]filter.ObservedRule, len(snapshot.Rules), len(snapshot.Rules)+len(changes)) - copy(current.Rules, snapshot.Rules) - plan := filter.BackendPlan{ - Provider: filter.ProviderNftables, Scope: snapshot.Scope, SnapshotRevision: snapshot.Revision, - Rules: make([]filter.NativeRulePlan, 0, len(changes)), +func compileDeleteBatch(snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, error) { + plan := filter.CommandBatch{Provider: filter.ProviderNftables, Scope: snapshot.Scope, CommandOnly: true} + var script strings.Builder + for _, change := range changes { + expected, previous, err := compileChange(snapshot, change) + if err != nil { + return filter.CommandBatch{}, err + } + handle := previous.Locator.NativeID + if _, err := strconv.ParseUint(handle, 10, 64); err != nil || change.Locator.NativeID != handle { + return filter.CommandBatch{}, filter.ErrRuleStale + } + fmt.Fprintf(&script, "delete rule %s %s %s handle %s\n", nftables_helper.TableFamily(snapshot.Scope.Family), nftables_helper.TableName, nativeChainName(snapshot.Scope), handle) + plan.Rules = append(plan.Rules, filter.RuleCommands{ + RuleUUID: ruleUUID(change), Operation: filter.ChangeDelete, Previous: previous, Expected: expected, + }) } + plan.Rules[0].Commands = []filter.NativeCommand{{Executable: "nft", Stdin: script.String()}} + return plan, nil +} + +func compileCreateBatch(snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, error) { + plan := filter.CommandBatch{Provider: filter.ProviderNftables, Scope: snapshot.Scope, CommandOnly: true} + var script strings.Builder for _, change := range changes { - if len(changes) > 1 && change.Operation != operation { - return filter.BackendPlan{}, fmt.Errorf("%w: nftables batch plan operations must be homogeneous", filter.ErrInvalidRule) + if change.After == nil { + return filter.CommandBatch{}, fmt.Errorf("%w: create rule is required", filter.ErrInvalidRule) } - rules, expected, previous, err := applyChange(current, change) + rule, err := filter.NormalizeRule(*change.After) if err != nil { - return filter.BackendPlan{}, err + return filter.CommandBatch{}, err } - plan.Rules = append(plan.Rules, filter.NativeRulePlan{ - RuleUUID: ruleUUID(change), Operation: change.Operation, Previous: previous, Expected: expected, + if rule.Scope.Key() != snapshot.Scope.Key() || rule.UUID == "" { + return filter.CommandBatch{}, fmt.Errorf("%w: invalid nftables creation rule", filter.ErrInvalidRule) + } + marker := "1panel-rule:" + rule.UUID + verb := "add" + if !change.Append { + if rule.OrderIndex == nil || *rule.OrderIndex != 1 { + return filter.CommandBatch{}, fmt.Errorf("%w: batch insertion requires the first position", filter.ErrInvalidRule) + } + verb = "insert" + } + fmt.Fprintf(&script, "%s rule %s %s %s %s\n", verb, nftables_helper.TableFamily(rule.Scope.Family), nftables_helper.TableName, nativeChainName(rule.Scope), strings.Join(compileExpressionArgs(rule, marker), " ")) + plan.Rules = append(plan.Rules, filter.RuleCommands{ + RuleUUID: rule.UUID, Operation: filter.ChangeCreate, + Expected: filter.ObservedRule{Rule: rule, Marker: marker, ParseStatus: filter.ParseStatusSupported}, }) - current.Rules = rules - } - if _, err := filter.NewSnapshot(snapshot.Scope, current.Rules); err != nil { - return filter.BackendPlan{}, err - } - applyCommand, err := rebuildCommand(snapshot.Scope, current.Rules) - if err != nil { - return filter.BackendPlan{}, err - } - rollbackCommand, err := rebuildCommand(snapshot.Scope, snapshot.Rules) - if err != nil { - return filter.BackendPlan{}, err } - plan.Rules[0].Commands = []filter.NativeCommand{applyCommand} - plan.Rules[0].RollbackCommands = []filter.NativeCommand{rollbackCommand} + plan.Rules[0].Commands = []filter.NativeCommand{{Executable: "nft", Stdin: script.String()}} return plan, nil } -func (a *Adapter) Apply(ctx context.Context, plan filter.BackendPlan) (filter.ApplyResult, error) { +func (a *Adapter) RunCommands(ctx context.Context, plan filter.CommandBatch) error { if err := validatePlan(plan); err != nil { - return filter.ApplyResult{}, err + return err } for _, rulePlan := range plan.Rules { for _, command := range rulePlan.Commands { if err := validateNativeCommand(command); err != nil { - return filter.ApplyResult{}, err + return err } } } @@ -125,62 +277,17 @@ func (a *Adapter) Apply(ctx context.Context, plan filter.BackendPlan) (filter.Ap for _, command := range rulePlan.Commands { if err := a.backend.Run(ctx, command); err != nil { if plan.CommandOnly { - return filter.ApplyResult{}, err + return err } - return filter.ApplyResult{}, a.compensate(ctx, plan, err) + return a.compensate(ctx, plan, err) } } } - if err := a.backend.Save(ctx); err != nil { - if plan.CommandOnly { - return filter.ApplyResult{}, err - } - return filter.ApplyResult{}, a.compensate(ctx, plan, err) - } - applied := make([]filter.ObservedRule, 0, len(plan.Rules)) - for _, rulePlan := range plan.Rules { - applied = append(applied, rulePlan.Expected) - } - return filter.ApplyResult{Applied: applied}, nil -} -func (a *Adapter) Verify(ctx context.Context, plan filter.BackendPlan) (filter.VerifyResult, error) { - if err := validatePlan(plan); err != nil { - return filter.VerifyResult{}, err - } - snapshot, err := a.Observe(ctx, plan.Scope) - if err != nil { - return filter.VerifyResult{}, err - } - byMarker := make(map[string][]int, len(snapshot.Rules)) - for index, observed := range snapshot.Rules { - byMarker[observed.Marker] = append(byMarker[observed.Marker], index) - } - for _, expected := range plan.Rules { - matches := 0 - for _, index := range byMarker[expected.Expected.Marker] { - observed := snapshot.Rules[index] - if observed.Marker == expected.Expected.Marker { - if expected.Operation == filter.ChangeDelete { - matches++ - continue - } - want, wantErr := filter.RuleKey(expected.Expected.Rule) - got, gotErr := filter.RuleKey(observed.Rule) - if wantErr == nil && gotErr == nil && want == got { - matches++ - } - } - } - if (expected.Operation == filter.ChangeDelete && matches != 0) || - (expected.Operation != filter.ChangeDelete && matches != 1) { - return filter.VerifyResult{Snapshot: snapshot, Matched: false}, nil - } - } - return filter.VerifyResult{Snapshot: snapshot, Matched: true}, nil + return nil } -func (a *Adapter) Rollback(ctx context.Context, plan filter.BackendPlan) error { +func (a *Adapter) Rollback(ctx context.Context, plan filter.CommandBatch) error { if err := validatePlan(plan); err != nil { return err } @@ -201,7 +308,7 @@ func (a *Adapter) Rollback(ctx context.Context, plan filter.BackendPlan) error { return a.backend.Save(ctx) } -func (a *Adapter) compensate(ctx context.Context, plan filter.BackendPlan, cause error) error { +func (a *Adapter) compensate(ctx context.Context, plan filter.CommandBatch, cause error) error { if plan.CreatesOnly() { return cause } @@ -237,7 +344,7 @@ func nativeChainName(scope filter.Scope) string { } } -func validatePlan(plan filter.BackendPlan) error { +func validatePlan(plan filter.CommandBatch) error { if plan.Provider != filter.ProviderNftables || len(plan.Rules) == 0 { return fmt.Errorf("%w: invalid nftables plan", filter.ErrInvalidRule) } @@ -251,70 +358,42 @@ func validateNativeCommand(command filter.NativeCommand) error { return nil } -func applyChange(snapshot filter.Snapshot, change filter.DesiredChange) ([]filter.ObservedRule, filter.ObservedRule, *filter.ObservedRule, error) { - rules := snapshot.Rules +func compileChange(snapshot filter.RuleSet, change filter.RuleChange) (filter.ObservedRule, *filter.ObservedRule, error) { rule := change.After if change.Operation == filter.ChangeDelete { rule = change.Before } if rule == nil { - return nil, filter.ObservedRule{}, nil, fmt.Errorf("%w: %s rule is required", filter.ErrInvalidRule, change.Operation) + return filter.ObservedRule{}, nil, fmt.Errorf("%w: %s rule is required", filter.ErrInvalidRule, change.Operation) } normalized, err := filter.NormalizeRule(*rule) if err != nil { - return nil, filter.ObservedRule{}, nil, err + return filter.ObservedRule{}, nil, err } if normalized.Scope.Key() != snapshot.Scope.Key() || normalized.UUID == "" { - return nil, filter.ObservedRule{}, nil, fmt.Errorf("%w: invalid nftables mutation rule", filter.ErrInvalidRule) + return filter.ObservedRule{}, nil, fmt.Errorf("%w: invalid nftables mutation rule", filter.ErrInvalidRule) } - position := len(rules) + 1 - marker := "1panel-rule:" + normalized.UUID - var previous *filter.ObservedRule - - if change.Operation != filter.ChangeCreate { - if change.Locator == nil || change.Locator.Position == nil { - return nil, filter.ObservedRule{}, nil, fmt.Errorf("%w: mutation requires a position locator", filter.ErrInvalidRule) - } - position = *change.Locator.Position - if position < 1 || position > len(rules) { - return nil, filter.ObservedRule{}, nil, filter.ErrRuleStale - } - selected := rules[position-1] - if selected.Protected { - return nil, filter.ObservedRule{}, nil, filter.ErrProtectedRule - } - previousCopy := selected - previous = &previousCopy - rules = append(rules[:position-1], rules[position:]...) + if change.Locator == nil || change.Locator.Position == nil { + return filter.ObservedRule{}, nil, fmt.Errorf("%w: mutation requires a position locator", filter.ErrInvalidRule) } - - target := position - if normalized.OrderIndex != nil { - target = int(*normalized.OrderIndex) + position := *change.Locator.Position + if position < 1 || position > len(snapshot.Rules) { + return filter.ObservedRule{}, nil, filter.ErrRuleStale } - if change.Operation == filter.ChangeCreate && normalized.OrderIndex == nil { - target = len(rules) + 1 + previous := snapshot.Rules[position-1] + if previous.Protected { + return filter.ObservedRule{}, nil, filter.ErrProtectedRule } - if change.Operation == filter.ChangeDelete { - expected := observedRule(normalized, marker, position, "") - for index := position - 1; index < len(rules); index++ { - position := index + 1 - rules[index].Locator.Position = &position - } - return rules, expected, previous, nil + target := position + if normalized.OrderIndex != nil && (change.Operation == filter.ChangeUpdate || change.Operation == filter.ChangeReorder) { + target = int(*normalized.OrderIndex) } - if target < 1 || target > len(rules)+1 { - return nil, filter.ObservedRule{}, nil, fmt.Errorf("%w: target position is out of range", filter.ErrInvalidRule) + if target < 1 || target > len(snapshot.Rules) { + return filter.ObservedRule{}, nil, fmt.Errorf("%w: target position is out of range", filter.ErrInvalidRule) } + marker := "1panel-rule:" + normalized.UUID expected := observedRule(normalized, marker, target, strings.Join(compileExpressionArgs(normalized, marker), " ")) - rules = append(rules, filter.ObservedRule{}) - copy(rules[target:], rules[target-1:]) - rules[target-1] = expected - for index := min(position, target) - 1; index < len(rules); index++ { - position := index + 1 - rules[index].Locator.Position = &position - } - return rules, expected, previous, nil + return expected, &previous, nil } func observedRule(rule filter.FirewallRule, marker string, position int, raw string) filter.ObservedRule { @@ -324,7 +403,7 @@ func observedRule(rule filter.FirewallRule, marker string, position int, raw str } } -func ruleUUID(change filter.DesiredChange) string { +func ruleUUID(change filter.RuleChange) string { if change.After != nil { return change.After.UUID } @@ -334,39 +413,6 @@ func ruleUUID(change filter.DesiredChange) string { return "" } -func rebuildCommand(scope filter.Scope, rules []filter.ObservedRule) (filter.NativeCommand, error) { - tableFamily := nftables_helper.TableFamily(scope.Family) - chain := nativeChainName(scope) - var script strings.Builder - script.WriteString("flush chain ") - script.WriteString(tableFamily) - script.WriteByte(' ') - script.WriteString(nftables_helper.TableName) - script.WriteByte(' ') - script.WriteString(chain) - script.WriteByte('\n') - for _, rule := range rules { - raw := strings.TrimSpace(rule.Raw) - if rule.ParseStatus == filter.ParseStatusSupported && rule.Marker != "" { - script.WriteString(strings.Join([]string{"add", "rule", tableFamily, nftables_helper.TableName, chain}, " ")) - script.WriteByte(' ') - script.WriteString(strings.Join(compileExpressionArgs(rule.Rule, rule.Marker), " ")) - script.WriteByte('\n') - continue - } - if raw != "" { - if strings.ContainsAny(raw, "\r\n") { - return filter.NativeCommand{}, fmt.Errorf("%w: invalid newline in native nftables rule", filter.ErrInvalidRule) - } - script.WriteString(strings.Join([]string{"add", "rule", tableFamily, nftables_helper.TableName, chain}, " ")) - script.WriteByte(' ') - script.WriteString(raw) - script.WriteByte('\n') - } - } - return filter.NativeCommand{Executable: "nft", Stdin: script.String()}, nil -} - func compileExpressionArgs(rule filter.FirewallRule, marker string) []string { parts := make([]string, 0, 24) if rule.Protocol != "all" { @@ -593,6 +639,10 @@ func numericSymbol(value string) string { type systemBackend struct{} +func (systemBackend) ReadRulesByComment(ctx context.Context, scope filter.Scope, comment string) (string, error) { + return filter.ReadRulesByComment(ctx, "nft", []string{"-a", "-n", "-n", "list", "chain", nftables_helper.TableFamily(scope.Family), nftables_helper.TableName, nativeChainName(scope)}, comment) +} + func (systemBackend) ListChain(ctx context.Context, scope filter.Scope) (string, error) { run := func(args ...string) (string, error) { return cmd.NewCommandMgr(cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second)).RunWithOptionalSudoAndStdout( @@ -602,6 +652,13 @@ func (systemBackend) ListChain(ctx context.Context, scope filter.Scope) (string, return nftables_helper.ReadChain(run, nftables_helper.TableFamily(scope.Family), nftables_helper.TableName, nativeChainName(scope)) } +func (systemBackend) ListTable(ctx context.Context, scope filter.Scope) (string, bool, error) { + run := func(args ...string) (string, error) { + return cmd.NewCommandMgr(cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second)).RunWithOptionalSudoAndStdout("nft", append([]string{"-n", "-n"}, args...)...) + } + return nftables_helper.ReadTable(run, nftables_helper.TableFamily(scope.Family), nftables_helper.TableName) +} + func (systemBackend) Run(ctx context.Context, command filter.NativeCommand) error { if command.Executable != "nft" { return fmt.Errorf("unexpected nftables executable %q", command.Executable) @@ -616,3 +673,7 @@ func (systemBackend) Run(ctx context.Context, command filter.NativeCommand) erro func (systemBackend) Save(ctx context.Context) error { return nftables_helper.PersistRuleset(ctx) } + +func (a *Adapter) SaveRules(ctx context.Context, scope filter.Scope) error { + return a.backend.Save(ctx) +} diff --git a/agent/utils/firewall/filter/providers/ufw/adapter.go b/agent/utils/firewall/filter/providers/ufw/adapter.go index 9c76c378eef8..f153e89e5d84 100644 --- a/agent/utils/firewall/filter/providers/ufw/adapter.go +++ b/agent/utils/firewall/filter/providers/ufw/adapter.go @@ -16,6 +16,7 @@ import ( ) type CommandReader interface { + filter.CommentRuleReader Read(context.Context, ...string) (string, error) } @@ -69,7 +70,7 @@ func (a *Adapter) AppendUnverified(ctx context.Context, rule filter.FirewallRule } comment = strings.TrimSpace(comment) if comment == "" || strings.ContainsAny(comment, "\r\n\x00") { - return fmt.Errorf("%w: invalid UFW fallback comment", filter.ErrInvalidRule) + return fmt.Errorf("%w: invalid UFW external rule comment", filter.ErrInvalidRule) } command := commentCommand(normalized, comment) if err := validateCommand(command); err != nil { @@ -85,15 +86,35 @@ func (a *Adapter) Capabilities(context.Context) (filter.Capabilities, error) { }, nil } -func (a *Adapter) Observe(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) { - snapshots, err := a.ObserveScopes(ctx, []filter.Scope{scope}) +func (a *Adapter) ListRulesByComment(ctx context.Context, scopes []filter.Scope, comment string) ([]filter.ObservedRule, error) { + var rules []filter.ObservedRule + var output string + for index, scope := range scopes { + scope = scope.Normalize() + if err := validateScope(scope); err != nil { + return nil, err + } + if index == 0 { + var err error + output, err = a.reader.ReadRulesByComment(ctx, scope, comment) + if err != nil { + return nil, err + } + } + rules = append(rules, parseNumberedRules(scope, output)...) + } + return rules, nil +} + +func (a *Adapter) ListRules(ctx context.Context, scope filter.Scope) (filter.RuleSet, error) { + snapshots, err := a.ListRuleScopes(ctx, []filter.Scope{scope}) if err != nil { - return filter.Snapshot{}, err + return filter.RuleSet{}, err } return snapshots[0], nil } -func (a *Adapter) ObserveScopes(ctx context.Context, scopes []filter.Scope) ([]filter.Snapshot, error) { +func (a *Adapter) ListRuleScopes(ctx context.Context, scopes []filter.Scope) ([]filter.RuleSet, error) { if len(scopes) == 0 { return nil, fmt.Errorf("%w: UFW observation requires at least one scope", filter.ErrInvalidScope) } @@ -113,14 +134,28 @@ func (a *Adapter) ObserveScopes(ctx context.Context, scopes []filter.Scope) ([]f return nil, fmt.Errorf("%w: read UFW numbered status: %w", filter.ErrInventoryUnavailable, err) } + lastPositions := make(map[filter.Family]int) + for _, line := range strings.Split(numbered, "\n") { + matches := re.UFWNumberedRuleRegex.FindStringSubmatch(strings.TrimSpace(line)) + if len(matches) != 0 { + position, err := strconv.Atoi(matches[1]) + if err == nil { + family := familyForNumberedRule(matches[2], matches[5]) + lastPositions[family] = max(lastPositions[family], position) + } + } else if observed, family, _, ok := parseUnrecognizedNumberedRule(normalizedScopes[0], strings.TrimSpace(line)); ok && observed.Locator.Position != nil { + lastPositions[family] = max(lastPositions[family], *observed.Locator.Position) + } + } notices := statusNotices(numbered) - snapshots := make([]filter.Snapshot, 0, len(normalizedScopes)) + snapshots := make([]filter.RuleSet, 0, len(normalizedScopes)) for _, scope := range normalizedScopes { - snapshot, err := filter.NewSnapshot(scope, parseNumberedRules(scope, numbered)) + snapshot, err := filter.NewRuleSet(scope, parseNumberedRules(scope, numbered)) if err != nil { return nil, err } snapshot.Notices = append([]filter.ScopeNotice(nil), notices...) + snapshot.LastPosition = lastPositions[scope.Family] snapshots = append(snapshots, snapshot) } return snapshots, nil @@ -145,46 +180,61 @@ func (a *Adapter) NativeDetail(ctx context.Context, profile string, _ bool) (str return info, nil } -func (a *Adapter) Compile(snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, error) { - if snapshot.Revision == "" { - return filter.BackendPlan{}, filter.ErrRuleStale - } +func (a *Adapter) BuildCommands(snapshot filter.RuleSet, changes []filter.RuleChange) (filter.CommandBatch, error) { if err := validateScope(snapshot.Scope); err != nil { - return filter.BackendPlan{}, err + return filter.CommandBatch{}, err } if hasScopeNotice(snapshot.Notices, filter.ScopeNoticeManagedScopeInactive) { - return filter.BackendPlan{}, fmt.Errorf("%w: ufw is inactive", filter.ErrProviderUnavailable) + return filter.CommandBatch{}, fmt.Errorf("%w: ufw is inactive", filter.ErrProviderUnavailable) } if len(changes) != 1 { - return filter.BackendPlan{}, fmt.Errorf("%w: ufw plans currently require exactly one change", filter.ErrInvalidRule) + return filter.CommandBatch{}, fmt.Errorf("%w: ufw plans currently require exactly one change", filter.ErrInvalidRule) } rulePlan, err := compileChange(snapshot, changes[0]) if err != nil { - return filter.BackendPlan{}, err + return filter.CommandBatch{}, err } - return filter.BackendPlan{ - Provider: filter.ProviderUFW, Scope: snapshot.Scope, SnapshotRevision: snapshot.Revision, - Rules: []filter.NativeRulePlan{rulePlan}, + return filter.CommandBatch{ + Provider: filter.ProviderUFW, Scope: snapshot.Scope, + Rules: []filter.RuleCommands{rulePlan}, }, nil } -func (a *Adapter) Apply(ctx context.Context, plan filter.BackendPlan) (filter.ApplyResult, error) { +func (a *Adapter) RunCommands(ctx context.Context, plan filter.CommandBatch) error { if err := validateBackendPlan(plan); err != nil { - return filter.ApplyResult{}, err + return err } if a.writer == nil { - return filter.ApplyResult{}, errors.New("ufw writer is required") + return errors.New("ufw writer is required") } for _, command := range append(append([]filter.NativeCommand(nil), plan.Rules[0].Commands...), plan.Rules[0].RollbackCommands...) { if err := validateCommand(command); err != nil { - return filter.ApplyResult{}, err + return err } } executed := 0 for index, command := range plan.Rules[0].Commands { - if err := a.writer.Run(ctx, command); err != nil { + err := a.writer.Run(ctx, command) + if err != nil && plan.CommandOnly && plan.CreatesOnly() && command.Args[0] == "insert" && + (strings.Contains(err.Error(), "Invalid position") || strings.Contains(err.Error(), "Cannot insert rule at position")) { + ipv4, ipv6 := plan.Scope, plan.Scope + ipv4.Family, ipv6.Family = filter.FamilyIPv4, filter.FamilyIPv6 + snapshots, readErr := a.ListRuleScopes(ctx, []filter.Scope{ipv4, ipv6}) + if readErr != nil { + return errors.Join(err, readErr) + } + lastPosition := maximumObservedPosition(snapshots[0]) + if plan.Scope.Family == filter.FamilyIPv6 { + lastPosition = max(lastPosition, maximumObservedPosition(snapshots[1])) + } + position, _ := strconv.Atoi(command.Args[1]) + if position == lastPosition+1 && !hasScopeNotice(snapshots[0].Notices, filter.ScopeNoticeManagedScopeInactive) { + err = a.writer.Run(ctx, commentCommand(plan.Rules[0].Expected.Rule, plan.Rules[0].Expected.Marker)) + } + } + if err != nil { if plan.CommandOnly { - return filter.ApplyResult{}, fmt.Errorf("execute UFW rule: %w", err) + return fmt.Errorf("execute UFW rule: %w", err) } if !plan.CreatesOnly() { probeCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) @@ -194,46 +244,22 @@ func (a *Adapter) Apply(ctx context.Context, plan filter.BackendPlan) (filter.Ap cancel() } cause := ufwApplyError(plan.Rules[0], fmt.Errorf("execute UFW rule: %w", err)) - return filter.ApplyResult{}, a.compensate(ctx, plan.Rules[0], executed, cause) + return a.compensate(ctx, plan.Rules[0], executed, cause) } executed = index + 1 } - if plan.CommandOnly || plan.CreatesOnly() { - return filter.ApplyResult{Applied: []filter.ObservedRule{plan.Rules[0].Expected}}, nil - } - verification, err := a.verify(ctx, plan) - if err != nil { - cause := ufwApplyError(plan.Rules[0], fmt.Errorf("verify UFW rule: %w", err)) - return filter.ApplyResult{}, a.compensate(ctx, plan.Rules[0], executed, cause) - } - if !verification.Matched { - cause := ufwApplyError(plan.Rules[0], fmt.Errorf( - "ufw write verification failed for marker %q in scope %s", - plan.Rules[0].Expected.Marker, - plan.Scope.Key(), - )) - return filter.ApplyResult{}, a.compensate( - ctx, - plan.Rules[0], - executed, - cause, - ) - } - return filter.ApplyResult{ - Applied: []filter.ObservedRule{plan.Rules[0].Expected}, - Verification: &verification, - }, nil + return nil } -func ufwApplyError(plan filter.NativeRulePlan, cause error) error { +func ufwApplyError(plan filter.RuleCommands, cause error) error { if plan.Operation == filter.ChangeAdopt { return buserr.WithDetail("ErrUFWRuleAdopt", cause.Error(), cause) } return cause } -func (a *Adapter) failedCommandApplied(ctx context.Context, plan filter.NativeRulePlan, commandIndex int) bool { - snapshot, err := a.Observe(ctx, plan.Expected.Rule.Scope) +func (a *Adapter) failedCommandApplied(ctx context.Context, plan filter.RuleCommands, commandIndex int) bool { + snapshot, err := a.ListRules(ctx, plan.Expected.Rule.Scope) if err != nil { return true } @@ -258,14 +284,7 @@ func (a *Adapter) failedCommandApplied(ctx context.Context, plan filter.NativeRu } } -func (a *Adapter) Verify(ctx context.Context, plan filter.BackendPlan) (filter.VerifyResult, error) { - if err := validateBackendPlan(plan); err != nil { - return filter.VerifyResult{}, err - } - return a.verify(ctx, plan) -} - -func (a *Adapter) Rollback(ctx context.Context, plan filter.BackendPlan) error { +func (a *Adapter) Rollback(ctx context.Context, plan filter.CommandBatch) error { if err := validateBackendPlan(plan); err != nil { return err } @@ -291,8 +310,8 @@ func validateScope(scope filter.Scope) error { return nil } -func validateBackendPlan(plan filter.BackendPlan) error { - if plan.Provider != filter.ProviderUFW || len(plan.Rules) != 1 || plan.SnapshotRevision == "" { +func validateBackendPlan(plan filter.CommandBatch) error { + if plan.Provider != filter.ProviderUFW || len(plan.Rules) != 1 { return fmt.Errorf("%w: invalid ufw backend plan", filter.ErrInvalidRule) } if err := validateScope(plan.Scope); err != nil { @@ -304,41 +323,53 @@ func validateBackendPlan(plan filter.BackendPlan) error { return nil } -func compileChange(snapshot filter.Snapshot, change filter.DesiredChange) (filter.NativeRulePlan, error) { +func compileChange(snapshot filter.RuleSet, change filter.RuleChange) (filter.RuleCommands, error) { rule := change.After if change.Operation == filter.ChangeDelete { rule = change.Before } if rule == nil { - return filter.NativeRulePlan{}, fmt.Errorf("%w: %s rule is required", filter.ErrInvalidRule, change.Operation) + return filter.RuleCommands{}, fmt.Errorf("%w: %s rule is required", filter.ErrInvalidRule, change.Operation) } normalized, err := filter.NormalizeRule(*rule) if err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } if normalized.Scope.Key() != snapshot.Scope.Key() { - return filter.NativeRulePlan{}, fmt.Errorf("%w: change scope %s", filter.ErrUnsupportedScope, normalized.Scope.Key()) + return filter.RuleCommands{}, fmt.Errorf("%w: change scope %s", filter.ErrUnsupportedScope, normalized.Scope.Key()) } if err := validateWritableRule(normalized); err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } if normalized.UUID == "" { - return filter.NativeRulePlan{}, fmt.Errorf("%w: rule UUID is required", filter.ErrInvalidRule) + return filter.RuleCommands{}, fmt.Errorf("%w: rule UUID is required", filter.ErrInvalidRule) } marker := "1panel-rule:" + normalized.UUID + if change.Operation == filter.ChangeDelete && change.CommandOnly && change.Locator == nil { + if change.UnmarkedAdopted || change.PreviousMarker != "" { + marker = observedComment(filter.ObservedRule{Rule: normalized, Marker: change.PreviousMarker}) + } + return filter.RuleCommands{ + RuleUUID: normalized.UUID, Operation: change.Operation, + Expected: filter.ObservedRule{Rule: normalized, Marker: marker, ParseStatus: filter.ParseStatusSupported}, + Commands: []filter.NativeCommand{deleteRuleCommand(normalized, marker)}, + RollbackCommands: []filter.NativeCommand{commentCommand(normalized, marker)}, + }, nil + } position := insertionPosition(snapshot, normalized) expected := observedForRule(normalized, marker, position) - plan := filter.NativeRulePlan{RuleUUID: normalized.UUID, Operation: change.Operation, Expected: expected} + plan := filter.RuleCommands{RuleUUID: normalized.UUID, Operation: change.Operation, Expected: expected} switch change.Operation { case filter.ChangeCreate: if normalized.OrderIndex != nil && *normalized.OrderIndex < 1 { - return filter.NativeRulePlan{}, fmt.Errorf("%w: create target is out of range", filter.ErrInvalidRule) + return filter.RuleCommands{}, fmt.Errorf("%w: create target is out of range", filter.ErrInvalidRule) } command := insertCommand(position, normalized, marker) - if !change.Append && normalized.OrderIndex != nil && position == 1 { + hasSnapshot := !change.CommandOnly || snapshot.Rules != nil + if !change.Append && normalized.OrderIndex != nil && position == 1 && hasSnapshot { command = filter.NativeCommand{Executable: "ufw", Args: append([]string{"prepend"}, compileRuleArgs(normalized, marker)...)} - } else if change.Append || position == maximumObservedPosition(snapshot)+1 { + } else if change.Append || normalized.OrderIndex == nil || hasSnapshot && position == maximumObservedPosition(snapshot)+1 { command = commentCommand(normalized, marker) } if command.Args[0] != "insert" { @@ -350,7 +381,7 @@ func compileChange(snapshot filter.Snapshot, change filter.DesiredChange) (filte case filter.ChangeAdopt: target, targetErr := validateMutationTarget(snapshot, change, normalized, marker, false) if targetErr != nil { - return filter.NativeRulePlan{}, targetErr + return filter.RuleCommands{}, targetErr } position = *target.Locator.Position appendAtEnd := position == maximumObservedPosition(snapshot) @@ -371,27 +402,27 @@ func compileChange(snapshot filter.Snapshot, change filter.DesiredChange) (filte } case filter.ChangeUpdate, filter.ChangeReorder: if change.Before == nil { - return filter.NativeRulePlan{}, fmt.Errorf("%w: previous ufw rule is required", filter.ErrInvalidRule) + return filter.RuleCommands{}, fmt.Errorf("%w: previous ufw rule is required", filter.ErrInvalidRule) } before, err := filter.NormalizeRule(*change.Before) if err != nil { - return filter.NativeRulePlan{}, err + return filter.RuleCommands{}, err } target, targetErr := validateMutationTarget(snapshot, change, before, marker, true) if targetErr != nil { - return filter.NativeRulePlan{}, targetErr + return filter.RuleCommands{}, targetErr } position = *target.Locator.Position targetPosition := position if normalized.OrderIndex != nil { if *normalized.OrderIndex < 1 { - return filter.NativeRulePlan{}, fmt.Errorf("%w: update target is out of range", filter.ErrInvalidRule) + return filter.RuleCommands{}, fmt.Errorf("%w: update target is out of range", filter.ErrInvalidRule) } targetPosition = int(*normalized.OrderIndex) } maximumPosition := maximumObservedPosition(snapshot) if targetPosition > maximumPosition { - return filter.NativeRulePlan{}, fmt.Errorf( + return filter.RuleCommands{}, fmt.Errorf( "%w: update target position %d is out of range 1-%d", filter.ErrInvalidRule, targetPosition, maximumPosition, ) @@ -415,18 +446,18 @@ func compileChange(snapshot filter.Snapshot, change filter.DesiredChange) (filte case filter.ChangeDelete: target, targetErr := validateMutationTarget(snapshot, change, normalized, marker, !change.UnmarkedAdopted) if targetErr != nil { - return filter.NativeRulePlan{}, targetErr + return filter.RuleCommands{}, targetErr } position = *target.Locator.Position plan.Previous = &target plan.Expected = target - plan.Commands = []filter.NativeCommand{deletePositionCommand(position)} + plan.Commands = []filter.NativeCommand{deleteRuleCommand(target.Rule, observedComment(target))} restoreAtEnd := change.RestoreAtEnd || position == maximumObservedPosition(snapshot) plan.RollbackCommands = []filter.NativeCommand{ positionedCommand(position, target.Rule, observedComment(target), restoreAtEnd), } default: - return filter.NativeRulePlan{}, fmt.Errorf("%w: unsupported operation %s", filter.ErrInvalidRule, change.Operation) + return filter.RuleCommands{}, fmt.Errorf("%w: unsupported operation %s", filter.ErrInvalidRule, change.Operation) } return plan, nil } @@ -445,7 +476,7 @@ func validateWritableRule(rule filter.FirewallRule) error { return nil } -func validateMutationTarget(snapshot filter.Snapshot, change filter.DesiredChange, desired filter.FirewallRule, marker string, requireOwned bool) (filter.ObservedRule, error) { +func validateMutationTarget(snapshot filter.RuleSet, change filter.RuleChange, desired filter.FirewallRule, marker string, requireOwned bool) (filter.ObservedRule, error) { if change.Locator == nil || change.Locator.Position == nil { return filter.ObservedRule{}, fmt.Errorf("%w: ufw mutation requires a numbered locator", filter.ErrInvalidRule) } @@ -515,7 +546,7 @@ func validateMutationTarget(snapshot filter.Snapshot, change filter.DesiredChang return target, nil } -func insertionPosition(snapshot filter.Snapshot, rule filter.FirewallRule) int { +func insertionPosition(snapshot filter.RuleSet, rule filter.FirewallRule) int { if rule.OrderIndex != nil && *rule.OrderIndex > 0 { return int(*rule.OrderIndex) } @@ -523,8 +554,8 @@ func insertionPosition(snapshot filter.Snapshot, rule filter.FirewallRule) int { return position } -func maximumObservedPosition(snapshot filter.Snapshot) int { - maximum := 0 +func maximumObservedPosition(snapshot filter.RuleSet) int { + maximum := snapshot.LastPosition for _, observed := range snapshot.Rules { if observed.Locator.Position != nil && *observed.Locator.Position > maximum { maximum = *observed.Locator.Position @@ -620,47 +651,7 @@ func observedComment(observed filter.ObservedRule) string { return observed.Rule.Description } -func (a *Adapter) verify(ctx context.Context, plan filter.BackendPlan) (filter.VerifyResult, error) { - scopes := append([]filter.Scope{plan.Scope}, relatedScopes(plan.Scope)...) - snapshots, err := a.ObserveScopes(ctx, scopes) - if err != nil { - return filter.VerifyResult{}, err - } - if len(snapshots) != len(scopes) { - return filter.VerifyResult{}, errors.New("ufw observation returned an incomplete scope set") - } - snapshot := snapshots[0] - rulePlan := plan.Rules[0] - marker := rulePlan.Expected.Marker - count := 0 - for _, observedSnapshot := range snapshots { - count += countMarker(observedSnapshot, marker) - } - if rulePlan.Operation == filter.ChangeDelete { - return filter.VerifyResult{Snapshot: snapshot, Matched: count == 0}, nil - } - if count != 1 { - return filter.VerifyResult{Snapshot: snapshot, Matched: false}, nil - } - for _, observed := range snapshot.Rules { - if observed.Marker != marker { - continue - } - positionMatches := rulePlan.Expected.Locator.Position == nil || - (observed.Locator.Position != nil && *observed.Locator.Position == *rulePlan.Expected.Locator.Position) - semanticMatches := filter.ObservedRuleMatchesExpected(observed, rulePlan.Expected.Rule) - if observed.ParseStatus == filter.ParseStatusOpaque { - semanticMatches = true - } - if !semanticMatches || !positionMatches { - return filter.VerifyResult{Snapshot: snapshot, Matched: false}, nil - } - return filter.VerifyResult{Snapshot: snapshot, Matched: true}, nil - } - return filter.VerifyResult{Snapshot: snapshot, Matched: false}, nil -} - -func (a *Adapter) compensate(ctx context.Context, plan filter.NativeRulePlan, executed int, cause error) error { +func (a *Adapter) compensate(ctx context.Context, plan filter.RuleCommands, executed int, cause error) error { if plan.Operation == filter.ChangeCreate { return cause } @@ -673,7 +664,7 @@ func (a *Adapter) compensate(ctx context.Context, plan filter.NativeRulePlan, ex return cause } -func (a *Adapter) rollback(ctx context.Context, plan filter.NativeRulePlan, executed int) error { +func (a *Adapter) rollback(ctx context.Context, plan filter.RuleCommands, executed int) error { var rollbackErr error for index := executed - 1; index >= 0; index-- { if index >= len(plan.RollbackCommands) { @@ -689,18 +680,7 @@ func (a *Adapter) rollback(ctx context.Context, plan filter.NativeRulePlan, exec return rollbackErr } -func relatedScopes(scope filter.Scope) []filter.Scope { - scope = scope.Normalize() - other := scope - if other.Family == filter.FamilyIPv4 { - other.Family = filter.FamilyIPv6 - } else { - other.Family = filter.FamilyIPv4 - } - return []filter.Scope{other} -} - -func countMarker(snapshot filter.Snapshot, marker string) int { +func countMarker(snapshot filter.RuleSet, marker string) int { count := 0 for _, observed := range snapshot.Rules { if observed.Marker == marker { @@ -710,7 +690,7 @@ func countMarker(snapshot filter.Snapshot, marker string) int { return count } -func containsObservedRule(snapshot filter.Snapshot, expected filter.ObservedRule) bool { +func containsObservedRule(snapshot filter.RuleSet, expected filter.ObservedRule) bool { for _, observed := range snapshot.Rules { if expected.Locator.Position != nil && (observed.Locator.Position == nil || *observed.Locator.Position != *expected.Locator.Position) { @@ -767,7 +747,7 @@ func parseNumberedRules(scope filter.Scope, output string) []filter.ObservedRule if err != nil || position < 1 { continue } - if !isInboundNumberedRule(matches[4], matches[5]) { + if matches[4] == "OUT" || matches[4] == "FWD" || strings.Contains(matches[5], "(out)") { continue } family := familyForNumberedRule(matches[2], matches[5]) @@ -785,7 +765,7 @@ func parseNumberedRule(scope filter.Scope, position int, destination, action, di positionCopy := position locator := filter.Locator{ Provider: filter.ProviderUFW, ScopeKey: scope.Key(), NativeID: strconv.Itoa(position), - Canonical: normalizedDisplay(raw), Position: &positionCopy, + Canonical: strings.Join(strings.Fields(raw), " "), Position: &positionCopy, } ruleAction := map[string]filter.Action{ "ALLOW": filter.ActionAccept, @@ -922,7 +902,7 @@ func looksLikeAddressedService(value string) bool { if len(tokens) < 2 { return false } - if isAnywhere(tokens[0]) { + if strings.EqualFold(strings.TrimSpace(tokens[0]), "Anywhere") { return true } _, ok := parseAddress(tokens[0]) @@ -992,7 +972,7 @@ func parseUnrecognizedNumberedRule(scope filter.Scope, raw string) (filter.Obser }, Locator: filter.Locator{ Provider: filter.ProviderUFW, ScopeKey: scope.Key(), NativeID: strconv.Itoa(position), - Canonical: normalizedDisplay(raw), Position: &positionCopy, + Canonical: strings.Join(strings.Fields(raw), " "), Position: &positionCopy, }, Marker: marker, ParseStatus: filter.ParseStatusOpaque, Raw: raw, Persistence: filter.PersistenceStatusConverged, }, family, inbound, true @@ -1036,7 +1016,7 @@ func parseDestination(value string) (address, port, protocol, iface, annotation tokens := strings.Fields(value) switch len(tokens) { case 1: - if isAnywhere(tokens[0]) { + if strings.EqualFold(strings.TrimSpace(tokens[0]), "Anywhere") { return "", "", "all", iface, annotation, true } if parsedAddress, parsedProtocol, endpointOK := parseAddressProtocol(tokens[0]); endpointOK { @@ -1068,7 +1048,7 @@ func parseAddressProtocol(value string) (string, string, bool) { return "", "", false } endpoint := strings.TrimSpace(value[:separator]) - if isAnywhere(endpoint) { + if strings.EqualFold(strings.TrimSpace(endpoint), "Anywhere") { return "", protocol, true } address, ok := parseAddress(endpoint) @@ -1101,7 +1081,7 @@ func splitDestinationAnnotation(value string) (string, string) { } func parseSource(value string) (string, bool) { - if isAnywhere(value) { + if strings.EqualFold(strings.TrimSpace(value), "Anywhere") { return "", true } if strings.Contains(value, " on ") || len(strings.Fields(value)) != 1 { @@ -1110,10 +1090,6 @@ func parseSource(value string) (string, bool) { return parseAddress(value) } -func isInboundNumberedRule(direction, source string) bool { - return direction != "OUT" && direction != "FWD" && !strings.Contains(source, "(out)") -} - func splitInterface(value string) (endpoint, iface string, ok bool) { const delimiter = " on " index := strings.LastIndex(value, delimiter) @@ -1168,7 +1144,7 @@ func parsePortProtocol(value string) (string, string, bool) { } func parseAddress(value string) (string, bool) { - if isAnywhere(value) { + if strings.EqualFold(strings.TrimSpace(value), "Anywhere") { return "", true } if prefix, err := netip.ParsePrefix(value); err == nil { @@ -1180,10 +1156,6 @@ func parseAddress(value string) (string, bool) { return "", false } -func isAnywhere(value string) bool { - return strings.EqualFold(strings.TrimSpace(value), "Anywhere") -} - func canonicalRule(rule filter.FirewallRule) string { return strings.Join([]string{ string(rule.Action), rule.Protocol, rule.SourceAddress, rule.DestinationAddress, @@ -1191,10 +1163,6 @@ func canonicalRule(rule filter.FirewallRule) string { }, "|") } -func normalizedDisplay(raw string) string { - return strings.Join(strings.Fields(raw), " ") -} - func statusNotices(numbered string) []filter.ScopeNotice { notices := make([]filter.ScopeNotice, 0, 1) if !statusActive(numbered) { @@ -1240,6 +1208,10 @@ func IsIPv6Unavailable(err error) bool { type systemBackend struct{} +func (systemBackend) ReadRulesByComment(ctx context.Context, _ filter.Scope, comment string) (string, error) { + return filter.ReadRulesByComment(ctx, "ufw", []string{"status", "numbered"}, comment) +} + func (systemBackend) Read(ctx context.Context, args ...string) (string, error) { return cmd.NewCommandMgr( cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second), cmd.WithEnv("LANGUAGE=en_US:en"), @@ -1250,7 +1222,14 @@ func (systemBackend) Run(ctx context.Context, command filter.NativeCommand) erro if err := validateCommand(command); err != nil { return err } - return cmd.NewCommandMgr( + output, err := cmd.NewCommandMgr( cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second), cmd.WithEnv("LANGUAGE=en_US:en"), - ).RunWithOptionalSudo(command.Executable, command.Args...) + ).RunWithOptionalSudoAndStdout(command.Executable, command.Args...) + if err != nil { + return err + } + if strings.Contains(output, "Could not delete non-existent rule") { + return fmt.Errorf("%w: %s", filter.ErrRuleStale, strings.TrimSpace(output)) + } + return nil } diff --git a/agent/utils/firewall/filter/runtime/inventory.go b/agent/utils/firewall/filter/runtime/inventory.go deleted file mode 100644 index 06cca2cb9483..000000000000 --- a/agent/utils/firewall/filter/runtime/inventory.go +++ /dev/null @@ -1,123 +0,0 @@ -package runtime - -import ( - "context" - "strings" - "sync" - "sync/atomic" - "time" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -const inventoryTTL = 2 * time.Second - -var inventoryGeneration atomic.Uint64 - -func InvalidateInventory() { inventoryGeneration.Add(1) } - -type inventoryEntry struct { - snapshots []filter.Snapshot - expires time.Time - generation uint64 -} - -type inventoryCache struct { - mu sync.Mutex - entries map[string]inventoryEntry - reading map[string]chan struct{} -} - -func (e *Engine) ObserveInventory(ctx context.Context, scope filter.Scope, refresh bool) (filter.Snapshot, error) { - snapshots, err := e.ObserveInventoryScopes(ctx, []filter.Scope{scope}, refresh) - if err != nil { - return filter.Snapshot{}, err - } - return snapshots[0], nil -} - -func (e *Engine) ObserveInventoryScopes(ctx context.Context, scopes []filter.Scope, refresh bool) ([]filter.Snapshot, error) { - read := func() ([]filter.Snapshot, error) { - if len(scopes) == 1 { - snapshot, err := e.adapter.Observe(ctx, scopes[0]) - return []filter.Snapshot{snapshot}, err - } - observer, ok := e.adapter.(filter.MultiScopeObserver) - if !ok { - return nil, filter.ErrAdapterUnavailable - } - return observer.ObserveScopes(ctx, scopes) - } - var snapshots []filter.Snapshot - var err error - if e.Provider() == filter.ProviderFirewalld || e.Provider() == filter.ProviderUFW { - keys := make([]string, len(scopes)) - for i, scope := range scopes { - keys[i] = scope.Normalize().Key() - } - snapshots, err = e.inventory.load(ctx, strings.Join(keys, "\n"), refresh, read) - } else { - snapshots, err = read() - } - if err != nil { - return nil, err - } - result := append([]filter.Snapshot(nil), snapshots...) - for i := range result { - result[i].Rules = append([]filter.ObservedRule(nil), snapshots[i].Rules...) - result[i].Notices = append([]filter.ScopeNotice(nil), snapshots[i].Notices...) - if e.policy != nil { - result[i], err = e.policy(ctx, result[i]) - if err != nil { - return nil, err - } - } - } - return result, nil -} - -func (c *inventoryCache) load(ctx context.Context, key string, refresh bool, read func() ([]filter.Snapshot, error)) ([]filter.Snapshot, error) { - for { - if err := ctx.Err(); err != nil { - return nil, err - } - c.mu.Lock() - generation := inventoryGeneration.Load() - if entry, ok := c.entries[key]; !refresh && ok && entry.generation == generation && time.Now().Before(entry.expires) { - c.mu.Unlock() - return entry.snapshots, nil - } - if done := c.reading[key]; done != nil { - c.mu.Unlock() - select { - case <-ctx.Done(): - return nil, ctx.Err() - case <-done: - refresh = false - continue - } - } - if c.reading == nil { - c.reading = make(map[string]chan struct{}) - } - done := make(chan struct{}) - c.reading[key] = done - delete(c.entries, key) - c.mu.Unlock() - snapshots, err := read() - c.mu.Lock() - if err == nil && generation == inventoryGeneration.Load() { - if len(c.entries) >= 32 { - clear(c.entries) - } - if c.entries == nil { - c.entries = make(map[string]inventoryEntry) - } - c.entries[key] = inventoryEntry{snapshots, time.Now().Add(inventoryTTL), generation} - } - delete(c.reading, key) - close(done) - c.mu.Unlock() - return snapshots, err - } -} diff --git a/agent/utils/firewall/filter/runtime/runtime.go b/agent/utils/firewall/filter/runtime/runtime.go deleted file mode 100644 index 75dce91b6b3f..000000000000 --- a/agent/utils/firewall/filter/runtime/runtime.go +++ /dev/null @@ -1,404 +0,0 @@ -package runtime - -import ( - "context" - "errors" - "fmt" - "time" - - "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 - inventory inventoryCache -} - -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) NewObservationSession() *Engine { - if factory, ok := e.adapter.(filter.ObservationSessionFactory); ok { - return New(factory.NewObservationSession(), e.policy) - } - return e -} - -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) AppendUnverified(ctx context.Context, rule filter.FirewallRule, comment string) error { - InvalidateInventory() - defer InvalidateInventory() - appender, ok := e.adapter.(filter.UnverifiedRuleAppender) - if !ok { - return fmt.Errorf("%w: %s does not support unverified rule appends", filter.ErrAdapterUnavailable, e.Provider()) - } - return appender.AppendUnverified(ctx, rule, comment) -} - -func (e *Engine) CompileDesired( - ctx context.Context, - policyUUID string, - origin filter.RuleOrigin, - rules []filter.FirewallRule, -) ([]filter.DesiredRule, error) { - capabilities, err := e.Capabilities(ctx) - if err != nil { - return nil, err - } - result := make([]filter.DesiredRule, 0, len(rules)) - for ordinal, 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 - } - prepared.UUID = compiledRuleUUID(policyUUID, ruleKey, ordinal) - desired := filter.DesiredRule{UUID: policyUUID, Rule: prepared, RuleKey: ruleKey, Origin: origin} - if capabilities.Marker { - desired.Marker = "1panel-rule:" + prepared.UUID - } - result = append(result, desired) - } - return result, nil -} - -func (e *Engine) ValidatePosition( - ctx context.Context, - snapshot filter.Snapshot, - rule filter.FirewallRule, - target int64, -) error { - if target < 1 { - return fmt.Errorf("%w: target position must be positive", filter.ErrInvalidRule) - } - 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) ExecuteCreate(ctx context.Context, snapshot filter.Snapshot, changes []filter.DesiredChange) error { - InvalidateInventory() - defer InvalidateInventory() - plan, err := e.adapter.Compile(snapshot, changes) - if err != nil { - return err - } - if !plan.CreatesOnly() { - return fmt.Errorf("%w: expected a create-only plan", filter.ErrInvalidRule) - } - _, err = e.adapter.Apply(ctx, plan) - return err -} - -func (e *Engine) ExecuteSync(ctx context.Context, snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.ApplyResult, error) { - InvalidateInventory() - defer InvalidateInventory() - if err := ctx.Err(); err != nil { - return filter.ApplyResult{}, err - } - changes = append([]filter.DesiredChange(nil), changes...) - for index := range changes { - changes[index].CommandOnly = true - } - plan, err := e.adapter.Compile(snapshot, changes) - if err != nil { - return filter.ApplyResult{}, err - } - plan.CommandOnly = true - return e.adapter.Apply(ctx, plan) -} - -func (e *Engine) Execute(ctx context.Context, snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, filter.VerifyResult, error) { - InvalidateInventory() - defer InvalidateInventory() - 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 && !plan.CreatesOnly() { - 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 { - if plan.CreatesOnly() { - return plan, verification, err - } - return plan, verification, e.rollback(ctx, plan, err) - } - if !verification.Matched && !plan.CreatesOnly() { - 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 { - InvalidateInventory() - defer InvalidateInventory() - 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()) - } - ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) - defer cancel() - 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) -} - -func (e *Engine) NewCreatePlanner(snapshot filter.Snapshot) (filter.CreatePlanner, error) { - factory, ok := e.adapter.(filter.CreatePlannerFactory) - if !ok { - return nil, filter.ErrAdapterUnavailable - } - return factory.NewCreatePlanner(snapshot), nil -} - -func (e *Engine) ExecutePlannedCreate(ctx context.Context, planner filter.CreatePlanner, change filter.DesiredChange) (filter.ObservedRule, error) { - if err := ctx.Err(); err != nil { - return filter.ObservedRule{}, err - } - plan, err := planner.Compile(change) - if err != nil { - return filter.ObservedRule{}, err - } - if !plan.CreatesOnly() { - return filter.ObservedRule{}, filter.ErrInvalidRule - } - InvalidateInventory() - defer InvalidateInventory() - plan.CommandOnly = change.CommandOnly - result, err := e.adapter.Apply(ctx, plan) - if err != nil { - return filter.ObservedRule{}, err - } - if len(result.Applied) != 1 { - return filter.ObservedRule{}, filter.ErrVerificationFailed - } - planner.Applied(result.Applied[0]) - return result.Applied[0], nil -} diff --git a/agent/utils/firewall/filter/safety.go b/agent/utils/firewall/filter/safety.go index 8eecd23129ba..aaa5dccc9dce 100644 --- a/agent/utils/firewall/filter/safety.go +++ b/agent/utils/firewall/filter/safety.go @@ -17,27 +17,65 @@ var ( var ErrVerificationFailed = errors.New("firewall rule verification failed") -func ProtectSnapshot(snapshot Snapshot, ports []PortWhitelist) (Snapshot, error) { +func ProtectRuleSet(snapshot RuleSet, ports []PortWhitelist) (RuleSet, error) { rules := append([]ObservedRule(nil), snapshot.Rules...) + whitelist := NewPortWhitelistIndex(ports) for index := range rules { - if rules[index].ParseStatus == ParseStatusSupported && RuleMatchesPortWhitelist(rules[index].Rule, ports) { + if rules[index].ParseStatus == ParseStatusSupported && whitelist.Matches(rules[index].Rule) { rules[index].Protected = true } } protected := snapshot protected.Rules = rules - if protected.Revision == "" { - var err error - protected, err = NewSnapshot(snapshot.Scope, rules) + + protected.LastPosition = snapshot.LastPosition + protected.Notices = append([]ScopeNotice(nil), snapshot.Notices...) + return protected, nil +} + +type portWhitelistKey struct { + family Family + protocol, port, source string +} + +type PortWhitelistIndex map[portWhitelistKey]bool + +func NewPortWhitelistIndex(ports []PortWhitelist) PortWhitelistIndex { + index := make(PortWhitelistIndex) + for _, port := range ports { + protocol, err := normalizeProtocol(port.Protocol) if err != nil { - return Snapshot{}, err + continue + } + portRange, err := normalizePortValue(port.Port, false) + if err != nil { + continue + } + portFamily := Family(strings.ToLower(strings.TrimSpace(port.Family))) + sources := port.Sources + if len(sources) == 0 { + sources = []string{""} + } + for _, family := range []Family{FamilyIPv4, FamilyIPv6} { + if portFamily != "" && portFamily != FamilyInet && portFamily != family { + continue + } + for _, source := range sources { + normalized, err := normalizeAddress(source, family) + if err == nil { + index[portWhitelistKey{family, protocol, portRange, normalized}] = true + } + } } } - protected.Notices = append([]ScopeNotice(nil), snapshot.Notices...) - return protected, nil + return index } func RuleMatchesPortWhitelist(rule FirewallRule, ports []PortWhitelist) bool { + return NewPortWhitelistIndex(ports).Matches(rule) +} + +func (index PortWhitelistIndex) Matches(rule FirewallRule) bool { rule, err := NormalizeRule(rule) if err != nil || rule.Action != ActionAccept || rule.SourcePort != "" || rule.DestinationAddress != "" || rule.Interface != "" || len(rule.ConnectionStates) != 0 { return false @@ -51,36 +89,7 @@ func RuleMatchesPortWhitelist(rule FirewallRule, ports []PortWhitelist) bool { families = []Family{FamilyIPv4, FamilyIPv6} } for _, family := range families { - matched := false - for _, port := range ports { - portFamily := Family(strings.ToLower(strings.TrimSpace(port.Family))) - if portFamily != "" && !familiesOverlap(family, portFamily) { - continue - } - protocol, err := normalizeProtocol(port.Protocol) - if err != nil || rule.Protocol != protocol { - continue - } - portRange, err := normalizePort(port.Port) - if err != nil || rule.DestinationPort != portRange { - continue - } - sources := port.Sources - if len(sources) == 0 { - sources = []string{""} - } - for _, source := range sources { - normalized, err := normalizeAddress(source, family) - if err == nil && normalized == rule.SourceAddress { - matched = true - break - } - } - if matched { - break - } - } - if !matched { + if !index[portWhitelistKey{family, rule.Protocol, rule.DestinationPort, rule.SourceAddress}] { return false } } @@ -153,27 +162,7 @@ func MatchObservedByRuleKey(observed []ObservedRule, rule FirewallRule) ([]Obser return matches, nil } -func ManagedObserved(snapshot Snapshot, desired DesiredRule) (ObservedRule, error) { - items, err := MergeInventory(InventoryMergeInput{Observed: snapshot.Rules, Desired: []DesiredRule{desired}}) - if err != nil { - return ObservedRule{}, err - } - for _, item := range items { - if item.Desired == nil || item.Desired.UUID != desired.UUID { - continue - } - if item.State == InventoryStateProtected || (item.Observed != nil && item.Observed.Protected) { - return ObservedRule{}, ErrProtectedRule - } - if item.Observed == nil || item.Match != InventoryMatchExact || item.State == InventoryStateDrifted { - return ObservedRule{}, ErrRuleStale - } - return *item.Observed, nil - } - return ObservedRule{}, ErrRuleStale -} - -func FindCommittedObserved(snapshot Snapshot, requested FirewallRule, plan BackendPlan) (ObservedRule, error) { +func FindCommittedObserved(snapshot RuleSet, requested FirewallRule, plan CommandBatch) (ObservedRule, error) { if len(plan.Rules) == 1 && plan.Rules[0].Expected.Marker != "" { matches := make([]ObservedRule, 0, 1) for _, observed := range snapshot.Rules { @@ -201,8 +190,8 @@ func RulesOverlap(left, right FirewallRule) bool { if leftErr != nil || rightErr != nil || left.Scope.Key() != right.Scope.Key() { return false } - return familiesOverlap(left.Scope.Family, right.Scope.Family) && - protocolsOverlap(left.Protocol, right.Protocol) && + return (left.Scope.Family == FamilyInet || right.Scope.Family == FamilyInet || left.Scope.Family == right.Scope.Family) && + (left.Protocol == "all" || right.Protocol == "all" || left.Protocol == right.Protocol) && addressesOverlap(left.SourceAddress, right.SourceAddress) && addressesOverlap(left.DestinationAddress, right.DestinationAddress) && portsOverlap(left.SourcePort, right.SourcePort) && @@ -210,14 +199,6 @@ func RulesOverlap(left, right FirewallRule) bool { (left.Interface == "" || right.Interface == "" || left.Interface == right.Interface) } -func familiesOverlap(left, right Family) bool { - return left == FamilyInet || right == FamilyInet || left == right -} - -func protocolsOverlap(left, right string) bool { - return left == "all" || right == "all" || left == right -} - func addressesOverlap(left, right string) bool { if left == "" || right == "" { return true @@ -278,25 +259,3 @@ func portInterval(value string) (int, int, error) { end, err := strconv.Atoi(parts[1]) return start, end, err } - -var ErrDuplicateAdoption = fmt.Errorf("%w: duplicate firewall rules prevent adoption; manually delete duplicate rules and retry", ErrRuleOperation) - -func CheckAdoptDuplicates(snapshot Snapshot, requested FirewallRule) error { - count := 0 - for _, observed := range snapshot.Rules { - if observed.ParseStatus != ParseStatusSupported { - continue - } - same, err := SameRuleContent(observed.Rule, requested) - if err != nil { - return err - } - if same { - count++ - if count > 1 { - return ErrDuplicateAdoption - } - } - } - return nil -} diff --git a/agent/utils/firewall/forwarding/forwarding.go b/agent/utils/firewall/forwarding/forwarding.go index 6b340e62f890..4adf4ba8090e 100644 --- a/agent/utils/firewall/forwarding/forwarding.go +++ b/agent/utils/firewall/forwarding/forwarding.go @@ -1,30 +1,25 @@ package forwarding import ( - "errors" + "context" "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 - + 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" + ForwardFile = "1panel_forward.rules" + PreRoutingFile = "1panel_forward_pre.rules" + PostRoutingFile = "1panel_forward_post.rules" ) type Rule struct { @@ -51,7 +46,9 @@ const ( type Adapter interface { Name() string List() ([]Rule, error) - Reconcile(rules []Rule) error + CreateRules(context.Context, []Rule) error + DeleteRules(context.Context, []Rule) error + ReplaceRules(rules []Rule) error Enable() error Cleanup() error InitStatus() (bool, bool, error) @@ -59,91 +56,6 @@ type Adapter interface { 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 New(provider string) (Adapter, error) { - switch provider { - case "iptables": - return newIptablesNATAdapter(provider), nil - case "nftables": - return newNftablesAdapter(), nil - default: - return nil, errors.New("unsupported forwarding provider: " + provider) - } -} - -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 == "" { diff --git a/agent/utils/firewall/forwarding/iptables.go b/agent/utils/firewall/forwarding/iptables.go index db6e15262676..cb29ffdd32b3 100644 --- a/agent/utils/firewall/forwarding/iptables.go +++ b/agent/utils/firewall/forwarding/iptables.go @@ -24,7 +24,7 @@ type iptablesBackend interface { RunWithStd(table string, args ...string) (string, error) RunIPv6(table string, args ...string) error RunIPv6WithStd(table string, args ...string) (string, error) - Restore(family, input string) error + Restore(ctx context.Context, family, input string) error LoadRulesFromFile(table, chain, fileName string) error LoadIPv6RulesFromFile(table, chain, fileName string) error } @@ -58,7 +58,7 @@ func (systemIptablesBackend) RunIPv6WithStd(table string, args ...string) (strin return iptables_helper.RunIPv6WithStd(table, args...) } -func (systemIptablesBackend) Restore(family, input string) error { +func (systemIptablesBackend) Restore(ctx context.Context, family, input string) error { commands, err := lifecycle.ResolveIptablesCommands() if err != nil { return err @@ -70,8 +70,12 @@ func (systemIptablesBackend) Restore(family, input string) error { return fmt.Errorf("ip6tables-restore command family is unavailable") } } - manager := cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second), cmd.WithStdin(strings.NewReader(input))) + var stderr strings.Builder + manager := cmd.NewCommandMgr(cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second), cmd.WithStdin(strings.NewReader(input)), cmd.WithStderr(&stderr)) err = manager.RunWithOptionalSudo(executable, "--noflush", "--wait") + if err == nil && strings.TrimSpace(stderr.String()) != "" { + err = fmt.Errorf("firewall command warning: %s", strings.TrimSpace(stderr.String())) + } return firewallutil.WrapBatchCommandError(executable+" --noflush --wait", input, err) } @@ -103,25 +107,25 @@ func (defaultForwardingSystem) RunWithOptionalSudo(name string, args ...string) return cmd.NewCommandMgr().RunWithOptionalSudo(name, args...) } -type iptablesNATAdapter struct { +type Iptables struct { provider string backend iptablesBackend system forwardingSystem } -func newIptablesNATAdapter(provider string) *iptablesNATAdapter { - return &iptablesNATAdapter{ +func NewIptables(provider string) *Iptables { + return &Iptables{ provider: provider, backend: systemIptablesBackend{}, system: defaultForwardingSystem{}, } } -func (l *iptablesNATAdapter) Name() string { +func (l *Iptables) Name() string { return l.provider } -func (l *iptablesNATAdapter) List() ([]Rule, error) { +func (l *Iptables) List() ([]Rule, error) { stdout, err := l.backend.RunWithStd(iptables_helper.NatTab, "-S") if err != nil { return nil, fmt.Errorf("failed to list NAT rules: %w", err) @@ -137,7 +141,7 @@ func (l *iptablesNATAdapter) List() ([]Rule, error) { return append(rules, parseIptablesRules(stdout, FamilyIPv6)...), nil } -func (l *iptablesNATAdapter) Reconcile(rules []Rule) error { +func (l *Iptables) ReplaceRules(rules []Rule) error { byFamily := map[string][]Rule{ FamilyIPv4: nil, FamilyIPv6: nil, @@ -159,24 +163,73 @@ func (l *iptablesNATAdapter) Reconcile(rules []Rule) error { if err := l.batchEnsureChains(family); err != nil { return err } - script, err := buildIptablesForwardRestoreScript(byFamily[family]) + script, err := buildIptablesForwardScript(byFamily[family], OperationAdd, true) if err != nil { return err } - if err := l.backend.Restore(family, script); err != nil { + if err := l.backend.Restore(context.Background(), family, script); err != nil { return fmt.Errorf("restore %s forwarding rules: %w", family, err) } } return nil } -func buildIptablesForwardRestoreScript(rules []Rule) (string, error) { - natRules := [][]string{{"-F", ChainPreRouting}, {"-F", ChainPostRouting}} - filterRules := [][]string{{"-F", ChainForward}} +func (l *Iptables) CreateRules(ctx context.Context, rules []Rule) error { + if len(rules) == 0 { + return nil + } + script, err := buildIptablesForwardScript(rules, OperationAdd, false) + if err != nil { + return err + } + family := rules[0].Family + if family == "" { + family = FamilyIPv4 + } + return l.backend.Restore(ctx, family, script) +} + +func (l *Iptables) DeleteRules(ctx context.Context, rules []Rule) error { + if len(rules) == 0 { + return nil + } + script, err := buildIptablesForwardScript(rules, OperationRemove, false) + if err != nil { + return err + } + family := rules[0].Family + if family == "" { + family = FamilyIPv4 + } + return l.backend.Restore(ctx, family, script) +} + +func buildIptablesForwardScript(rules []Rule, operation OperationType, replace bool) (string, error) { + var natRules, filterRules [][]string + if replace { + natRules = [][]string{{"-F", ChainPreRouting}, {"-F", ChainPostRouting}} + filterRules = [][]string{{"-F", ChainForward}} + } + verb := "-A" + if operation == OperationRemove { + verb = "-D" + } else if operation != OperationAdd { + return "", fmt.Errorf("unsupported forwarding operation %q", operation) + } + family := "" for _, rule := range rules { + normalized, err := NormalizeRule(rule) + if err != nil { + return "", err + } + rule = normalized + if family != "" && family != rule.Family { + return "", fmt.Errorf("iptables forwarding batch must use one address family") + } + family = rule.Family sourcePort := strings.ReplaceAll(rule.Port, "-", ":") targetPort := strings.ReplaceAll(rule.TargetPort, "-", ":") - preRouting := []string{"-A", ChainPreRouting} + preRouting := []string{verb, ChainPreRouting} if rule.Interface != "" { preRouting = append(preRouting, "-i", rule.Interface) } @@ -187,11 +240,11 @@ func buildIptablesForwardRestoreScript(rules []Rule) (string, error) { } natRules = append(natRules, append(preRouting, "-j", "DNAT", "--to-destination", forwardingTarget(rule)), - []string{"-A", ChainPostRouting, "-d", rule.TargetIP, "-p", rule.Protocol, "--dport", targetPort, "-j", "MASQUERADE"}, + []string{verb, ChainPostRouting, "-d", rule.TargetIP, "-p", rule.Protocol, "--dport", targetPort, "-j", "MASQUERADE"}, ) filterRules = append(filterRules, - []string{"-A", ChainForward, "-d", rule.TargetIP, "-p", rule.Protocol, "--dport", targetPort, "-j", "ACCEPT"}, - []string{"-A", ChainForward, "-s", rule.TargetIP, "-p", rule.Protocol, "--sport", targetPort, "-j", "ACCEPT"}, + []string{verb, ChainForward, "-d", rule.TargetIP, "-p", rule.Protocol, "--dport", targetPort, "-j", "ACCEPT"}, + []string{verb, ChainForward, "-s", rule.TargetIP, "-p", rule.Protocol, "--sport", targetPort, "-j", "ACCEPT"}, ) } var script strings.Builder @@ -233,7 +286,7 @@ func isRemoteTarget(family, target string) bool { return target != "" && target != "127.0.0.1" && target != "localhost" } -func (l *iptablesNATAdapter) Enable() error { +func (l *Iptables) Enable() error { if err := ensureForwardingSysctls(l.system, l.backend.IPv6Available()); err != nil { return err } @@ -249,7 +302,7 @@ func (l *iptablesNATAdapter) Enable() error { return nil } -func (l *iptablesNATAdapter) batchEnsureChains(family string) error { +func (l *Iptables) batchEnsureChains(family string) error { list := l.backend.RunWithStd if family == FamilyIPv6 { list = l.backend.RunIPv6WithStd @@ -266,13 +319,13 @@ func (l *iptablesNATAdapter) batchEnsureChains(family string) error { if script == "" { return nil } - if err := l.backend.Restore(family, script); err != nil { + if err := l.backend.Restore(context.Background(), family, script); err != nil { return fmt.Errorf("batch initialize %s forwarding chains: %w", family, err) } return nil } -func (l *iptablesNATAdapter) Cleanup() error { +func (l *Iptables) Cleanup() error { for _, family := range []string{FamilyIPv4, FamilyIPv6} { if family == FamilyIPv6 && !l.backend.IPv6Available() { continue @@ -291,7 +344,7 @@ func (l *iptablesNATAdapter) Cleanup() error { } script := buildIptablesForwardLifecycleScript(outputs, false) if script != "" { - if err := l.backend.Restore(family, script); err != nil { + if err := l.backend.Restore(context.Background(), family, script); err != nil { return fmt.Errorf("batch delete %s forwarding chains: %w", family, err) } } @@ -441,7 +494,7 @@ func containsExactLine(output, want string) bool { return false } -func (l *iptablesNATAdapter) InitStatus() (bool, bool, error) { +func (l *Iptables) InitStatus() (bool, bool, error) { ipv4Init, ipv4Bind, err := l.familyInitStatus(FamilyIPv4) if err != nil { return false, false, err @@ -456,7 +509,7 @@ func (l *iptablesNATAdapter) InitStatus() (bool, bool, error) { return ipv4Init && ipv6Init, ipv4Bind && ipv6Bind, nil } -func (l *iptablesNATAdapter) familyInitStatus(family string) (bool, bool, error) { +func (l *Iptables) familyInitStatus(family string) (bool, bool, error) { sysctlPath := "/proc/sys/net/ipv4/ip_forward" label := "IPv4" list := l.backend.RunWithStd @@ -495,7 +548,7 @@ func (l *iptablesNATAdapter) familyInitStatus(family string) (bool, bool, error) return natInit && filterInit, forwardingEnabled && natBind && filterInit && filterBind, nil } -func (l *iptablesNATAdapter) FamilyStatus(family string) (bool, bool, error) { +func (l *Iptables) FamilyStatus(family string) (bool, bool, error) { if family == FamilyIPv6 && !l.backend.IPv6Available() { return false, false, nil } @@ -525,7 +578,7 @@ func containsExactRule(lines []string, rule string) bool { return false } -func (l *iptablesNATAdapter) Replay() error { +func (l *Iptables) Replay() error { for _, family := range []string{FamilyIPv4, FamilyIPv6} { if family == FamilyIPv6 && !l.backend.IPv6Available() { continue diff --git a/agent/utils/firewall/forwarding/nftables.go b/agent/utils/firewall/forwarding/nftables.go index c9f31d2ef136..789d5b410262 100644 --- a/agent/utils/firewall/forwarding/nftables.go +++ b/agent/utils/firewall/forwarding/nftables.go @@ -1,6 +1,7 @@ package forwarding import ( + "context" "encoding/base64" "errors" "fmt" @@ -12,6 +13,7 @@ import ( "github.com/1Panel-dev/1Panel/agent/global" "github.com/1Panel-dev/1Panel/agent/utils/cmd" + firewallutil "github.com/1Panel-dev/1Panel/agent/utils/firewall" "github.com/1Panel-dev/1Panel/agent/utils/firewall/nftables_helper" ) @@ -22,18 +24,18 @@ const ( nftForwardMarker = "1panel-forward:" ) -type nftablesAdapter struct{ system forwardingSystem } +type Nftables struct{ system forwardingSystem } -func newNftablesAdapter() *nftablesAdapter { - return &nftablesAdapter{system: defaultForwardingSystem{}} +func NewNftables() *Nftables { + return &Nftables{system: defaultForwardingSystem{}} } -func (n *nftablesAdapter) Name() string { return "nftables" } +func (n *Nftables) Name() string { return "nftables" } -func (n *nftablesAdapter) List() ([]Rule, error) { +func (n *Nftables) List() ([]Rule, error) { rules := make([]Rule, 0) for _, family := range []string{FamilyIPv4, FamilyIPv6} { - stdout, err := nftables_helper.ReadChain(nftRun, nftTableFamily(family), nftForwardTable, nftForwardChain(ChainPreRouting)) + stdout, err := nftables_helper.ReadChain(nftRun, nftTableFamily(family), nftForwardTable, "NFT_"+ChainPreRouting) if errors.Is(err, nftables_helper.ErrChainNotFound) { continue } @@ -45,7 +47,7 @@ func (n *nftablesAdapter) List() ([]Rule, error) { return rules, nil } -func (n *nftablesAdapter) Reconcile(rules []Rule) error { +func (n *Nftables) ReplaceRules(rules []Rule) error { if err := ensureNftForwardTables(); err != nil { return fmt.Errorf("initialize nftables forwarding table: %w", err) } @@ -53,10 +55,64 @@ func (n *nftablesAdapter) Reconcile(rules []Rule) error { if err != nil { return err } - return nftRunCommands(commands) + return nftRunCommands(context.Background(), commands) } -func (n *nftablesAdapter) Enable() error { +func (n *Nftables) CreateRules(ctx context.Context, rules []Rule) error { + if len(rules) == 0 { + return nil + } + commands, err := createNftForwardCommands(rules) + if err != nil { + return err + } + return nftRunCommands(ctx, commands) +} + +func (n *Nftables) DeleteRules(ctx context.Context, rules []Rule) error { + wanted := make(map[string]map[string]bool) + for _, rule := range rules { + normalized, err := NormalizeRule(rule) + if err != nil { + return err + } + family := nftTableFamily(normalized.Family) + if wanted[family] == nil { + wanted[family] = make(map[string]bool) + } + wanted[family][normalized.Identity()] = true + } + var commands [][]string + for _, family := range []string{"ip", "ip6"} { + if len(wanted[family]) == 0 { + continue + } + output, _, err := nftables_helper.ReadTable(func(args ...string) (string, error) { + return cmd.NewCommandMgr(cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second)).RunWithOptionalSudoAndStdout("nft", args...) + }, family, nftForwardTable) + if err != nil { + return err + } + chains := nftables_helper.ParseTableChains(output) + for _, chain := range []string{ChainPreRouting, ChainPostRouting, ChainForward} { + for _, rule := range parseNftForwardRules(chains["NFT_"+chain]) { + if !wanted[family][rule.Identity()] { + continue + } + if _, err := strconv.ParseUint(rule.Num, 10, 64); err != nil { + return fmt.Errorf("invalid nftables forwarding handle %q", rule.Num) + } + commands = append(commands, []string{"delete", "rule", family, nftForwardTable, "NFT_" + chain, "handle", rule.Num}) + } + } + } + if len(commands) == 0 { + return nil + } + return nftRunCommands(ctx, commands) +} + +func (n *Nftables) Enable() error { if err := ensureForwardingSysctls(n.system, true); err != nil { return err } @@ -66,7 +122,7 @@ func (n *nftablesAdapter) Enable() error { return nil } -func (n *nftablesAdapter) Cleanup() error { +func (n *Nftables) Cleanup() error { commands := make([][]string, 0, 2) for _, family := range []string{FamilyIPv4, FamilyIPv6} { tableFamily := nftTableFamily(family) @@ -76,7 +132,7 @@ func (n *nftablesAdapter) Cleanup() error { commands = append(commands, []string{"delete", "table", tableFamily, nftForwardTable}) } if len(commands) > 0 { - if err := nftRunCommands(commands); err != nil { + if err := nftRunCommands(context.Background(), commands); err != nil { return err } } @@ -87,7 +143,7 @@ func (n *nftablesAdapter) Cleanup() error { return nil } -func (n *nftablesAdapter) InitStatus() (bool, bool, error) { +func (n *Nftables) InitStatus() (bool, bool, error) { for _, family := range []string{FamilyIPv4, FamilyIPv6} { initialized, bound, err := n.FamilyStatus(family) if err != nil || !initialized || !bound { @@ -97,7 +153,7 @@ func (n *nftablesAdapter) InitStatus() (bool, bool, error) { return true, true, nil } -func (n *nftablesAdapter) FamilyStatus(family string) (bool, bool, error) { +func (n *Nftables) FamilyStatus(family string) (bool, bool, error) { sysctlPath := "/proc/sys/net/ipv4/ip_forward" if family == FamilyIPv6 { sysctlPath = "/proc/sys/net/ipv6/conf/all/forwarding" @@ -106,15 +162,23 @@ func (n *nftablesAdapter) FamilyStatus(family string) (bool, bool, error) { if err != nil { return false, false, fmt.Errorf("read %s forwarding status: %w", family, err) } + output, exists, err := nftables_helper.ReadTable(nftRun, nftTableFamily(family), nftForwardTable) + if err != nil { + return false, false, err + } + if !exists { + return false, false, nil + } + chains := nftables_helper.ParseTableChains(output) for _, chain := range []string{ChainPreRouting, ChainPostRouting, ChainForward} { - if _, err := nftRun("list", "chain", nftTableFamily(family), nftForwardTable, nftForwardChain(chain)); err != nil { + if _, exists := chains["NFT_"+chain]; !exists { return false, false, nil } } return true, strings.TrimSpace(string(data)) != "0", nil } -func (n *nftablesAdapter) Replay() error { +func (n *Nftables) Replay() error { file := filepath.Join(global.Dir.FirewallDir, nftForwardFile) if _, err := os.Stat(file); errors.Is(err, os.ErrNotExist) { return nil @@ -137,23 +201,24 @@ func ensureNftForwardTables() error { commands := make([][]string, 0, 8) for _, family := range []string{FamilyIPv4, FamilyIPv6} { tableFamily := nftTableFamily(family) - tableExists := true - if _, err := nftRun("list", "table", tableFamily, nftForwardTable); err != nil { - tableExists = false + output, tableExists, err := nftables_helper.ReadTable(nftRun, tableFamily, nftForwardTable) + if err != nil { + return err + } + existingChains := nftables_helper.ParseTableChains(output) + if !tableExists { commands = append(commands, []string{"add", "table", tableFamily, nftForwardTable}) } chains := []struct { name, chainType, hook, priority string }{ - {nftForwardChain(ChainPreRouting), "nat", "prerouting", "-100"}, - {nftForwardChain(ChainPostRouting), "nat", "postrouting", "100"}, - {nftForwardChain(ChainForward), "filter", "forward", "0"}, + {"NFT_" + ChainPreRouting, "nat", "prerouting", "-100"}, + {"NFT_" + ChainPostRouting, "nat", "postrouting", "100"}, + {"NFT_" + ChainForward, "filter", "forward", "0"}, } for _, chain := range chains { - if tableExists { - if _, err := nftRun("list", "chain", tableFamily, nftForwardTable, chain.name); err == nil { - continue - } + if _, exists := existingChains[chain.name]; exists { + continue } commands = append(commands, []string{ "add", "chain", tableFamily, nftForwardTable, chain.name, @@ -164,16 +229,22 @@ func ensureNftForwardTables() error { if len(commands) == 0 { return nil } - return nftRunCommands(commands) + return nftRunCommands(context.Background(), commands) } func rebuildNftForwardCommands(rules []Rule) ([][]string, error) { commands := make([][]string, 0, 6+len(rules)*4) for _, family := range []string{FamilyIPv4, FamilyIPv6} { for _, chain := range []string{ChainPreRouting, ChainPostRouting, ChainForward} { - commands = append(commands, []string{"flush", "chain", nftTableFamily(family), nftForwardTable, nftForwardChain(chain)}) + commands = append(commands, []string{"flush", "chain", nftTableFamily(family), nftForwardTable, "NFT_" + chain}) } } + additions, err := createNftForwardCommands(rules) + return append(commands, additions...), err +} + +func createNftForwardCommands(rules []Rule) ([][]string, error) { + commands := make([][]string, 0, len(rules)*4) for _, rule := range rules { normalized, err := NormalizeRule(rule) if err != nil { @@ -181,25 +252,25 @@ func rebuildNftForwardCommands(rules []Rule) ([][]string, error) { } rule = normalized tableFamily := nftTableFamily(rule.Family) - addressKeyword := nftAddressKeyword(rule.Family) + addressKeyword := tableFamily comment := strconv.Quote(encodeNftForwardRule(rule)) interfaceMatch := make([]string, 0, 2) if rule.Interface != "" { interfaceMatch = append(interfaceMatch, "iifname", strconv.Quote(rule.Interface)) } if isRemoteTarget(rule.Family, rule.TargetIP) { - preRouting := []string{"add", "rule", tableFamily, nftForwardTable, nftForwardChain(ChainPreRouting)} + preRouting := []string{"add", "rule", tableFamily, nftForwardTable, "NFT_" + ChainPreRouting} preRouting = append(preRouting, interfaceMatch...) preRouting = append(preRouting, "meta", "l4proto", rule.Protocol, rule.Protocol, "dport", rule.Port, "dnat", "to", forwardingTarget(rule), "comment", comment) commands = append(commands, preRouting, - []string{"add", "rule", tableFamily, nftForwardTable, nftForwardChain(ChainPostRouting), addressKeyword, "daddr", rule.TargetIP, "meta", "l4proto", rule.Protocol, rule.Protocol, "dport", rule.TargetPort, "masquerade", "comment", comment}, - []string{"add", "rule", tableFamily, nftForwardTable, nftForwardChain(ChainForward), addressKeyword, "daddr", rule.TargetIP, "meta", "l4proto", rule.Protocol, rule.Protocol, "dport", rule.TargetPort, "accept", "comment", comment}, - []string{"add", "rule", tableFamily, nftForwardTable, nftForwardChain(ChainForward), addressKeyword, "saddr", rule.TargetIP, "meta", "l4proto", rule.Protocol, rule.Protocol, "sport", rule.TargetPort, "accept", "comment", comment}, + []string{"add", "rule", tableFamily, nftForwardTable, "NFT_" + ChainPostRouting, addressKeyword, "daddr", rule.TargetIP, "meta", "l4proto", rule.Protocol, rule.Protocol, "dport", rule.TargetPort, "masquerade", "comment", comment}, + []string{"add", "rule", tableFamily, nftForwardTable, "NFT_" + ChainForward, addressKeyword, "daddr", rule.TargetIP, "meta", "l4proto", rule.Protocol, rule.Protocol, "dport", rule.TargetPort, "accept", "comment", comment}, + []string{"add", "rule", tableFamily, nftForwardTable, "NFT_" + ChainForward, addressKeyword, "saddr", rule.TargetIP, "meta", "l4proto", rule.Protocol, rule.Protocol, "sport", rule.TargetPort, "accept", "comment", comment}, ) continue } - preRouting := []string{"add", "rule", tableFamily, nftForwardTable, nftForwardChain(ChainPreRouting)} + preRouting := []string{"add", "rule", tableFamily, nftForwardTable, "NFT_" + ChainPreRouting} preRouting = append(preRouting, interfaceMatch...) preRouting = append(preRouting, "meta", "l4proto", rule.Protocol, rule.Protocol, "dport", rule.Port, "redirect", "to", ":"+rule.TargetPort, "comment", comment) commands = append(commands, preRouting) @@ -214,13 +285,6 @@ func nftTableFamily(family string) string { return nftForwardFamily } -func nftAddressKeyword(family string) string { - if family == FamilyIPv6 { - return "ip6" - } - return "ip" -} - func encodeNftForwardRule(rule Rule) string { family, protocol := "4", "t" if rule.Family == FamilyIPv6 { @@ -312,10 +376,6 @@ func parseNftForwardRules(stdout string) []Rule { return result } -func nftForwardChain(logical string) string { - return "NFT_" + logical -} - func nftRun(args ...string) (string, error) { stdout, err := cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudoAndStdout("nft", args...) if err != nil { @@ -332,12 +392,17 @@ func nftRunCommand(args ...string) error { return nil } -func nftRunCommands(commands [][]string) error { +func nftRunCommands(ctx context.Context, commands [][]string) error { script, err := nftCommandsScript(commands) if err != nil { return err } - return nftables_helper.RunScript(script) + var stderr strings.Builder + err = cmd.NewCommandMgr(cmd.WithContext(ctx), cmd.WithTimeout(60*time.Second), cmd.WithStdin(strings.NewReader(script)), cmd.WithStderr(&stderr)).RunWithOptionalSudo("nft", "-f", "-") + if err == nil && strings.TrimSpace(stderr.String()) != "" { + err = fmt.Errorf("firewall command warning: %s", strings.TrimSpace(stderr.String())) + } + return firewallutil.WrapBatchCommandError("nft -f -", script, err) } func nftCommandsScript(commands [][]string) (string, error) { diff --git a/agent/utils/firewall/iptables_helper/ipv6.go b/agent/utils/firewall/iptables_helper/ipv6.go index 0e253d87feee..908163f9da67 100644 --- a/agent/utils/firewall/iptables_helper/ipv6.go +++ b/agent/utils/firewall/iptables_helper/ipv6.go @@ -10,14 +10,6 @@ import ( "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" ) -func (m *Manager) EnsureIPv6BaseChains() error { - ports, err := m.loadRequiredPorts() - if err != nil { - return err - } - return EnsureIPv6BaseChains(ports) -} - func EnsureIPv6BaseChains(ports []firewall.PortWhitelist) error { commands, err := lifecycle.ResolveIptablesCommands() if err != nil || !commands.IPv6Available() { @@ -40,16 +32,7 @@ func EnsureIPv6BaseChains(ports []firewall.PortWhitelist) error { if err := setBaseChainBindings(true, true); err != nil { return err } - for _, chain := range []struct{ name, file string }{ - {BasicBeforeChain, IPv6FileName(BasicBeforeFileName)}, - {BasicChain, IPv6FileName(BasicFileName)}, - {BasicAfterChain, IPv6FileName(BasicAfterFileName)}, - } { - if err := SaveIPv6RulesToFile(FilterTab, chain.name, chain.file); err != nil { - return err - } - } - return nil + return saveBaseChainsFamily(true) } func UnbindIPv6BaseChains() error { diff --git a/agent/utils/firewall/iptables_helper/manager.go b/agent/utils/firewall/iptables_helper/manager.go index 37c4a159b274..c5b4d0590e32 100644 --- a/agent/utils/firewall/iptables_helper/manager.go +++ b/agent/utils/firewall/iptables_helper/manager.go @@ -17,13 +17,8 @@ import ( "github.com/mattn/go-shellwords" ) -type Manager struct { - UpdateSetting func(key, value string) error - LoadRequiredPorts func() ([]firewall.PortWhitelist, error) -} - -func (m *Manager) Cleanup() error { - if err := m.disableBase(); err != nil { +func Cleanup() error { + if err := disableBase(); err != nil { return err } if err := cleanupBaseChains(false); err != nil { @@ -43,31 +38,31 @@ func (m *Manager) Cleanup() error { return nil } -func (m *Manager) Operate(operation firewall.BaseOperation) error { +func Operate(operation firewall.BaseOperation, requiredPorts []firewall.PortWhitelist) error { switch operation { case firewall.BaseOperationInit, firewall.BaseOperationBind: if _, err := lifecycle.ResolveIptablesCommands(); err != nil { return fmt.Errorf("failed to find iptables") } - return m.enableBase(true) + return enableBase(true, requiredPorts) case firewall.BaseOperationBindWithoutInit: - return m.enableBase(false) + return enableBase(false, requiredPorts) case firewall.BaseOperationUnbind: - return m.disableBase() + return disableBase() default: return fmt.Errorf("unsupported iptables base operation %q", operation) } } -func (m *Manager) enableBase(prepare bool) error { +func enableBase(prepare bool, requiredPorts []firewall.PortWhitelist) error { if prepare { if err := ensureBaseChainsFamily(false); err != nil { return err } - if err := m.initPreRules(); err != nil { + if err := applyRequiredFirewallPortWhiteListRules(requiredPorts, false, true, false); err != nil { return err } - if err := saveBaseChains(); err != nil { + if err := saveBaseChainsFamily(false); err != nil { return err } } @@ -75,26 +70,32 @@ func (m *Manager) enableBase(prepare bool) error { return err } if prepare { - if err := m.ensureIPv6BaseChains(); err != nil { + commands, err := lifecycle.ResolveIptablesCommands() + if err != nil { return err } - if err := m.SyncRequiredPorts(true); err != nil { + if commands.IPv6Available() { + if err := EnsureIPv6BaseChains(requiredPorts); err != nil { + return err + } + } + if err := syncRequiredPorts(requiredPorts, true); err != nil { return err } } else if err := BindIPv6BaseChains(); err != nil { return err } - return m.updateSetting("IptablesStatus", constant.StatusEnable) + return nil } -func (m *Manager) disableBase() error { +func disableBase() error { if err := setBaseChainBindings(false, false); err != nil { return err } if err := UnbindIPv6BaseChains(); err != nil && !errors.Is(err, filter.ErrFamilyUnavailable) { return err } - return m.updateSetting("IptablesStatus", constant.StatusDisable) + return nil } func ensureBaseChainsFamily(ipv6 bool) error { @@ -224,13 +225,28 @@ func baseChainBindingCommands(output string, bind bool) []string { return lines } -func saveBaseChains() error { +func saveBaseChainsFamily(ipv6 bool) error { + read := RunWithStd + if ipv6 { + read = RunIPv6WithStd + } + output, err := read(FilterTab, "-S") + if err != nil { + return err + } for _, item := range []struct{ chain, file string }{ {BasicBeforeChain, BasicBeforeFileName}, {BasicChain, BasicFileName}, {BasicAfterChain, BasicAfterFileName}, } { - if err := SaveRulesToFile(FilterTab, item.chain, item.file); err != nil { + if !containsIptablesRule(output, "-N "+item.chain) { + return fmt.Errorf("cannot save missing iptables chain %s", item.chain) + } + file := item.file + if ipv6 { + file = IPv6FileName(file) + } + if err := writeChainRules(output, item.chain, file); err != nil { return err } } @@ -322,27 +338,7 @@ func buildBaseChainsRestoreScript(firewallDir string, ipv6 bool, requiredPorts . return script.String(), nil } -func (m *Manager) initPreRules() error { - requiredPorts, err := m.loadRequiredPorts() - if err != nil { - return err - } - return applyRequiredFirewallPortWhiteListRules(requiredPorts, false, true, false) -} - -func (m *Manager) ensureIPv6BaseChains() error { - commands, err := lifecycle.ResolveIptablesCommands() - if err != nil || !commands.IPv6Available() { - return nil - } - return m.EnsureIPv6BaseChains() -} - -func (m *Manager) SyncRequiredPorts(withSave bool) error { - requiredPorts, err := m.loadRequiredPorts() - if err != nil { - return err - } +func syncRequiredPorts(requiredPorts []firewall.PortWhitelist, withSave bool) error { commands, err := lifecycle.ResolveIptablesCommands() if err != nil { return err @@ -422,12 +418,7 @@ func applyRequiredFirewallPortWhiteListRules(portWhiteList []firewall.PortWhitel return save(FilterTab, BasicAfterChain, afterFile) } -func buildRequiredPortsRestoreScript( - desired []firewall.SystemPort, - family string, - beforeRaw, afterRaw string, - includeDefaults bool, -) string { +func buildRequiredPortsRestoreScript(desired []firewall.SystemPort, family string, beforeRaw, afterRaw string, includeDefaults bool) string { var commands []string for _, line := range []string{"-A " + BasicBeforeChain + " " + IoRuleIn, "-A " + BasicBeforeChain + " " + EstablishedRule} { if !containsIptablesRule(beforeRaw, line) { @@ -500,17 +491,3 @@ func countIptablesRule(output, rule string) int { } return count } - -func (m *Manager) updateSetting(key, value string) error { - if m != nil && m.UpdateSetting != nil { - return m.UpdateSetting(key, value) - } - return nil -} - -func (m *Manager) loadRequiredPorts() ([]firewall.PortWhitelist, error) { - if m != nil && m.LoadRequiredPorts != nil { - return m.LoadRequiredPorts() - } - return nil, fmt.Errorf("load required firewall ports is not configured") -} diff --git a/agent/utils/firewall/iptables_helper/persistence.go b/agent/utils/firewall/iptables_helper/persistence.go index 827d1895cf5d..2f6044360002 100644 --- a/agent/utils/firewall/iptables_helper/persistence.go +++ b/agent/utils/firewall/iptables_helper/persistence.go @@ -53,8 +53,6 @@ func IPv6FileName(fileName string) string { } func saveRulesToFile(ctx context.Context, executable, tab, chain, fileName string) error { - rulesFile := path.Join(global.Dir.FirewallDir, fileName) - var stdout string var err error if strings.HasPrefix(path.Base(executable), "ip6tables") { @@ -65,11 +63,16 @@ func saveRulesToFile(ctx context.Context, executable, tab, chain, fileName strin if err != nil { return fmt.Errorf("failed to list %s rules: %w", chain, err) } + return writeChainRules(stdout, chain, fileName) +} + +func writeChainRules(stdout, chain, fileName string) error { + rulesFile := path.Join(global.Dir.FirewallDir, fileName) var rules []string lines := strings.Split(stdout, "\n") for _, line := range lines { line = strings.TrimSpace(line) - if strings.HasPrefix(line, fmt.Sprintf("-A %s", chain)) { + if strings.HasPrefix(line, "-A "+chain+" ") { rules = append(rules, line) } } diff --git a/agent/utils/firewall/iptables_helper/repair.go b/agent/utils/firewall/iptables_helper/repair.go index 60b3dd788833..676eaa9363bd 100644 --- a/agent/utils/firewall/iptables_helper/repair.go +++ b/agent/utils/firewall/iptables_helper/repair.go @@ -9,11 +9,12 @@ import ( "github.com/1Panel-dev/1Panel/agent/constant" "github.com/1Panel-dev/1Panel/agent/global" + "github.com/1Panel-dev/1Panel/agent/utils/firewall" "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle" ) -func (m *Manager) RepairBaseChains() error { +func RepairBaseChains(ports []firewall.PortWhitelist) error { commands, err := lifecycle.ResolveIptablesCommands() if err != nil { return err @@ -35,10 +36,6 @@ func (m *Manager) RepairBaseChains() error { return err } script, err := buildBaseChainsRepairScript(global.Dir.FirewallDir, output, ipv6, func() ([]string, error) { - ports, err := m.loadRequiredPorts() - if err != nil { - return nil, err - } return baseDefaultRules(ports, family) }) if err != nil { diff --git a/agent/utils/firewall/lifecycle/lifecycle.go b/agent/utils/firewall/lifecycle/lifecycle.go index 656821fcb7bb..d9b8e678603b 100644 --- a/agent/utils/firewall/lifecycle/lifecycle.go +++ b/agent/utils/firewall/lifecycle/lifecycle.go @@ -10,6 +10,14 @@ import ( "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle/providers" ) +type Operation string + +const ( + OperationStart Operation = "start" + OperationStop Operation = "stop" + OperationRestart Operation = "restart" +) + const ( ProviderFirewalld = constant.FirewallProviderFirewalld ProviderUFW = constant.FirewallProviderUFW @@ -109,15 +117,14 @@ type PreStopResetter interface { ResetBeforeStop() error } -func NewClient() (Client, error) { - runtime, err := DetectRuntime() - if err != nil { - return nil, err +func NewClient(provider string) (Client, error) { + if provider == "" { + runtime, err := DetectRuntime() + if err != nil { + return nil, err + } + provider = runtime.Provider } - return NewClientFor(runtime.Provider) -} - -func NewClientFor(provider string) (Client, error) { switch provider { case "firewalld": if !which("firewalld") { diff --git a/agent/utils/firewall/lifecycle/operator.go b/agent/utils/firewall/lifecycle/operator.go deleted file mode 100644 index d1cc45540aee..000000000000 --- a/agent/utils/firewall/lifecycle/operator.go +++ /dev/null @@ -1,178 +0,0 @@ -package lifecycle - -import ( - "errors" - "fmt" - "os" - - "github.com/1Panel-dev/1Panel/agent/global" - "github.com/1Panel-dev/1Panel/agent/utils/controller" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle/providers" -) - -const fail2BanRestoreWithFirewallMarker = "/run/1panel_fail2ban_restore_with_firewall" - -type Operation string - -const ( - OperationStart Operation = "start" - OperationStop Operation = "stop" - OperationRestart Operation = "restart" -) - -type Operator struct { - client Client - RunAction func(operation, name string, action func() error) error -} - -// DockerRestartError reports that the requested firewall operation completed, -// but rebuilding Docker's firewall rules failed. -type DockerRestartError struct { - Err error -} - -func (e *DockerRestartError) Error() string { - return fmt.Sprintf("failed to restart Docker: %v", e.Err) -} - -func (e *DockerRestartError) Unwrap() error { - return e.Err -} - -type CompletedOperationError struct { - Operation Operation - Err error -} - -func (e *CompletedOperationError) Error() string { - return fmt.Sprintf("firewall %s completed with recovery errors: %v", e.Operation, e.Err) -} - -func (e *CompletedOperationError) Unwrap() error { - return e.Err -} - -func NewOperator(client Client) *Operator { - return &Operator{client: client} -} - -func (o *Operator) runAction(operation, name string, action func() error) error { - if o.RunAction != nil { - return o.RunAction(operation, name, action) - } - return action() -} - -func (o *Operator) Operate(operation Operation, withDockerRestart bool, prepareStart func(Client) error) error { - var recoveryErrors []error - switch operation { - case OperationStart: - if err := o.runAction("Start", o.client.Name(), o.client.Start); err != nil { - return err - } - if prepareStart != nil { - if err := o.prepareAfterStart(prepareStart); err != nil { - recoveryErrors = append(recoveryErrors, fmt.Errorf("prepare firewall after start: %w", err)) - } - } - case OperationStop: - return o.StopWithPrepare(withDockerRestart, nil) - case OperationRestart: - if err := o.runAction("TaskRestart", o.client.Name(), o.client.Restart); err != nil { - return err - } - if prepareStart != nil { - if err := o.prepareAfterStart(prepareStart); err != nil { - recoveryErrors = append(recoveryErrors, fmt.Errorf("prepare firewall after restart: %w", err)) - } - } - default: - return fmt.Errorf("not supported operation: %s", operation) - } - - if withDockerRestart { - if err := o.runAction("TaskRestart", "Docker", func() error { return controller.HandleRestart("docker") }); err != nil { - recoveryErrors = append(recoveryErrors, &DockerRestartError{Err: err}) - } - } - if o.client.Name() == ProviderFirewalld && operation == OperationStart { - if err := o.runAction("TaskRecover", "Fail2Ban", restoreFail2BanAfterFirewallStart); err != nil { - recoveryErrors = append(recoveryErrors, err) - } - } - if err := errors.Join(recoveryErrors...); err != nil { - return &CompletedOperationError{Operation: operation, Err: err} - } - return nil -} - -func (o *Operator) prepareAfterStart(prepare func(Client) error) error { - if err := prepare(o.client); err != nil { - return err - } - if o.client.Name() == ProviderFirewalld { - return providers.RemoveFirewalldSSHService() - } - return nil -} - -// StopWithPrepare records dependent service state, runs preparation, stops the -// firewall, and optionally restarts Docker in that order. -func (o *Operator) StopWithPrepare(withDockerRestart bool, prepareStop func() error) error { - if o.client.Name() == ProviderFirewalld { - if err := rememberFail2BanBeforeFirewallStop(); err != nil { - return err - } - } - if prepareStop != nil { - if err := prepareStop(); err != nil { - return err - } - } - if err := o.runAction("Stop", o.client.Name(), o.client.Stop); err != nil { - return err - } - if withDockerRestart { - if err := o.runAction("TaskRestart", "Docker", func() error { return controller.HandleRestart("docker") }); err != nil { - return &DockerRestartError{Err: err} - } - } - return nil -} - -func rememberFail2BanBeforeFirewallStop() error { - exists, err := controller.CheckExist("fail2ban.service") - if err != nil { - global.LOG.Warnf("check fail2ban.service installation before stopping the firewall failed: %v", err) - } - if !exists { - return nil - } - active, err := controller.CheckActive("fail2ban.service") - if err != nil { - global.LOG.Warnf("check fail2ban.service status before stopping the firewall failed: %v", err) - } - if !active { - return nil - } - if err := os.WriteFile(fail2BanRestoreWithFirewallMarker, nil, 0600); err != nil { - return fmt.Errorf("mark Fail2Ban for restoration with the firewall: %w", err) - } - return nil -} - -func restoreFail2BanAfterFirewallStart() error { - if _, err := os.Stat(fail2BanRestoreWithFirewallMarker); err != nil { - if os.IsNotExist(err) { - return nil - } - return fmt.Errorf("load Fail2Ban restore marker after starting the firewall: %w", err) - } - if err := controller.HandleStart("fail2ban.service"); err != nil { - return fmt.Errorf("restore Fail2Ban after starting the firewall: %w", err) - } - if err := os.Remove(fail2BanRestoreWithFirewallMarker); err != nil && !os.IsNotExist(err) { - return fmt.Errorf("clear Fail2Ban firewall restore status: %w", err) - } - return nil -} diff --git a/agent/utils/firewall/lifecycle/providers/firewalld.go b/agent/utils/firewall/lifecycle/providers/firewalld.go index baa9bbf5491f..9a88bb50986e 100644 --- a/agent/utils/firewall/lifecycle/providers/firewalld.go +++ b/agent/utils/firewall/lifecycle/providers/firewalld.go @@ -38,7 +38,8 @@ func (f *Firewalld) Name() string { func (f *Firewalld) Status() (bool, error) { stdout, err := cmd.NewCommandMgr(cmd.WithEnv("LANGUAGE=en_US:en")).RunWithStdout("firewall-cmd", "--state") if err != nil { - if firewalldStopped(stdout, err) { + message := strings.ToLower(strings.TrimSpace(stdout)) + " " + strings.ToLower(err.Error()) + if strings.Contains(message, "not running") { return false, nil } return false, fmt.Errorf("load firewall status failed: %w", err) @@ -46,14 +47,6 @@ func (f *Firewalld) Status() (bool, error) { return strings.TrimSpace(stdout) == "running", nil } -func firewalldStopped(stdout string, err error) bool { - message := strings.ToLower(strings.TrimSpace(stdout)) - if err != nil { - message += " " + strings.ToLower(err.Error()) - } - return strings.Contains(message, "not running") -} - func (f *Firewalld) Version() (string, error) { stdout, err := cmd.NewCommandMgr(cmd.WithEnv("LANGUAGE=en_US:en")).RunWithStdout("firewall-cmd", "--version") if err != nil { @@ -70,15 +63,13 @@ func (f *Firewalld) Start() error { } func RemoveFirewalldSSHService() error { - for _, permanent := range []bool{true, false} { - args := []string{"--zone=" + filter.FirewalldInputZone, "--remove-service=ssh"} - configuration := "runtime" - if permanent { - args = append(args, "--permanent") - configuration = "permanent" - } - if _, err := cmd.NewCommandMgr(cmd.WithEnv("LANGUAGE=en_US:en")).RunWithStdout("firewall-cmd", args...); err != nil { - return fmt.Errorf("remove firewalld SSH service from %s configuration: %w", configuration, err) + manager := cmd.NewCommandMgr(cmd.WithEnv("LANGUAGE=en_US:en")) + for _, args := range [][]string{ + {"--zone=" + filter.FirewalldInputZone, "--remove-service=ssh"}, + {"--permanent", "--zone=" + filter.FirewalldInputZone, "--remove-service=ssh"}, + } { + if err := manager.RunWithOptionalSudo("firewall-cmd", args...); err != nil { + return fmt.Errorf("remove firewalld SSH service: %w", err) } } return nil @@ -121,12 +112,7 @@ func (f *Firewalld) ResetBeforeStop() error { return nil } -func replaceFirewalldConfig( - configDir string, - backupDir string, - prepare func(string) error, - validate func() error, -) (func() error, error) { +func replaceFirewalldConfig(configDir string, backupDir string, prepare func(string) error, validate func() error) (func() error, error) { info, err := os.Lstat(configDir) hadConfig := err == nil if err != nil && !os.IsNotExist(err) { diff --git a/agent/utils/firewall/nftables_helper/manager.go b/agent/utils/firewall/nftables_helper/manager.go index 4a70db197785..9a109e7b55c0 100644 --- a/agent/utils/firewall/nftables_helper/manager.go +++ b/agent/utils/firewall/nftables_helper/manager.go @@ -4,27 +4,20 @@ import ( "context" "errors" "fmt" + "github.com/1Panel-dev/1Panel/agent/global" + "github.com/1Panel-dev/1Panel/agent/utils/firewall" + "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" "net/netip" "os" "path/filepath" "slices" "strconv" "strings" - - "github.com/1Panel-dev/1Panel/agent/constant" - "github.com/1Panel-dev/1Panel/agent/global" - "github.com/1Panel-dev/1Panel/agent/utils/firewall" - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" ) const requiredPortComment = "1Panel Port Whitelist" -type Manager struct { - UpdateSetting func(key, value string) error - LoadRequiredPorts func() ([]firewall.PortWhitelist, error) -} - -func (m *Manager) Cleanup() error { +func Cleanup() error { commands := make([][]string, 0, 2) for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} { tableFamily := TableFamily(family) @@ -40,69 +33,58 @@ func (m *Manager) Cleanup() error { if err := os.Remove(file); err != nil && !errors.Is(err, os.ErrNotExist) { return err } - return m.updateSetting("IptablesStatus", constant.StatusDisable) + return nil } -func (m *Manager) Operate(operation firewall.BaseOperation) error { +func Operate(operation firewall.BaseOperation, requiredPorts []firewall.PortWhitelist) error { switch operation { case firewall.BaseOperationInit, firewall.BaseOperationBind: - return m.enableBase(true) + return enableBase(true, requiredPorts) case firewall.BaseOperationBindWithoutInit: - return m.enableBase(false) + return enableBase(false, requiredPorts) case firewall.BaseOperationUnbind: - return m.disableBase() + return Unbind() default: return fmt.Errorf("unsupported nftables base operation %q", operation) } } -func (m *Manager) enableBase(prepare bool) error { +func enableBase(prepare bool, requiredPorts []firewall.PortWhitelist) error { if prepare { - if err := m.ensureBaseChains(); err != nil { + if err := ensureBaseChains(); err != nil { return err } - if err := m.initPreRules(); err != nil { + if err := initPreRules(requiredPorts); err != nil { return err } } if err := Bind(); err != nil { return err } - return m.updateSetting("IptablesStatus", constant.StatusEnable) -} - -func (m *Manager) disableBase() error { - if err := Unbind(); err != nil { - return err - } - return m.updateSetting("IptablesStatus", constant.StatusDisable) + return nil } -func (m *Manager) ensureBaseChains() error { +func ensureBaseChains() error { commands := make([][]string, 0, 10) for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} { tableFamily := TableFamily(family) - tableExists := true - if _, err := run("list", "table", tableFamily, TableName); err != nil { - tableExists = false - commands = append(commands, []string{"add", "table", tableFamily, TableName}) + output, tableExists, err := ReadTable(run, tableFamily, TableName) + if err != nil { + return err } + chains := ParseTableChains(output) if !tableExists { - commands = append(commands, []string{ - "add", "chain", tableFamily, TableName, InputChain, - "{", "type", "filter", "hook", "input", "priority", "0", ";", "policy", "accept", ";", "}", - }) - } else if _, err := run("list", "chain", tableFamily, TableName, InputChain); err != nil { + commands = append(commands, []string{"add", "table", tableFamily, TableName}) + } + if _, exists := chains[InputChain]; !exists { commands = append(commands, []string{ "add", "chain", tableFamily, TableName, InputChain, "{", "type", "filter", "hook", "input", "priority", "0", ";", "policy", "accept", ";", "}", }) } for _, nativeChain := range BasicChains() { - if tableExists { - if _, err := run("list", "chain", tableFamily, TableName, nativeChain); err == nil { - continue - } + if _, exists := chains[nativeChain]; exists { + continue } commands = append(commands, []string{"add", "chain", tableFamily, TableName, nativeChain}) } @@ -124,12 +106,8 @@ func requiredPortCommand(tableFamily string, rule firewall.SystemPort) []string "accept", "comment", `"`+requiredPortComment+`"`) } -func (m *Manager) initPreRules() error { - ports, err := m.loadRequiredPorts() - if err != nil { - return err - } - ports, err = firewall.NormalizeRequiredPorts(ports) +func initPreRules(requiredPorts []firewall.PortWhitelist) error { + ports, err := firewall.NormalizeRequiredPorts(requiredPorts) if err != nil { return err } @@ -218,20 +196,6 @@ func containsRequiredPortRule(output, expression string) bool { return false } -func (m *Manager) updateSetting(key, value string) error { - if m != nil && m.UpdateSetting != nil { - return m.UpdateSetting(key, value) - } - return nil -} - -func (m *Manager) loadRequiredPorts() ([]firewall.PortWhitelist, error) { - if m != nil && m.LoadRequiredPorts != nil { - return m.LoadRequiredPorts() - } - return nil, fmt.Errorf("load required firewall ports is not configured") -} - func Bind() error { for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} { tableFamily := TableFamily(family) diff --git a/agent/utils/firewall/nftables_helper/runtime.go b/agent/utils/firewall/nftables_helper/runtime.go index b2362f94dd35..c7336b38ecd9 100644 --- a/agent/utils/firewall/nftables_helper/runtime.go +++ b/agent/utils/firewall/nftables_helper/runtime.go @@ -169,23 +169,53 @@ func hasBaseChainBinding(output string) bool { } func loadFamilyInitStatus(family filter.Family) (bool, bool, error) { + output, exists, err := ReadTable(run, TableFamily(family), TableName) + if err != nil || !exists { + return false, false, err + } + chains := ParseTableChains(output) for _, chain := range BasicChains() { - if _, exists, err := readNftObject(run, "list", "chain", TableFamily(family), TableName, chain); err != nil || !exists { - return false, false, err + if _, exists := chains[chain]; !exists { + return false, false, nil } } - stdout, exists, err := readNftObject(run, "list", "chain", TableFamily(family), TableName, InputChain) - if err != nil || !exists { - return false, false, err + input, exists := chains[InputChain] + if !exists { + return false, false, nil } for _, chain := range BasicChains() { - if !strings.Contains(stdout, "jump "+chain) { + if !strings.Contains(input, "jump "+chain) { return true, false, nil } } return true, true, nil } +func ReadTable(run func(...string) (string, error), family, table string) (string, bool, error) { + return readNftObject(run, "-a", "list", "table", family, table) +} + +func ParseTableChains(output string) map[string]string { + chains := make(map[string]string) + lines := strings.Split(output, "\n") + name, indent, start := "", "", 0 + for index, line := range lines { + trimmed := strings.TrimSpace(line) + if name == "" { + fields := strings.Fields(trimmed) + if len(fields) >= 3 && fields[0] == "chain" && fields[2] == "{" { + name, indent, start = fields[1], line[:len(line)-len(strings.TrimLeft(line, " \t"))], index + } + continue + } + if strings.HasPrefix(line, indent+"}") { + chains[name] = strings.Join(lines[start:index+1], "\n") + name = "" + } + } + return chains +} + var ErrChainNotFound = errors.New("nftables chain is not initialized") func ReadChain(run func(...string) (string, error), family, table, chain string) (string, error) { diff --git a/agent/utils/firewall/sync/diff.go b/agent/utils/firewall/sync/diff.go index cef0ca0cdefd..2360af79bb93 100644 --- a/agent/utils/firewall/sync/diff.go +++ b/agent/utils/firewall/sync/diff.go @@ -19,96 +19,3 @@ const ( ReasonUnsafeRemoval ReasonCode = "unsafe_managed_rule_removal" ReasonReadOnlyRule ReasonCode = "read_only_rule" ) - -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" - case ReasonReadOnlyRule: - return "read-only runtime rule is preserved but cannot be synchronized" - 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 deleted file mode 100644 index 5d9d3254110e..000000000000 --- a/agent/utils/firewall/sync/order.go +++ /dev/null @@ -1,136 +0,0 @@ -package sync - -import ( - "slices" - "strconv" - "strings" - - "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter" -) - -func RuleOrder(snapshot filter.Snapshot, ordered []filter.InventoryItem) map[string]bool { - drifted := make(map[string]bool) - if snapshot.Scope.Provider == filter.ProviderFirewalld { - return drifted - } - markers := make([]string, 0, len(ordered)) - projected := append([]filter.ObservedRule(nil), snapshot.Rules...) - markerSet := make(map[string]bool, len(ordered)) - locations := make(map[string][]int, len(projected)) - for index, rule := range projected { - key := locatorBucket(rule.Locator) - locations[key] = append(locations[key], index) - } - for _, item := range ordered { - if item.Desired == nil || item.Desired.Marker == "" { - continue - } - markers = append(markers, item.Desired.Marker) - markerSet[item.Desired.Marker] = true - if item.Observed != nil { - for _, index := range locations[locatorBucket(item.Observed.Locator)] { - if filter.SameLocator(projected[index].Locator, item.Observed.Locator) { - projected[index].Marker = item.Desired.Marker - break - } - } - } - } - actual := make([]string, 0, len(markers)) - presentMarkers := make(map[string]bool, len(markers)) - for _, rule := range projected { - if markerSet[rule.Marker] { - actual = append(actual, rule.Marker) - presentMarkers[rule.Marker] = true - } - } - present := make([]string, 0, len(actual)) - for _, marker := range markers { - if presentMarkers[marker] { - present = append(present, marker) - } - } - for index, marker := range present { - if actual[index] != marker { - drifted[marker] = true - drifted[actual[index]] = true - } - } - return drifted -} - -func InsertionPosition(snapshot filter.Snapshot, markers []string, target string) *int64 { - if snapshot.Scope.Provider == filter.ProviderFirewalld { - return nil - } - positions := make(map[string]int, len(snapshot.Rules)) - for _, rule := range snapshot.Rules { - if _, exists := positions[rule.Marker]; !exists && rule.Locator.Position != nil { - positions[rule.Marker] = *rule.Locator.Position - } - } - targetIndex := slices.Index(markers, target) - for index := targetIndex - 1; index >= 0; index-- { - if previous, ok := positions[markers[index]]; ok { - position := int64(previous + 1) - return &position - } - } - if targetIndex >= 0 { - for _, marker := range markers[targetIndex+1:] { - if next, ok := positions[marker]; ok { - position := int64(next) - return &position - } - } - } - return nil -} - -func DeleteChange(snapshot filter.Snapshot, previous filter.ObservedRule, desired filter.DesiredRule) (filter.DesiredChange, error) { - beforeKey, err := filter.RuleKey(previous.Rule) - if err != nil { - return filter.DesiredChange{}, err - } - matches := make([]filter.ObservedRule, 0, 1) - for _, current := range snapshot.Rules { - if previous.Marker != "" { - if current.Marker != previous.Marker { - continue - } - } else if current.Marker != "" || !filter.SameLocator(current.Locator, previous.Locator) { - continue - } - key, err := filter.RuleKey(current.Rule) - if err == nil && key == beforeKey { - matches = append(matches, current) - } - } - if len(matches) != 1 { - return filter.DesiredChange{}, filter.ErrRuleStale - } - current := matches[0] - if err := filter.GuardMutation(current); err != nil { - return filter.DesiredChange{}, err - } - before := ObservedRule(current) - if before.UUID == "" { - before.UUID = desired.Rule.UUID - } - return filter.DesiredChange{Operation: filter.ChangeDelete, Before: &before, Locator: ¤t.Locator, UnmarkedAdopted: current.Marker == "" && desired.Origin == filter.RuleOriginAdopted}, nil -} - -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 -} - -func locatorBucket(locator filter.Locator) string { - if locator.Position != nil { - return locator.ScopeKey + "\x00position:" + strconv.Itoa(*locator.Position) - } - return locator.ScopeKey + "\x00canonical:" + locator.Canonical -} diff --git a/frontend/src/api/interface/firewall.ts b/frontend/src/api/interface/firewall.ts index 644e38c7540e..beec8786813e 100644 --- a/frontend/src/api/interface/firewall.ts +++ b/frontend/src/api/interface/firewall.ts @@ -22,6 +22,7 @@ export namespace Firewall { initialized: boolean; bound: boolean; reason?: string; + forwardPolicy?: 'ACCEPT' | 'DROP'; } export interface BackendGroup { selected: string; @@ -243,7 +244,9 @@ export namespace Firewall { export interface AdoptRequest { scope: Scope; - instanceKey: string; + instanceKey?: string; + rule?: Rule; + marker?: string; } export interface CreateItem { diff --git a/frontend/src/lang/modules/en.ts b/frontend/src/lang/modules/en.ts index e71b69a96f96..428a587cd569 100644 --- a/frontend/src/lang/modules/en.ts +++ b/frontend/src/lang/modules/en.ts @@ -4204,6 +4204,8 @@ const message = { dockerGuard: 'Container Port Guard', systemFirewall: 'Host firewall', systemFirewallHelper: 'Controls host port access and inbound firewall rules.', + forwardPolicyDropWarning: + 'The default FORWARD policy for {0} is DROP. Traffic forwarded to other hosts may be blocked unless explicitly allowed by iptables/ip6tables rules. Check the allow rules for the corresponding IP version.', forwardingHelper: 'Manages port-forwarding rules.', dockerFirewallHelper: 'Selects how 1Panel manages container port protection.', dockerNftablesRequirement: 'Docker ≥ 29.0.0, experimental', @@ -4211,7 +4213,8 @@ const message = { configuredRules: '{0} rules configured', addressFamily: 'IP version', portOrRange: 'Port / range', - importBackendHelper: 'Imported rules are converted for the current {0} backend. Source rules are not changed.', + batchLimit: 'Create or import at most {0} rules per batch (after expansion).', + importLimit: 'Import up to {0} rules (counted after expansion), with a file size of at most {1} KB.', ruleSyncTitle: 'Synchronize rules', ruleSyncAction: 'Sync rules', ruleSyncDatabase: '1Panel database', diff --git a/frontend/src/lang/modules/es-es.ts b/frontend/src/lang/modules/es-es.ts index 78ef441c00fd..e226cb875e2d 100644 --- a/frontend/src/lang/modules/es-es.ts +++ b/frontend/src/lang/modules/es-es.ts @@ -4251,6 +4251,8 @@ const message = { dockerGuard: 'Protección de puertos de contenedores', systemFirewall: 'Firewall del host', systemFirewallHelper: 'Controla el acceso a los puertos del host y las reglas de entrada.', + forwardPolicyDropWarning: + 'La política predeterminada de FORWARD para {0} es DROP. El tráfico reenviado a otros hosts puede bloquearse si las reglas de iptables/ip6tables no lo permiten explícitamente. Revise las reglas de permiso de la versión IP correspondiente.', forwardingHelper: 'Gestiona las reglas de reenvío de puertos.', dockerFirewallHelper: 'Selecciona cómo gestiona 1Panel la protección de puertos de contenedores.', dockerNftablesRequirement: 'Docker ≥ 29.0.0, experimental', @@ -4258,8 +4260,8 @@ const message = { configuredRules: '{0} reglas configuradas', addressFamily: 'Versión de IP', portOrRange: 'Puerto / rango', - importBackendHelper: - 'Las reglas importadas se convierten para el backend actual {0}. Las reglas de origen no se modifican.', + batchLimit: 'Puede crear o importar hasta {0} reglas por lote, después de expandirlas.', + importLimit: 'Importe hasta {0} reglas (contadas después de expandirlas), con un archivo de hasta {1} KB.', ruleSyncTitle: 'Sincronizar reglas', ruleSyncAction: 'Sincronizar reglas', ruleSyncDatabase: 'Base de datos de 1Panel', diff --git a/frontend/src/lang/modules/fa.ts b/frontend/src/lang/modules/fa.ts index 9e410d382911..2ed484008d34 100644 --- a/frontend/src/lang/modules/fa.ts +++ b/frontend/src/lang/modules/fa.ts @@ -4159,6 +4159,8 @@ const message = { dockerGuard: 'محافظت از پورت کانتینر', systemFirewall: 'فایروال میزبان', systemFirewallHelper: 'دسترسی پورت‌های میزبان و قوانین ورودی را کنترل می‌کند.', + forwardPolicyDropWarning: + 'سیاست پیش‌فرض FORWARD برای {0} برابر DROP است. ترافیک هدایت‌شده به میزبان‌های دیگر ممکن است مسدود شود، مگر اینکه قواعد iptables/ip6tables صریحاً آن را مجاز کنند. قواعد مجازکننده نسخه IP مربوطه را بررسی کنید.', forwardingHelper: 'قوانین انتقال پورت را مدیریت می‌کند.', dockerFirewallHelper: 'روش مدیریت محافظت از پورت کانتینر در 1Panel را انتخاب می‌کند.', dockerNftablesRequirement: 'Docker ≥ 29.0.0، آزمایشی', @@ -4166,7 +4168,8 @@ const message = { configuredRules: '{0} قانون پیکربندی شده', addressFamily: 'نسخه IP', portOrRange: 'پورت / بازه', - importBackendHelper: 'قوانین واردشده برای بک‌اند فعلی {0} تبدیل می‌شوند. قوانین مبدأ تغییر نمی‌کنند.', + batchLimit: 'در هر نوبت حداکثر {0} قانون پس از گسترش ایجاد یا وارد کنید.', + importLimit: 'حداکثر {0} قانون (پس از گسترش) وارد کنید؛ حجم فایل نباید از {1} کیلوبایت بیشتر باشد.', ruleSyncTitle: 'همگام‌سازی قوانین', ruleSyncAction: 'همگام‌سازی قوانین', ruleSyncDatabase: 'پایگاه داده 1Panel', diff --git a/frontend/src/lang/modules/ja.ts b/frontend/src/lang/modules/ja.ts index 112dc5487dd6..ff21ac0fe37f 100644 --- a/frontend/src/lang/modules/ja.ts +++ b/frontend/src/lang/modules/ja.ts @@ -4186,6 +4186,8 @@ const message = { dockerGuard: 'コンテナポート保護', systemFirewall: 'ホストファイアウォール', systemFirewallHelper: 'ホストのポートアクセスと受信ルールを管理します。', + forwardPolicyDropWarning: + '{0} の FORWARD のデフォルトポリシーは DROP です。iptables/ip6tables のルールで明示的に許可されていない他のホストへの転送トラフィックは、遮断される可能性があります。該当する IP バージョンの許可ルールを確認してください。', forwardingHelper: 'ポート転送ルールを管理します。', dockerFirewallHelper: '1Panel のコンテナポート保護の管理方法を選択します。', dockerNftablesRequirement: 'Docker ≥ 29.0.0、実験的', @@ -4193,8 +4195,8 @@ const message = { configuredRules: '{0} 件設定済み', addressFamily: 'IP バージョン', portOrRange: 'ポート / 範囲', - importBackendHelper: - 'インポートしたルールは現在の {0} バックエンド向けに変換されます。移行元のルールは変更されません。', + batchLimit: '一度に作成またはインポートできるルールは、展開後の件数で最大 {0} 件です。', + importLimit: '最大 {0} 件のルール(展開後の件数)をインポートできます。ファイルサイズの上限は {1} KB です。', ruleSyncTitle: 'ルールを同期', ruleSyncAction: 'ルールを同期', ruleSyncDatabase: '1Panel データベース', diff --git a/frontend/src/lang/modules/ko.ts b/frontend/src/lang/modules/ko.ts index 3578e408293f..70c66c2c34a2 100644 --- a/frontend/src/lang/modules/ko.ts +++ b/frontend/src/lang/modules/ko.ts @@ -4109,6 +4109,8 @@ const message = { dockerGuard: '컨테이너 포트 보호', systemFirewall: '호스트 방화벽', systemFirewallHelper: '호스트 포트 접근과 인바운드 규칙을 관리합니다.', + forwardPolicyDropWarning: + '{0}의 FORWARD 기본 정책은 DROP입니다. iptables/ip6tables 규칙에서 명시적으로 허용하지 않은 다른 호스트로의 전달 트래픽은 차단될 수 있습니다. 해당 IP 버전의 허용 규칙을 확인하세요.', forwardingHelper: '포트 포워딩 규칙을 관리합니다.', dockerFirewallHelper: '1Panel 컨테이너 포트 보호 관리 방식을 선택합니다.', dockerNftablesRequirement: 'Docker ≥ 29.0.0, 실험적', @@ -4116,7 +4118,8 @@ const message = { configuredRules: '{0}개 규칙 설정됨', addressFamily: 'IP 버전', portOrRange: '포트 / 범위', - importBackendHelper: '가져온 규칙은 현재 {0} 백엔드에 맞게 변환됩니다. 원본 백엔드의 규칙은 변경되지 않습니다.', + batchLimit: '한 번에 생성하거나 가져올 수 있는 규칙은 확장 후 기준으로 최대 {0}개입니다.', + importLimit: '최대 {0}개의 규칙(확장 후 기준)을 가져올 수 있으며, 파일 크기는 {1} KB를 초과할 수 없습니다.', ruleSyncTitle: '규칙 동기화', ruleSyncAction: '규칙 동기화', ruleSyncDatabase: '1Panel 데이터베이스', diff --git a/frontend/src/lang/modules/lo.ts b/frontend/src/lang/modules/lo.ts index 1d5b2b5a58f8..b2f68fe903f1 100644 --- a/frontend/src/lang/modules/lo.ts +++ b/frontend/src/lang/modules/lo.ts @@ -4076,6 +4076,8 @@ const message = { dockerGuard: 'ການປ້ອງກັນພອດຄອນເທນເນີ', systemFirewall: 'ໄຟວໍໂຮສ', systemFirewallHelper: 'ຄວບຄຸມການເຂົ້າເຖິງພອດໂຮສ ແລະ ກົດຂາເຂົ້າ.', + forwardPolicyDropWarning: + 'ນະໂຍບາຍ FORWARD ເລີ່ມຕົ້ນສຳລັບ {0} ແມ່ນ DROP. ການຈະລາຈອນທີ່ສົ່ງຕໍ່ໄປຍັງໂຮສອື່ນອາດຖືກບລັອກ ຖ້າບໍ່ມີກົດ iptables/ip6tables ອະນຸຍາດຢ່າງຊັດເຈນ. ກະລຸນາກວດສອບກົດອະນຸຍາດສຳລັບເວີຊັນ IP ທີ່ກ່ຽວຂ້ອງ.', forwardingHelper: 'ຈັດການກົດການສົ່ງຕໍ່ພອດ.', dockerFirewallHelper: 'ເລືອກວິທີທີ່ 1Panel ຈັດການການປ້ອງກັນພອດຄອນເທນເນີ.', dockerNftablesRequirement: 'Docker ≥ 29.0.0, ທົດລອງ', @@ -4083,7 +4085,8 @@ const message = { configuredRules: 'ຕັ້ງຄ່າແລ້ວ {0} ກົດ', addressFamily: 'ລຸ້ນ IP', portOrRange: 'ພອດ / ຊ່ວງ', - importBackendHelper: 'ກົດທີ່ນຳເຂົ້າຈະຖືກປ່ຽນໃຫ້ເໝາະກັບແບັກເອນ {0} ປັດຈຸບັນ. ກົດຕົ້ນທາງຈະບໍ່ຖືກປ່ຽນ.', + batchLimit: 'ແຕ່ລະຄັ້ງສາມາດສ້າງ ຫຼື ນຳເຂົ້າໄດ້ສູງສຸດ {0} ກົດ ໂດຍນັບຫຼັງຈາກຂະຫຍາຍແລ້ວ.', + importLimit: 'ນຳເຂົ້າໄດ້ສູງສຸດ {0} ກົດ (ນັບຫຼັງຈາກຂະຫຍາຍແລ້ວ), ຂະໜາດໄຟລ໌ບໍ່ເກີນ {1} KB.', ruleSyncTitle: 'ຊິງກົດ', ruleSyncAction: 'ຊິງກົດ', ruleSyncDatabase: 'ຖານຂໍ້ມູນ 1Panel', diff --git a/frontend/src/lang/modules/ms.ts b/frontend/src/lang/modules/ms.ts index 51935ab382ad..daaf32f7aa97 100644 --- a/frontend/src/lang/modules/ms.ts +++ b/frontend/src/lang/modules/ms.ts @@ -4273,6 +4273,8 @@ const message = { dockerGuard: 'Perlindungan port bekas', systemFirewall: 'Tembok api hos', systemFirewallHelper: 'Mengawal akses port hos dan peraturan masuk.', + forwardPolicyDropWarning: + 'Dasar FORWARD lalai untuk {0} ialah DROP. Trafik yang dimajukan ke hos lain mungkin disekat melainkan dibenarkan secara jelas oleh peraturan iptables/ip6tables. Semak peraturan kebenaran untuk versi IP yang berkenaan.', forwardingHelper: 'Mengurus peraturan pemajuan port.', dockerFirewallHelper: 'Memilih cara 1Panel mengurus perlindungan port bekas.', dockerNftablesRequirement: 'Docker ≥ 29.0.0, percubaan', @@ -4280,8 +4282,9 @@ const message = { configuredRules: '{0} peraturan dikonfigurasi', addressFamily: 'Versi IP', portOrRange: 'Port / julat', - importBackendHelper: - 'Peraturan yang diimport ditukar untuk bahagian belakang semasa {0}. Peraturan sumber tidak diubah.', + batchLimit: 'Cipta atau import maksimum {0} peraturan setiap kelompok selepas pengembangan.', + importLimit: + 'Import sehingga {0} peraturan (dikira selepas pengembangan), dengan saiz fail tidak melebihi {1} KB.', ruleSyncTitle: 'Segerakkan peraturan', ruleSyncAction: 'Segerakkan peraturan', ruleSyncDatabase: 'Pangkalan data 1Panel', diff --git a/frontend/src/lang/modules/pt-br.ts b/frontend/src/lang/modules/pt-br.ts index 90e81dff3a47..ab1573a04567 100644 --- a/frontend/src/lang/modules/pt-br.ts +++ b/frontend/src/lang/modules/pt-br.ts @@ -4289,6 +4289,8 @@ const message = { dockerGuard: 'Proteção de portas de contêineres', systemFirewall: 'Firewall do host', systemFirewallHelper: 'Controla o acesso às portas do host e as regras de entrada.', + forwardPolicyDropWarning: + 'A política padrão de FORWARD para {0} é DROP. O tráfego encaminhado para outros hosts pode ser bloqueado se não for permitido explicitamente pelas regras do iptables/ip6tables. Verifique as regras de permissão da versão IP correspondente.', forwardingHelper: 'Gerencia regras de encaminhamento de portas.', dockerFirewallHelper: 'Seleciona como o 1Panel gerencia a proteção de portas de contêineres.', dockerNftablesRequirement: 'Docker ≥ 29.0.0, experimental', @@ -4296,8 +4298,8 @@ const message = { configuredRules: '{0} regras configuradas', addressFamily: 'Versão do IP', portOrRange: 'Porta / intervalo', - importBackendHelper: - 'As regras importadas são convertidas para o backend atual {0}. As regras de origem não são alteradas.', + batchLimit: 'Crie ou importe no máximo {0} regras por lote, após a expansão.', + importLimit: 'Importe até {0} regras (contadas após a expansão), com um arquivo de no máximo {1} KB.', ruleSyncTitle: 'Sincronizar regras', ruleSyncAction: 'Sincronizar regras', ruleSyncDatabase: 'Banco de dados do 1Panel', diff --git a/frontend/src/lang/modules/ru.ts b/frontend/src/lang/modules/ru.ts index d761cd20fc46..a28525ad951f 100644 --- a/frontend/src/lang/modules/ru.ts +++ b/frontend/src/lang/modules/ru.ts @@ -4257,6 +4257,8 @@ const message = { dockerGuard: 'Защита портов контейнеров', systemFirewall: 'Брандмауэр хоста', systemFirewallHelper: 'Управляет доступом к портам хоста и входящими правилами.', + forwardPolicyDropWarning: + 'Политика FORWARD по умолчанию для {0} — DROP. Трафик, пересылаемый на другие узлы, может блокироваться, если он явно не разрешён правилами iptables/ip6tables. Проверьте разрешающие правила для соответствующей версии IP.', forwardingHelper: 'Управляет правилами перенаправления портов.', dockerFirewallHelper: 'Выбирает способ управления защитой портов контейнеров в 1Panel.', dockerNftablesRequirement: 'Docker ≥ 29.0.0, экспериментально', @@ -4264,8 +4266,9 @@ const message = { configuredRules: 'Настроено правил: {0}', addressFamily: 'Версия IP', portOrRange: 'Порт / диапазон', - importBackendHelper: - 'Импортируемые правила преобразуются для текущего бэкенда {0}. Исходные правила не изменяются.', + batchLimit: 'За один раз можно создать или импортировать не более {0} правил после развёртывания.', + importLimit: + 'Можно импортировать до {0} правил (после развёртывания). Размер файла не должен превышать {1} КБ.', ruleSyncTitle: 'Синхронизация правил', ruleSyncAction: 'Синхронизировать правила', ruleSyncDatabase: 'База данных 1Panel', diff --git a/frontend/src/lang/modules/tr.ts b/frontend/src/lang/modules/tr.ts index 50cc8c6fbfa7..982222a41de6 100644 --- a/frontend/src/lang/modules/tr.ts +++ b/frontend/src/lang/modules/tr.ts @@ -4274,6 +4274,8 @@ const message = { dockerGuard: 'Konteyner portu koruması', systemFirewall: 'Ana makine güvenlik duvarı', systemFirewallHelper: 'Ana makine port erişimini ve gelen kuralları yönetir.', + forwardPolicyDropWarning: + '{0} için varsayılan FORWARD ilkesi DROP olarak ayarlanmış. Diğer ana bilgisayarlara yönlendirilen trafik, iptables/ip6tables kurallarıyla açıkça izin verilmediği sürece engellenebilir. İlgili IP sürümünün izin kurallarını kontrol edin.', forwardingHelper: 'Port yönlendirme kurallarını yönetir.', dockerFirewallHelper: '1Panel konteyner portu korumasının nasıl yönetileceğini seçer.', dockerNftablesRequirement: 'Docker ≥ 29.0.0, deneysel', @@ -4281,8 +4283,9 @@ const message = { configuredRules: '{0} kural yapılandırıldı', addressFamily: 'IP sürümü', portOrRange: 'Port / aralık', - importBackendHelper: - 'İçe aktarılan kurallar geçerli {0} arka ucu için dönüştürülür. Kaynak kurallar değiştirilmez.', + batchLimit: 'Genişletme sonrasında bir defada en fazla {0} kural oluşturabilir veya içe aktarabilirsiniz.', + importLimit: + 'En fazla {0} kural (genişletme sonrası sayıya göre) içe aktarılabilir; dosya boyutu {1} KB değerini aşmamalıdır.', ruleSyncTitle: 'Kuralları eşitle', ruleSyncAction: 'Kuralları eşitle', ruleSyncDatabase: '1Panel veritabanı', diff --git a/frontend/src/lang/modules/zh-Hant.ts b/frontend/src/lang/modules/zh-Hant.ts index 53a5c5b15049..0325050839ed 100644 --- a/frontend/src/lang/modules/zh-Hant.ts +++ b/frontend/src/lang/modules/zh-Hant.ts @@ -3921,6 +3921,8 @@ const message = { dockerGuard: '容器連接埠防護', systemFirewall: '主機防火牆', systemFirewallHelper: '用於管理主機連接埠存取和入站規則。', + forwardPolicyDropWarning: + '{0} 的 FORWARD 預設策略為 DROP,未被 iptables/ip6tables 規則明確允許的跨主機轉送流量可能遭到阻擋,請檢查對應位址族的允許規則。', forwardingHelper: '用於管理連接埠轉發規則。', dockerFirewallHelper: '用於選擇 1Panel 容器連接埠防護的管理方式。', dockerNftablesRequirement: 'Docker ≥ 29.0.0,實驗性', @@ -3928,7 +3930,8 @@ const message = { configuredRules: '已設定 {0} 條規則', addressFamily: 'IP 版本', portOrRange: '連接埠 / 範圍', - importBackendHelper: '匯入規則會轉換並寫入目前後端 {0},來源後端規則不會被修改。', + batchLimit: '每次最多建立或匯入 {0} 條規則(依展開後的數量計算)。', + importLimit: '最多匯入 {0} 條規則(依展開後的數量計算),檔案大小不得超過 {1} KB。', ruleSyncTitle: '同步規則', ruleSyncAction: '同步規則', ruleSyncDatabase: '1Panel 資料庫', diff --git a/frontend/src/lang/modules/zh.ts b/frontend/src/lang/modules/zh.ts index 0bc50cb17faf..c8c3cf82ad92 100644 --- a/frontend/src/lang/modules/zh.ts +++ b/frontend/src/lang/modules/zh.ts @@ -3969,6 +3969,8 @@ const message = { dockerGuard: '容器端口防护', systemFirewall: '主机防火墙', systemFirewallHelper: '用于管理主机端口访问和入站规则。', + forwardPolicyDropWarning: + '{0} 的 FORWARD 默认策略为 DROP,未被 iptables/ip6tables 规则明确放行的跨主机转发流量可能被阻断,请检查对应地址族的放行规则。', forwardingHelper: '用于管理端口转发规则。', dockerFirewallHelper: '用于选择 1Panel 容器端口防护的管理方式。', dockerNftablesRequirement: 'Docker ≥ 29.0.0,实验性', @@ -3976,7 +3978,8 @@ const message = { configuredRules: '已配置 {0} 条规则', addressFamily: 'IP 版本', portOrRange: '端口 / 范围', - importBackendHelper: '导入规则会转换并写入当前后端 {0},源后端规则不会被修改。', + batchLimit: '每次最多创建或导入 {0} 条规则(按展开后的数量计算)。', + importLimit: '最多导入 {0} 条规则(按展开后的数量计算),文件大小不超过 {1} KB。', ruleSyncTitle: '同步规则', ruleSyncAction: '同步规则', ruleSyncDatabase: '1Panel 数据库', diff --git a/frontend/src/views/host/firewall/docker/detail/index.vue b/frontend/src/views/host/firewall/docker/detail/index.vue index 7df53fb65fc4..df82f39aa785 100644 --- a/frontend/src/views/host/firewall/docker/detail/index.vue +++ b/frontend/src/views/host/firewall/docker/detail/index.vue @@ -48,14 +48,14 @@ v-for="group in filteredPortGroups" :key="group.key" class="port-detail-card" - :class="{ 'is-selected': selectedGroupKeys.includes(group.key) }" + :class="{ 'is-selected': selectedGroupKeys.has(group.key) }" shadow="never" @click="toggleSelection(group.key)" >
@@ -188,7 +188,12 @@ import { dockerGuardManagementTarget, isValidDockerGuardSource, } from '@/views/host/firewall/docker/model'; -import { formatHostAddress, formatHostAddressList, splitTagValues } from '@/views/host/firewall/utils/validation'; +import { + FIREWALL_BATCH_LIMIT, + formatHostAddress, + formatHostAddressList, + splitTagValues, +} from '@/views/host/firewall/utils/validation'; const props = defineProps<{ base: Firewall.DockerGuardBase; containers: Firewall.DockerGuardContainer[] }>(); const emit = defineEmits<{ search: []; created: [taskID: string] }>(); @@ -197,7 +202,7 @@ const drawerVisible = ref(false); const policyVisible = ref(false); const savingPolicy = ref(false); const activeContainerKey = ref(''); -const selectedGroupKeys = ref([]); +const selectedGroupKeys = ref(new Set()); const policyEndpoints = ref([]); const familyFilter = ref<'all' | Firewall.DockerGuardEndpoint['family']>('all'); const formRef = ref(); @@ -219,14 +224,14 @@ const filteredPortGroups = computed(() => ); const selectedEndpoints = computed(() => filteredPortGroups.value - .filter((group) => selectedGroupKeys.value.includes(group.key)) + .filter((group) => selectedGroupKeys.value.has(group.key)) .flatMap((group) => group.endpoints), ); const allSelected = computed( - () => filteredPortGroups.value.length > 0 && selectedGroupKeys.value.length === filteredPortGroups.value.length, + () => filteredPortGroups.value.length > 0 && selectedGroupKeys.value.size === filteredPortGroups.value.length, ); const selectionIndeterminate = computed( - () => selectedGroupKeys.value.length > 0 && selectedGroupKeys.value.length < filteredPortGroups.value.length, + () => selectedGroupKeys.value.size > 0 && selectedGroupKeys.value.size < filteredPortGroups.value.length, ); const hasMixedFamilies = computed(() => new Set(policyEndpoints.value.map((endpoint) => endpoint.family)).size > 1); const normalizeSources = (sources: string[]) => @@ -281,22 +286,24 @@ const rules = reactive({ const acceptParams = (container: Firewall.DockerGuardContainer) => { activeContainerKey.value = container.key; familyFilter.value = 'all'; - selectedGroupKeys.value = []; + selectedGroupKeys.value.clear(); drawerVisible.value = true; }; const changeFamilyFilter = () => { - selectedGroupKeys.value = []; + selectedGroupKeys.value.clear(); }; const changeAllSelection = (checked: boolean) => { - selectedGroupKeys.value = checked ? filteredPortGroups.value.map((group) => group.key) : []; + selectedGroupKeys.value = new Set(checked ? filteredPortGroups.value.map((group) => group.key) : []); }; const changeSelection = (key: string, checked: boolean) => { - selectedGroupKeys.value = checked - ? [...selectedGroupKeys.value, key] - : selectedGroupKeys.value.filter((item) => item !== key); + if (checked) { + selectedGroupKeys.value.add(key); + } else { + selectedGroupKeys.value.delete(key); + } }; const toggleSelection = (key: string) => { - changeSelection(key, !selectedGroupKeys.value.includes(key)); + changeSelection(key, !selectedGroupKeys.value.has(key)); }; const openPolicy = (endpoints: Firewall.DockerGuardEndpoint[]) => { if (!endpoints.length) return; @@ -350,6 +357,10 @@ const submitPolicy = async () => { if (!valid || !form.mode) return; const mode = form.mode; const sources = mode === 'deny_all' ? [] : splitTagValues(form.sources); + if (policyEndpoints.value.length > FIREWALL_BATCH_LIMIT) { + MsgError(i18n.global.t('firewall.batchLimit', [FIREWALL_BATCH_LIMIT])); + return; + } savingPolicy.value = true; try { const result = ( @@ -371,7 +382,7 @@ const submitPolicy = async () => { } policyVisible.value = false; drawerVisible.value = false; - selectedGroupKeys.value = []; + selectedGroupKeys.value.clear(); emit('created', result.taskID); } catch (error) { MsgError( @@ -404,7 +415,7 @@ const remove = async (endpoints: Firewall.DockerGuardEndpoint[], batch: boolean) MsgError(i18n.global.t('commons.msg.operationFailed')); return; } - selectedGroupKeys.value = []; + selectedGroupKeys.value.clear(); drawerVisible.value = false; emit('created', result.taskID); } catch (error) { diff --git a/frontend/src/views/host/firewall/docker/import/index.vue b/frontend/src/views/host/firewall/docker/import/index.vue index 2ac17804e724..b160b73ac32c 100644 --- a/frontend/src/views/host/firewall/docker/import/index.vue +++ b/frontend/src/views/host/firewall/docker/import/index.vue @@ -1,6 +1,11 @@ - - + + + + + @@ -50,7 +69,7 @@ @@ -67,14 +86,18 @@ import { isAxiosError } from 'axios'; import { genFileId, type UploadFile, type UploadFiles, type UploadProps, type UploadRawFile } from 'element-plus'; import { ref } from 'vue'; import { dockerGuardEndpointKey, normalizeDockerGuardPolicy } from '@/views/host/firewall/docker/model'; -import { formatHostAddressList } from '@/views/host/firewall/utils/validation'; +import { + FIREWALL_BATCH_LIMIT, + FIREWALL_IMPORT_MAX_SIZE, + formatHostAddressList, +} from '@/views/host/firewall/utils/validation'; import { Document } from '@element-plus/icons-vue'; const emit = defineEmits<{ (event: 'created', taskID: string): void }>(); const visible = ref(false); const loading = ref(false); const policies = ref([]); -const selects = ref([]); +const selects = ref(new Set()); const uploadRef = ref(); const uploaderFiles = ref([]); const submitError = ref(''); @@ -83,15 +106,29 @@ const displaySources = (policy: Firewall.DockerGuardPolicy) => formatHostAddress const fileOnChange = (uploadFile: UploadFile, uploadFiles: UploadFiles) => { if (!uploadFile.raw) return; loading.value = true; + policies.value = []; - selects.value = []; + selects.value = new Set(); submitError.value = ''; uploaderFiles.value = uploadFiles; + if (uploadFile.raw.size > FIREWALL_IMPORT_MAX_SIZE) { + uploadRef.value?.clearFiles(); + uploaderFiles.value = []; + loading.value = false; + MsgError(i18n.global.t('firewall.importLimit', [FIREWALL_BATCH_LIMIT, FIREWALL_IMPORT_MAX_SIZE / 1024])); + return; + } const reader = new FileReader(); reader.onload = (event) => { try { const parsed: unknown = JSON.parse(String(event.target?.result || '')); if (!Array.isArray(parsed)) throw new Error(); + if (parsed.length > FIREWALL_BATCH_LIMIT) { + MsgError( + i18n.global.t('firewall.importLimit', [FIREWALL_BATCH_LIMIT, FIREWALL_IMPORT_MAX_SIZE / 1024]), + ); + return; + } const normalized = parsed.map(normalizeDockerGuardPolicy); if (normalized.some((policy) => !policy)) throw new Error(); const byEndpoint = new Map(); @@ -99,10 +136,10 @@ const fileOnChange = (uploadFile: UploadFile, uploadFiles: UploadFiles) => { byEndpoint.set(dockerGuardEndpointKey(policy), policy); } policies.value = [...byEndpoint.values()]; - selects.value = [...policies.value]; + selects.value = new Set(policies.value); } catch { policies.value = []; - selects.value = []; + selects.value = new Set(); MsgError(i18n.global.t('commons.msg.errImportFormat')); } finally { loading.value = false; @@ -119,11 +156,19 @@ const handleExceed: UploadProps['onExceed'] = (files) => { }; const onImport = async () => { - if (loading.value || selects.value.length === 0) return; + if (loading.value || selects.value.size === 0) return; + if (selects.value.size > FIREWALL_BATCH_LIMIT) { + MsgError(i18n.global.t('firewall.importLimit', [FIREWALL_BATCH_LIMIT, FIREWALL_IMPORT_MAX_SIZE / 1024])); + return; + } loading.value = true; submitError.value = ''; try { - const result = (await upsertDockerPortGuardPolicies({ policies: selects.value })).data; + const result = ( + await upsertDockerPortGuardPolicies({ + policies: policies.value.filter((policy) => selects.value.has(policy)), + }) + ).data; if (!result.taskID || !result.queued) { submitError.value = i18n.global.t('commons.msg.operationFailed'); return; @@ -148,8 +193,9 @@ const modeLabel = (mode: Firewall.DockerGuardPolicy['mode']) => { const acceptParams = () => { loading.value = false; + policies.value = []; - selects.value = []; + selects.value = new Set(); uploaderFiles.value = []; submitError.value = ''; uploadRef.value?.clearFiles(); diff --git a/frontend/src/views/host/firewall/docker/index.vue b/frontend/src/views/host/firewall/docker/index.vue index 5404c86e2234..7ce4c684308b 100644 --- a/frontend/src/views/host/firewall/docker/index.vue +++ b/frontend/src/views/host/firewall/docker/index.vue @@ -55,7 +55,13 @@