refactor: simplify firewall service structure (#13646)

This commit is contained in:
ssongliu
2026-08-27 14:00:43 +08:00
committed by GitHub
parent ddfb816ef1
commit 53f75826d8
75 changed files with 1862 additions and 12696 deletions
@@ -1,49 +0,0 @@
package v2
import (
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/dto"
"github.com/1Panel-dev/1Panel/agent/app/service"
agenti18n "github.com/1Panel-dev/1Panel/agent/i18n"
"github.com/gin-gonic/gin"
)
func TestHandleDockerPortGuardErrorReturnsStableBusinessCode(t *testing.T) {
agenti18n.Init()
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
context, _ := gin.CreateTestContext(recorder)
handleDockerPortGuardError(context, fmt.Errorf("normalize policy: %w", service.ErrDockerGuardInvalid))
var response dto.Response
if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
t.Fatalf("decode response: %v", err)
}
if response.Code != http.StatusBadRequest || response.ErrorCode != "FW_DOCKER_GUARD_INVALID" {
t.Fatalf("unexpected Docker guard error response: %#v", response)
}
}
func TestHandleDockerPortGuardErrorLocalizesDockerUnavailable(t *testing.T) {
agenti18n.Init()
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
context, _ := gin.CreateTestContext(recorder)
handleDockerPortGuardError(context, fmt.Errorf("inspect Docker: %w", service.ErrDockerUnavailable))
var response dto.Response
if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
t.Fatalf("decode response: %v", err)
}
if response.Code != http.StatusServiceUnavailable || response.ErrorCode != "FW_DOCKER_UNAVAILABLE" {
t.Fatalf("unexpected Docker unavailable response: %#v", response)
}
if response.Message != agenti18n.Get("ErrDockerFailed") {
t.Fatalf("message = %q, want localized Docker failure", response.Message)
}
}
-70
View File
@@ -1,70 +0,0 @@
package v2
import (
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/dto"
"github.com/1Panel-dev/1Panel/agent/app/repo"
agenti18n "github.com/1Panel-dev/1Panel/agent/i18n"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
"github.com/gin-gonic/gin"
)
func TestHandleFirewallRuleErrorReturnsStableBusinessCode(t *testing.T) {
agenti18n.Init()
gin.SetMode(gin.TestMode)
tests := []struct {
name string
err error
code string
}{
{name: "stale rule", err: filter.ErrRuleStale, code: "FW_RULE_STALE"},
{name: "check required", err: filter.ErrRuleCheckRequired, code: "FW_RULE_CHECK_REQUIRED"},
{name: "revision conflict", err: fmt.Errorf("persist rule: %w", repo.ErrFirewallRuleRevisionConflict), code: "FW_RULE_REVISION_CONFLICT"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
recorder := httptest.NewRecorder()
context, _ := gin.CreateTestContext(recorder)
handleFirewallRuleError(context, test.err)
var response dto.Response
if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
t.Fatalf("decode response: %v", err)
}
if response.Code != 409 || response.ErrorCode != test.code {
t.Fatalf("unexpected firewall error response: %#v", response)
}
})
}
}
func TestNormalizeFirewallRuleUUID(t *testing.T) {
gin.SetMode(gin.TestMode)
t.Run("trims valid UUID", func(t *testing.T) {
recorder := httptest.NewRecorder()
context, _ := gin.CreateTestContext(recorder)
value := " managed-rule "
if !normalizeFirewallRuleUUID(context, &value) || value != "managed-rule" {
t.Fatalf("normalized UUID = %q", value)
}
})
t.Run("rejects blank UUID", func(t *testing.T) {
recorder := httptest.NewRecorder()
context, _ := gin.CreateTestContext(recorder)
value := " "
if normalizeFirewallRuleUUID(context, &value) {
t.Fatal("blank UUID was accepted")
}
var response dto.Response
if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
t.Fatalf("decode response: %v", err)
}
if recorder.Code != http.StatusOK || response.Code != http.StatusBadRequest {
t.Fatalf("transport status = %d, response = %#v", recorder.Code, response)
}
})
}
+6 -2
View File
@@ -1,6 +1,9 @@
package dto
import "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
import (
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
firewallsync "github.com/1Panel-dev/1Panel/agent/utils/firewall/sync"
)
type FirewallSubsystemStatus struct {
Name string `json:"name"`
@@ -245,7 +248,8 @@ type FirewallRuleSyncItem struct {
Rule *filter.FirewallRule `json:"rule,omitempty"`
ForwardRule *ForwardRule `json:"forwardRule,omitempty"`
DockerRule *DockerPortGuardEndpoint `json:"dockerRule,omitempty"`
Status string `json:"status"`
Status firewallsync.Status `json:"status"`
ReasonCode firewallsync.ReasonCode `json:"reasonCode,omitempty"`
Reason string `json:"reason,omitempty"`
}
+110
View File
@@ -5,6 +5,7 @@ import (
"encoding/hex"
"encoding/json"
"fmt"
"sort"
"strings"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
@@ -67,6 +68,19 @@ func FirewallRuleOwner(sourceKind, sourceID string) string {
return sourceKind + ":" + sourceID
}
func FirewallRulesRevision(rules []FirewallRule) (string, error) {
ordered := append([]FirewallRule(nil), rules...)
sort.Slice(ordered, func(i, j int) bool {
return ordered[i].UUID < ordered[j].UUID
})
payload, err := json.Marshal(ordered)
if err != nil {
return "", err
}
sum := sha256.Sum256(payload)
return hex.EncodeToString(sum[:]), nil
}
func FirewallRuleFromDomain(rule filter.FirewallRule) (FirewallRule, error) {
normalized, err := filter.NormalizeRule(rule)
if err != nil {
@@ -115,3 +129,99 @@ func (rule FirewallRule) PolicyKey() string {
sum := sha256.Sum256(payload)
return hex.EncodeToString(sum[:])
}
func (rule FirewallRule) RulesForProvider(provider filter.Provider) ([]filter.FirewallRule, error) {
if rule.CompatibilityError != "" {
return nil, fmt.Errorf("%w: %s", filter.ErrUnsupportedScope, rule.CompatibilityError)
}
connectionStates := make([]string, 0)
if rule.ConnectionStates != "" {
connectionStates = strings.Split(rule.ConnectionStates, ",")
}
base := filter.FirewallRule{
Protocol: rule.Protocol, SourceAddress: rule.SourceAddress, SourcePort: rule.SourcePort,
DestinationAddress: rule.DestinationAddress, DestinationPort: rule.DestinationPort,
Interface: rule.Interface, ConnectionStates: connectionStates,
Action: filter.Action(rule.Action), Description: rule.Description,
}
if provider == filter.ProviderFirewalld {
base.Priority = rule.Priority
}
families := []filter.Family{filter.Family(rule.Family)}
if provider != filter.ProviderFirewalld && families[0] == filter.FamilyInet {
hasIPv4, hasIPv6 := ruleAddressFamilies(base)
switch {
case hasIPv4 && hasIPv6:
return nil, fmt.Errorf("%w: inet policy contains both IPv4 and IPv6 addresses", filter.ErrUnsupportedScope)
case hasIPv6 || strings.EqualFold(base.Protocol, "icmpv6"):
families = []filter.Family{filter.FamilyIPv6}
case hasIPv4:
families = []filter.Family{filter.FamilyIPv4}
default:
families = []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6}
}
}
result := make([]filter.FirewallRule, 0, len(families))
for _, family := range families {
compiled := base
compiled.Scope = filter.Scope{Provider: provider, Family: family, Direction: filter.DirectionInput}
switch provider {
case filter.ProviderIptables, filter.ProviderNftables:
compiled.Scope.Table, compiled.Scope.Chain = "filter", filter.IptablesInputChain
case filter.ProviderFirewalld:
compiled.Scope.Zone = filter.FirewalldInputZone
case filter.ProviderUFW:
compiled.Scope.Chain = filter.UFWInputChain
default:
return nil, fmt.Errorf("%w: unsupported firewall provider %q", filter.ErrProviderUnavailable, provider)
}
expanded, err := filter.ExpandAtomicRules(compiled)
if err != nil {
return nil, err
}
result = append(result, expanded...)
}
return result, nil
}
func SortFirewallRules(rules []FirewallRule, provider filter.Provider) {
sort.SliceStable(rules, func(i, j int) bool {
left, right := rules[i], rules[j]
if provider == filter.ProviderFirewalld {
switch {
case left.Priority == nil && right.Priority != nil:
return false
case left.Priority != nil && right.Priority == nil:
return true
case left.Priority != nil && right.Priority != nil && *left.Priority != *right.Priority:
return *left.Priority < *right.Priority
}
} else {
switch {
case left.Sequence == nil && right.Sequence != nil:
return false
case left.Sequence != nil && right.Sequence == nil:
return true
case left.Sequence != nil && right.Sequence != nil && *left.Sequence != *right.Sequence:
return *left.Sequence < *right.Sequence
}
}
return left.UUID < right.UUID
})
}
func ruleAddressFamilies(rule filter.FirewallRule) (bool, bool) {
hasIPv4, hasIPv6 := false, false
for _, address := range []string{rule.SourceAddress, rule.DestinationAddress} {
address = strings.TrimSpace(address)
if address == "" {
continue
}
if strings.Contains(address, ":") {
hasIPv6 = true
} else {
hasIPv4 = true
}
}
return hasIPv4, hasIPv6
}
-78
View File
@@ -1,78 +0,0 @@
package model
import (
"errors"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
func TestFirewallRuleFromDomainUsesProviderNeutralIdentity(t *testing.T) {
rule := filter.FirewallRule{
Scope: filter.Scope{
Provider: filter.ProviderIptables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
},
Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept,
}
iptables, err := FirewallRuleFromDomain(rule)
if err != nil {
t.Fatal(err)
}
rule.Scope.Provider = filter.ProviderNftables
nftables, err := FirewallRuleFromDomain(rule)
if err != nil {
t.Fatal(err)
}
if iptables.PolicyKey() != nftables.PolicyKey() {
t.Fatalf("provider leaked into desired-rule identity: iptables=%#v nftables=%#v", iptables, nftables)
}
}
func TestFirewallRuleFromDomainRejectsProviderNativeOnlyRule(t *testing.T) {
_, err := FirewallRuleFromDomain(filter.FirewallRule{
Scope: filter.Scope{
Provider: filter.ProviderFirewalld, Family: filter.FamilyInet,
Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput,
},
NativeKind: filter.NativeKindZoneService,
Protocol: "all",
Action: filter.ActionAccept,
})
if !errors.Is(err, filter.ErrUnsupportedScope) {
t.Fatalf("provider-native rule error = %v, want unsupported scope", err)
}
}
func TestFirewallRuleFromDomainPersistsOnlyFirewalldPriority(t *testing.T) {
priority := -100
rule := filter.FirewallRule{
Scope: filter.Scope{
Provider: filter.ProviderFirewalld, Family: filter.FamilyInet,
Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput,
},
NativeKind: filter.NativeKindRichRule, Protocol: "tcp", DestinationPort: "443",
Action: filter.ActionAccept, Priority: &priority,
}
firewalld, err := FirewallRuleFromDomain(rule)
if err != nil {
t.Fatal(err)
}
if firewalld.Priority == nil || *firewalld.Priority != priority || firewalld.Sequence != nil {
t.Fatalf("firewalld placement was not persisted: %#v", firewalld)
}
rule.Scope = filter.Scope{
Provider: filter.ProviderIptables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
}
rule.NativeKind = filter.NativeKindRule
rule.Priority = nil
iptables, err := FirewallRuleFromDomain(rule)
if err != nil {
t.Fatal(err)
}
if iptables.Priority != nil || iptables.Sequence != nil {
t.Fatalf("positional backend persisted provider priority: %#v", iptables)
}
}
-120
View File
@@ -1,120 +0,0 @@
package repo
import (
"context"
"errors"
"fmt"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/glebarez/sqlite"
"github.com/google/uuid"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func TestFirewallRuleRepoRevision(t *testing.T) {
db := newFirewallRepoTestDB(t)
repository := NewFirewallRuleRepo(db)
ctx := context.Background()
rule := newFirewallRuleModel()
if err := repository.Create(ctx, &rule); err != nil {
t.Fatalf("create rule: %v", err)
}
if rule.UUID == "" || rule.Revision != 1 {
t.Fatalf("rule defaults were not applied: %#v", rule)
}
if err := repository.UpdateWithRevision(ctx, rule.UUID, 2, map[string]interface{}{"description": "updated"}); !errors.Is(err, ErrFirewallRuleRevisionConflict) {
t.Fatalf("expected revision conflict, got %v", err)
}
if err := repository.UpdateWithRevision(ctx, rule.UUID, 1, map[string]interface{}{"description": "updated"}); err != nil {
t.Fatalf("update rule: %v", err)
}
updated, err := repository.GetByUUID(ctx, rule.UUID)
if err != nil {
t.Fatalf("get updated rule: %v", err)
}
if updated.Revision != 2 || updated.Description != "updated" {
t.Fatalf("unexpected updated rule: %#v", updated)
}
if err := repository.DeleteWithRevision(ctx, rule.UUID, updated.Revision-1); !errors.Is(err, ErrFirewallRuleRevisionConflict) {
t.Fatalf("expected revision conflict, got %v", err)
}
if err := repository.DeleteWithRevision(ctx, rule.UUID, updated.Revision); err != nil {
t.Fatalf("delete rule: %v", err)
}
if _, err := repository.GetByUUID(ctx, rule.UUID); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("expected deleted rule to be absent, got %v", err)
}
var deletedCount int64
if err := db.Model(&model.FirewallRule{}).Where("uuid = ?", rule.UUID).Count(&deletedCount).Error; err != nil || deletedCount != 0 {
t.Fatalf("rule was not hard deleted: count=%d err=%v", deletedCount, err)
}
}
func TestFirewallRepositoriesRejectIncompleteRecords(t *testing.T) {
db := newFirewallRepoTestDB(t)
ctx := context.Background()
if err := NewFirewallRuleRepo(db).Create(ctx, &model.FirewallRule{}); !errors.Is(err, ErrFirewallPersistenceInvalid) {
t.Fatalf("expected invalid rule error, got %v", err)
}
}
func TestFirewallRepositoriesUseContextTransaction(t *testing.T) {
db := newFirewallRepoTestDB(t)
ruleRepo := NewFirewallRuleRepo(db)
wantRollback := errors.New("rollback")
err := db.Transaction(func(tx *gorm.DB) error {
ctx := context.WithValue(context.Background(), constant.DB, tx)
rule := newFirewallRuleModel()
if err := ruleRepo.Create(ctx, &rule); err != nil {
return err
}
return wantRollback
})
if !errors.Is(err, wantRollback) {
t.Fatalf("expected rollback error, got %v", err)
}
var ruleCount int64
if err := db.Model(&model.FirewallRule{}).Count(&ruleCount).Error; err != nil {
t.Fatalf("count rules: %v", err)
}
if ruleCount != 0 {
t.Fatalf("transaction did not roll back: rules=%d", ruleCount)
}
}
func newFirewallRepoTestDB(t *testing.T) *gorm.DB {
t.Helper()
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString())
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("load sql db: %v", err)
}
sqlDB.SetMaxOpenConns(1)
t.Cleanup(func() { _ = sqlDB.Close() })
if err := db.AutoMigrate(&model.FirewallRule{}); err != nil {
t.Fatalf("migrate models: %v", err)
}
return db
}
func newFirewallRuleModel() model.FirewallRule {
return model.FirewallRule{
Family: "ipv4",
Protocol: "tcp",
DestinationPort: "22",
Action: "accept",
}
}
-47
View File
@@ -1,47 +0,0 @@
package repo
import (
"context"
"fmt"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/global"
"github.com/glebarez/sqlite"
"github.com/google/uuid"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func TestForwardingRuleRepoReplaceAll(t *testing.T) {
previousDB := global.DB
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString())
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.ForwardingRule{}); err != nil {
t.Fatal(err)
}
global.DB = db
t.Cleanup(func() { global.DB = previousDB })
repository := NewIForwardingRuleRepo()
first := []model.ForwardingRule{{Family: "ipv4", Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80"}}
if err := repository.ReplaceAll(context.Background(), first); err != nil {
t.Fatal(err)
}
rules, err := repository.List(context.Background())
if err != nil || len(rules) != 1 || rules[0].Port != "8080" {
t.Fatalf("rules = %#v, err = %v", rules, err)
}
second := []model.ForwardingRule{{Family: "ipv6", Protocol: "udp", Port: "5353", TargetIP: "::1", TargetPort: "53", Interface: "eth0"}}
if err := repository.ReplaceAll(context.Background(), second); err != nil {
t.Fatal(err)
}
rules, err = repository.List(context.Background())
if err != nil || len(rules) != 1 || rules[0].Family != "ipv6" || rules[0].Port != "5353" {
t.Fatalf("rules = %#v, err = %v", rules, err)
}
}
@@ -2,11 +2,6 @@ package service
import (
"context"
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"os"
@@ -22,10 +17,7 @@ import (
"github.com/1Panel-dev/1Panel/agent/global"
"github.com/1Panel-dev/1Panel/agent/utils/firewall"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
filterfirewalld "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/firewalld"
filteriptables "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/iptables"
filternftables "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/nftables"
filterufw "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/ufw"
filterruntime "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/runtime"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/iptables_helper"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/nftables_helper"
@@ -35,7 +27,7 @@ import (
type FirewallService struct {
rules repo.IFirewallRuleRepo
adapters firewallRuleRuntimeRegistry
adapters firewallRuleRuntimeResolver
forwardingSync firewallDatabaseSyncAdapter
dockerSync firewallDatabaseSyncAdapter
selectedProvider func(context.Context) (filter.Provider, error)
@@ -48,23 +40,28 @@ type FirewallService struct {
installedProviders func() []string
}
type firewallRuleRuntimeResolver interface {
Resolve(filter.Provider) (*firewallRuleRuntime, error)
Providers() []filter.Provider
}
var firewallRuleMutationMu sync.Mutex
type IFirewallService interface {
LoadBaseInfo(chainGroup string) (dto.FirewallSubsystemStatus, error)
OperateFirewall(request dto.FirewallLifecycleOperation) error
OperateFilterChain(request dto.FilterChainOperation) error
Reset(context.Context, dto.FirewallRuleReset) (dto.FirewallRuleResetResponse, error)
Inventory(context.Context, dto.FirewallRuleInventory) (dto.FirewallRuleInventoryResponse, error)
LoadFirewallNativeDetail(context.Context, dto.FirewallNativeDetail) (string, error)
Check(context.Context, string, dto.FirewallRuleCheck) (dto.FirewallRuleCheckResponse, error)
Create(context.Context, dto.FirewallRuleCreate) (dto.FirewallRuleCreateResponse, error)
PreviewRuleSync(context.Context, string, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncPreview, error)
SyncRules(context.Context, string, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error)
CurrentRuleSyncTask() (dto.FirewallRuleSyncTask, error)
Delete(context.Context, dto.FirewallRuleDelete) (dto.FirewallRuleDeleteResponse, error)
Update(context.Context, string, dto.FirewallRuleUpdate) error
Reorder(context.Context, string, dto.FirewallRuleReorder) error
Reset(context.Context, dto.FirewallRuleReset) (dto.FirewallRuleResetResponse, error)
PreviewRuleSync(context.Context, string, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncPreview, error)
SyncRules(context.Context, string, dto.FirewallRuleSyncRequest) (dto.FirewallRuleSyncResult, error)
CurrentRuleSyncTask() (dto.FirewallRuleSyncTask, error)
}
func NewIFirewallService() IFirewallService {
@@ -425,11 +422,7 @@ func (s *FirewallService) LoadFirewallNativeDetail(ctx context.Context, request
if err != nil {
return "", err
}
informer, ok := runtime.adapter.(filter.NativeDetailReader)
if !ok {
return "", fmt.Errorf("%w: native details for %s", filter.ErrAdapterUnavailable, provider)
}
return informer.NativeDetail(ctx, request.Name, request.Permanent)
return runtime.NativeDetail(ctx, request.Name, request.Permanent)
}
func (s *FirewallService) checkUpdate(
@@ -531,7 +524,7 @@ func (s *FirewallService) Check(
if desiredErr != nil {
return dto.FirewallRuleCheckResponse{}, desiredErr
}
managedRevision, revisionErr := firewallManagedRevision(stored)
managedRevision, revisionErr := model.FirewallRulesRevision(stored)
if revisionErr != nil {
return dto.FirewallRuleCheckResponse{}, revisionErr
}
@@ -742,7 +735,7 @@ func (s *FirewallService) prepareCreate(
if listErr != nil {
return nil, index, listErr
}
managedRevision, revisionErr := firewallManagedRevision(stored)
managedRevision, revisionErr := model.FirewallRulesRevision(stored)
if revisionErr != nil {
return nil, index, revisionErr
}
@@ -773,7 +766,7 @@ func nativeCreateBatchEnd(prepared []preparedFirewallRuleCreate, start int) int
return start
}
first := prepared[start]
if first.runtime == nil || first.runtime.adapter == nil || !supportsNativeRuleBatch(first.runtime.adapter.Provider()) ||
if first.runtime == nil || !supportsNativeRuleBatch(first.runtime.Provider()) ||
first.authorization.Operation != filter.ChangeCreate || first.request.Rule.OrderIndex != nil {
return start + 1
}
@@ -1160,14 +1153,14 @@ func (s *FirewallService) createRule(
domainRule := request.Rule
appendRule := false
if authorization.Operation == filter.ChangeCreate && domainRule.Scope.Provider == filter.ProviderUFW && domainRule.OrderIndex == nil {
appendPosition, err := ufwAppendPosition(ctx, runtime, snapshot, domainRule)
appendPosition, err := runtime.AppendPosition(ctx, snapshot, domainRule)
if err != nil {
return err
}
domainRule.OrderIndex = &appendPosition
appendRule = true
} else if authorization.Operation == filter.ChangeCreate && domainRule.OrderIndex != nil {
maxPosition, err := maxPositionForRule(ctx, runtime, snapshot, domainRule)
maxPosition, err := runtime.MaxPosition(ctx, snapshot, domainRule)
if err != nil {
return err
}
@@ -1286,7 +1279,7 @@ func (s *FirewallService) deleteRule(ctx context.Context, ruleUUID string) error
}
restoreAtEnd := false
if desired.Rule.Scope.Provider == filter.ProviderUFW && observed.Locator.Position != nil {
maxPosition, maxErr := maxPositionForRule(ctx, runtime, snapshot, desired.Rule)
maxPosition, maxErr := runtime.MaxPosition(ctx, snapshot, desired.Rule)
if maxErr != nil {
return rollback(maxErr)
}
@@ -1378,7 +1371,7 @@ func (s *FirewallService) reorderRule(ctx context.Context, clientIP, ruleUUID st
if targetPosition == nil || *targetPosition < 1 {
return fmt.Errorf("%w: target position is required", filter.ErrInvalidRule)
}
if err := validatePositionTarget(ctx, runtime, snapshot, before.Rule, *targetPosition); err != nil {
if err := runtime.ValidatePosition(ctx, snapshot, before.Rule, *targetPosition); err != nil {
return err
}
after.OrderIndex = targetPosition
@@ -1408,101 +1401,6 @@ func (s *FirewallService) reorderRule(ctx context.Context, clientIP, ruleUUID st
})
}
func validatePositionTarget(
ctx context.Context,
runtime *firewallRuleRuntime,
snapshot filter.Snapshot,
rule filter.FirewallRule,
targetPosition int64,
) error {
if rule.Scope.Provider == filter.ProviderUFW {
minPosition, maxPosition := snapshotPositionBounds(snapshot)
if targetPosition < minPosition || targetPosition > maxPosition {
return fmt.Errorf(
"%w: target position %d is outside the %s range %d-%d",
filter.ErrInvalidRule, targetPosition, rule.Scope.Family, minPosition, maxPosition,
)
}
return nil
}
maxPosition, err := maxPositionForRule(ctx, runtime, snapshot, rule)
if err != nil {
return err
}
if targetPosition > maxPosition {
return fmt.Errorf("%w: target position %d is out of range 1-%d", filter.ErrInvalidRule, targetPosition, maxPosition)
}
return nil
}
func ufwAppendPosition(
ctx context.Context,
runtime *firewallRuleRuntime,
snapshot filter.Snapshot,
rule filter.FirewallRule,
) (int64, error) {
if rule.Scope.Family == filter.FamilyIPv4 {
return snapshotMaxPosition(snapshot) + 1, nil
}
maxPosition, err := maxPositionForRule(ctx, runtime, snapshot, rule)
if err != nil {
return 0, err
}
return maxPosition + 1, nil
}
func snapshotPositionBounds(snapshot filter.Snapshot) (int64, int64) {
minPosition, maxPosition := int64(0), int64(0)
for _, observed := range snapshot.Rules {
if observed.Locator.Position == nil {
continue
}
position := int64(*observed.Locator.Position)
if minPosition == 0 || position < minPosition {
minPosition = position
}
if position > maxPosition {
maxPosition = position
}
}
return minPosition, maxPosition
}
func maxPositionForRule(
ctx context.Context,
runtime *firewallRuleRuntime,
snapshot filter.Snapshot,
rule filter.FirewallRule,
) (int64, error) {
maxPosition := snapshotMaxPosition(snapshot)
if rule.Scope.Provider == filter.ProviderUFW {
relatedScope := rule.Scope
if relatedScope.Family == filter.FamilyIPv4 {
relatedScope.Family = filter.FamilyIPv6
} else {
relatedScope.Family = filter.FamilyIPv4
}
relatedSnapshot, err := runtime.ObserveMutation(ctx, relatedScope)
if err != nil {
return 0, err
}
if relatedMax := snapshotMaxPosition(relatedSnapshot); relatedMax > maxPosition {
maxPosition = relatedMax
}
}
return maxPosition, nil
}
func snapshotMaxPosition(snapshot filter.Snapshot) int64 {
var maxPosition int64
for _, observed := range snapshot.Rules {
if observed.Locator.Position != nil && int64(*observed.Locator.Position) > maxPosition {
maxPosition = int64(*observed.Locator.Position)
}
}
return maxPosition
}
type managedMutationRequest struct {
Stored model.FirewallRule
Before filter.FirewallRule
@@ -1566,7 +1464,7 @@ func (s *FirewallService) prepareManagedUpdate(
if after.OrderIndex == nil {
after.OrderIndex = &currentPosition
} else if *after.OrderIndex != currentPosition {
if err := validatePositionTarget(ctx, runtime, snapshot, before.Rule, *after.OrderIndex); err != nil {
if err := runtime.ValidatePosition(ctx, snapshot, before.Rule, *after.OrderIndex); err != nil {
return preparedManagedUpdate{}, err
}
}
@@ -1630,9 +1528,10 @@ func (s *FirewallService) selectedProviderForStoredRule(
if s.selectedProvider != nil {
return s.selectedProvider(ctx)
}
if len(s.adapters) == 1 {
for provider := range s.adapters {
return provider, nil
if s.adapters != nil {
providers := s.adapters.Providers()
if len(providers) == 1 {
return providers[0], nil
}
}
return "", fmt.Errorf("%w: selected provider is unavailable", filter.ErrProviderUnavailable)
@@ -1649,7 +1548,7 @@ func (s *FirewallService) executeManagedMutation(ctx context.Context, request ma
}
appendRule, restoreAtEnd := false, false
if after.Scope.Provider == filter.ProviderUFW && request.AdapterOperation == filter.ChangeUpdate {
maxPosition, maxErr := maxPositionForRule(ctx, request.Runtime, request.Snapshot, after)
maxPosition, maxErr := request.Runtime.MaxPosition(ctx, request.Snapshot, after)
if maxErr != nil {
return maxErr
}
@@ -1910,71 +1809,39 @@ func (s *FirewallService) adoptExternalSystemPort(ctx context.Context, port dto.
}
func systemPortRule(provider filter.Provider, port dto.FirewallSystemPort) filter.FirewallRule {
scope := filter.Scope{Provider: provider, Direction: filter.DirectionInput}
family := filter.Family(strings.ToLower(strings.TrimSpace(port.Family)))
switch provider {
case filter.ProviderIptables, filter.ProviderNftables:
if family != filter.FamilyIPv6 {
family = filter.FamilyIPv4
}
scope.Family = family
scope.Table = "filter"
case filter.ProviderFirewalld:
if family != filter.FamilyIPv4 && family != filter.FamilyIPv6 {
family = filter.FamilyInet
}
scope.Family = family
scope.Zone = filter.FirewalldInputZone
case filter.ProviderUFW:
if family != filter.FamilyIPv6 {
family = filter.FamilyIPv4
}
scope.Family = family
}
return filter.FirewallRule{
Scope: scope, Protocol: port.Protocol, DestinationPort: port.Port,
Action: filter.ActionAccept, Description: "1Panel managed accepted port",
}
return firewall.RuleForSystemPort(provider, firewall.SystemPort(port))
}
func normalizeSystemPorts(ports []dto.FirewallSystemPort) (map[string]dto.FirewallSystemPort, error) {
result := make(map[string]dto.FirewallSystemPort, len(ports))
domainPorts := make([]firewall.SystemPort, 0, len(ports))
for _, port := range ports {
normalized, err := filter.NormalizeRule(systemPortRule(filter.ProviderIptables, port))
if err != nil {
return nil, err
}
family := strings.ToLower(strings.TrimSpace(port.Family))
if family != "" {
family = string(normalized.Scope.Family)
}
item := dto.FirewallSystemPort{
Family: family, Port: normalized.DestinationPort, Protocol: normalized.Protocol,
}
result[systemPortKey(item)] = item
domainPorts = append(domainPorts, firewall.SystemPort(port))
}
normalized, err := firewall.NormalizeSystemPorts(domainPorts)
if err != nil {
return nil, err
}
result := make(map[string]dto.FirewallSystemPort, len(normalized))
for key, port := range normalized {
result[key] = dto.FirewallSystemPort(port)
}
return result, nil
}
func systemPortKey(port dto.FirewallSystemPort) string {
key := legacySystemPortKey(port)
if family := strings.ToLower(strings.TrimSpace(port.Family)); family != "" {
return family + "/" + key
}
return key
return firewall.SystemPortKey(firewall.SystemPort(port))
}
func legacySystemPortKey(port dto.FirewallSystemPort) string {
return strings.ToLower(strings.TrimSpace(port.Protocol)) + "/" + strings.TrimSpace(port.Port)
return firewall.LegacySystemPortKey(firewall.SystemPort(port))
}
func sortedSystemPortKeys(ports map[string]dto.FirewallSystemPort) []string {
keys := make([]string, 0, len(ports))
for key := range ports {
keys = append(keys, key)
domainPorts := make(map[string]firewall.SystemPort, len(ports))
for key, port := range ports {
domainPorts[key] = firewall.SystemPort(port)
}
sort.Strings(keys)
return keys
return firewall.SortedSystemPortKeys(domainPorts)
}
func firewallRuleModelForCreate(rule filter.FirewallRule, request dto.FirewallRuleCreateItem, origin string) (model.FirewallRule, error) {
@@ -2162,7 +2029,7 @@ func mergeFirewallInventory(
}
func desiredFirewallRuleFromModel(stored model.FirewallRule) (filter.DesiredRule, error) {
rules, err := firewallPolicyRulesForProvider(stored, filter.ProviderIptables)
rules, err := stored.RulesForProvider(filter.ProviderIptables)
if err != nil {
return filter.DesiredRule{}, err
}
@@ -2188,7 +2055,7 @@ func (s *FirewallService) compileStoredFirewallRules(
stored model.FirewallRule,
target filter.Provider,
) ([]filter.DesiredRule, error) {
rules, err := firewallPolicyRulesForProvider(stored, target)
rules, err := stored.RulesForProvider(target)
if err != nil {
return nil, err
}
@@ -2196,41 +2063,7 @@ func (s *FirewallService) compileStoredFirewallRules(
if err != nil {
return nil, err
}
result := make([]filter.DesiredRule, 0, len(rules))
scopeOrdinals := make(map[string]int)
for _, rule := range rules {
rule, err = runtime.Prepare(rule)
if err != nil {
return nil, err
}
if err = runtime.CheckRule(ctx, rule); err != nil {
return nil, err
}
ruleKey, keyErr := filter.RuleKey(rule)
if keyErr != nil {
return nil, keyErr
}
scopeKey := rule.Scope.Key()
ordinal := scopeOrdinals[scopeKey]
scopeOrdinals[scopeKey] = ordinal + 1
rule.UUID = compiledFirewallRuleUUID(stored.UUID, ruleKey, ordinal)
result = append(result, filter.DesiredRule{
UUID: stored.UUID, Rule: rule, RuleKey: ruleKey, Origin: filter.RuleOrigin(stored.Origin),
Marker: "1panel-rule:" + rule.UUID,
})
}
return result, nil
}
func compiledFirewallRuleUUID(policyUUID, ruleKey string, scopeOrdinal int) string {
if scopeOrdinal == 0 {
return policyUUID
}
const suffixLength = 12
if len(ruleKey) > suffixLength {
ruleKey = ruleKey[:suffixLength]
}
return fmt.Sprintf("%s-%d-%s", policyUUID, scopeOrdinal+1, ruleKey)
return runtime.CompileDesired(ctx, stored.UUID, filter.RuleOrigin(stored.Origin), rules)
}
func (s *FirewallService) desiredFirewallRulesForScope(
@@ -2238,7 +2071,7 @@ func (s *FirewallService) desiredFirewallRulesForScope(
stored []model.FirewallRule,
scope filter.Scope,
) ([]filter.DesiredRule, error) {
sortFirewallPolicies(stored, scope.Provider)
model.SortFirewallRules(stored, scope.Provider)
desired := make([]filter.DesiredRule, 0, len(stored))
for _, record := range stored {
compiled, err := s.compileStoredFirewallRules(ctx, record, scope.Provider)
@@ -2266,6 +2099,18 @@ func firewallRuleSelectedProvider(context.Context) (filter.Provider, error) {
return selectedRuleProvider()
}
type firewallSnapshotPolicy = filterruntime.SnapshotPolicy
type firewallRuleRuntime = filterruntime.Engine
type firewallRuleRuntimeRegistry = filterruntime.Registry
func newFirewallRuleRuntimeRegistry(policy firewallSnapshotPolicy) firewallRuleRuntimeRegistry {
return filterruntime.NewRegistry(policy)
}
func newFirewallRuleRuntime(adapter filter.Adapter, policy firewallSnapshotPolicy) *firewallRuleRuntime {
return filterruntime.New(adapter, policy)
}
func rollbackFirewallPlan(ctx context.Context, runtime *firewallRuleRuntime, plan filter.BackendPlan, cause error) error {
if runtime == nil {
return cause
@@ -2276,138 +2121,6 @@ func rollbackFirewallPlan(ctx context.Context, runtime *firewallRuleRuntime, pla
return cause
}
type firewallSnapshotPolicy func(context.Context, filter.Snapshot) (filter.Snapshot, error)
type firewallRuleRuntime struct {
adapter filter.Adapter
policy firewallSnapshotPolicy
}
type firewallRuleRuntimeRegistry map[filter.Provider]*firewallRuleRuntime
func newFirewallRuleRuntimeRegistry(policy firewallSnapshotPolicy) firewallRuleRuntimeRegistry {
return firewallRuleRuntimeRegistry{
filter.ProviderIptables: newFirewallRuleRuntime(filteriptables.NewAdapter(), policy),
filter.ProviderNftables: newFirewallRuleRuntime(filternftables.NewAdapter(), policy),
filter.ProviderFirewalld: newFirewallRuleRuntime(filterfirewalld.NewAdapter(), policy),
filter.ProviderUFW: newFirewallRuleRuntime(filterufw.NewAdapter(), policy),
}
}
func newFirewallRuleRuntime(adapter filter.Adapter, policy firewallSnapshotPolicy) *firewallRuleRuntime {
return &firewallRuleRuntime{adapter: adapter, policy: policy}
}
func (r firewallRuleRuntimeRegistry) Resolve(provider filter.Provider) (*firewallRuleRuntime, error) {
runtime, exists := r[provider]
if !exists || runtime == nil || runtime.adapter == nil {
return nil, fmt.Errorf("%w: %s", filter.ErrAdapterUnavailable, provider)
}
return runtime, nil
}
func (r *firewallRuleRuntime) Observe(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) {
snapshot, err := r.adapter.Observe(ctx, scope)
if err != nil {
return filter.Snapshot{}, err
}
if r.policy == nil {
return snapshot, nil
}
return r.policy(ctx, snapshot)
}
func (r *firewallRuleRuntime) ObserveScopes(ctx context.Context, scopes []filter.Scope) ([]filter.Snapshot, error) {
observer, ok := r.adapter.(filter.MultiScopeObserver)
if !ok {
return nil, fmt.Errorf("%w: %s multi-scope inventory", filter.ErrAdapterUnavailable, r.adapter.Provider())
}
snapshots, err := observer.ObserveScopes(ctx, scopes)
if err != nil {
return nil, err
}
if r.policy == nil {
return snapshots, nil
}
for index := range snapshots {
snapshots[index], err = r.policy(ctx, snapshots[index])
if err != nil {
return nil, err
}
}
return snapshots, nil
}
func (r *firewallRuleRuntime) ObserveMutation(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) {
snapshot, err := r.Observe(ctx, scope)
if err != nil {
return filter.Snapshot{}, err
}
for _, notice := range snapshot.Notices {
if notice.Code == filter.ScopeNoticeManagedScopeInactive || notice.Code == filter.ScopeNoticeManagedScopeMissing {
return filter.Snapshot{}, fmt.Errorf("%w: managed firewall scope is unavailable", filter.ErrProviderUnavailable)
}
}
return snapshot, nil
}
func (r *firewallRuleRuntime) Prepare(rule filter.FirewallRule) (filter.FirewallRule, error) {
preparer, ok := r.adapter.(filter.RulePreparer)
if !ok {
return rule, nil
}
return preparer.PrepareRule(rule)
}
func (r *firewallRuleRuntime) CheckRule(ctx context.Context, rule filter.FirewallRule) error {
checker, ok := r.adapter.(filter.RuleChecker)
if !ok {
return nil
}
return checker.CheckRule(ctx, rule)
}
func (r *firewallRuleRuntime) Capabilities(ctx context.Context) (filter.Capabilities, error) {
return r.adapter.Capabilities(ctx)
}
func (r *firewallRuleRuntime) Execute(ctx context.Context, snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, filter.VerifyResult, error) {
plan, err := r.adapter.Compile(snapshot, changes)
if err != nil {
return filter.BackendPlan{}, filter.VerifyResult{}, err
}
result, err := r.adapter.Apply(ctx, plan)
if err != nil {
return plan, filter.VerifyResult{}, err
}
if result.Verification != nil {
if !result.Verification.Matched {
if rollbackErr := r.Rollback(ctx, plan); rollbackErr != nil {
return plan, *result.Verification, errors.Join(filter.ErrVerificationFailed, rollbackErr)
}
}
return plan, *result.Verification, nil
}
verification, err := r.adapter.Verify(ctx, plan)
if err != nil {
return plan, verification, rollbackFirewallPlan(ctx, r, plan, err)
}
if !verification.Matched {
if rollbackErr := r.Rollback(ctx, plan); rollbackErr != nil {
return plan, verification, errors.Join(filter.ErrVerificationFailed, rollbackErr)
}
}
return plan, verification, nil
}
func (r *firewallRuleRuntime) Rollback(ctx context.Context, plan filter.BackendPlan) error {
rollbacker, ok := r.adapter.(filter.PlanRollbacker)
if !ok {
return fmt.Errorf("%w: provider %s does not support applied-plan rollback", filter.ErrAdapterUnavailable, r.adapter.Provider())
}
return rollbacker.Rollback(ctx, plan)
}
func selectedRuleProvider() (filter.Provider, error) {
provider, err := selectedSystemFirewallProvider()
if err != nil {
@@ -2416,28 +2129,7 @@ func selectedRuleProvider() (filter.Provider, error) {
return filter.Provider(provider), nil
}
type firewallRuleCheckClaims struct {
Version int `json:"version"`
Provider filter.Provider `json:"provider"`
ScopeKey string `json:"scopeKey"`
RuleDigest string `json:"ruleDigest"`
SnapshotRevision string `json:"snapshotRevision"`
ManagedRevision string `json:"managedRevision"`
Decision filter.CheckDecision `json:"decision"`
Classification filter.CheckClassification `json:"classification"`
AllowedActions []filter.CheckAction `json:"allowedActions"`
AdoptionCandidates []firewallAdoptionCandidate `json:"adoptionCandidates,omitempty"`
}
type firewallAdoptionCandidate struct {
InstanceKey string `json:"instanceKey"`
Locator filter.Locator `json:"locator"`
}
type firewallRuleCreateAuthorization struct {
Operation filter.ChangeOperation
Locator *filter.Locator
}
type firewallRuleCreateAuthorization = filter.CreateAuthorization
func refreshCreateAuthorization(
snapshot filter.Snapshot,
@@ -2486,36 +2178,7 @@ func refreshCreateAuthorization(
}
func signFirewallRuleCheck(result filter.RuleCheckResult, snapshot filter.Snapshot, managedRevision string) (string, error) {
ruleDigest, err := firewallRuleDigest(result.RequestedRule)
if err != nil {
return "", err
}
claims := firewallRuleCheckClaims{
Version: constant.FirewallRuleCheckVersion,
Provider: result.RequestedRule.Scope.Provider,
ScopeKey: result.RequestedRule.Scope.Key(),
RuleDigest: ruleDigest,
SnapshotRevision: snapshot.Revision,
ManagedRevision: managedRevision,
Decision: result.Decision,
Classification: result.Classification,
AllowedActions: result.AllowedActions,
}
if result.Classification == filter.CheckClassificationExactExternal {
claims.AdoptionCandidates = make([]firewallAdoptionCandidate, 0, len(result.Candidates))
for _, candidate := range result.Candidates {
claims.AdoptionCandidates = append(claims.AdoptionCandidates, firewallAdoptionCandidate{
InstanceKey: candidate.InstanceKey,
Locator: candidate.Locator,
})
}
}
payload, err := json.Marshal(claims)
if err != nil {
return "", err
}
signature := firewallRuleCheckSignature(payload)
return base64.RawURLEncoding.EncodeToString(payload) + "." + base64.RawURLEncoding.EncodeToString(signature), nil
return firewallCheckFlagCodec().Sign(result, snapshot, managedRevision)
}
func authorizeFirewallRuleCreate(
@@ -2526,94 +2189,12 @@ func authorizeFirewallRuleCreate(
snapshot filter.Snapshot,
managedRevision string,
) (firewallRuleCreateAuthorization, error) {
claims, err := parseFirewallRuleCheck(checkFlag)
if err != nil {
return firewallRuleCreateAuthorization{}, err
}
ruleDigest, err := firewallRuleDigest(rule)
if err != nil {
return firewallRuleCreateAuthorization{}, err
}
if claims.Version != constant.FirewallRuleCheckVersion ||
claims.Provider != rule.Scope.Provider ||
claims.ScopeKey != rule.Scope.Key() ||
claims.RuleDigest != ruleDigest ||
claims.SnapshotRevision != snapshot.Revision ||
claims.ManagedRevision != managedRevision {
return firewallRuleCreateAuthorization{}, fmt.Errorf("%w: firewall or managed rules changed", filter.ErrRuleCheckRequired)
}
if claims.Decision != filter.CheckDecisionReady && claims.Decision != filter.CheckDecisionConfirmationRequired {
return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation
}
if !containsFirewallCheckAction(claims.AllowedActions, action) {
return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation
}
switch action {
case filter.CheckActionCreate, filter.CheckActionCreateAnyway:
if strings.TrimSpace(adoptInstanceKey) != "" {
return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation
}
return firewallRuleCreateAuthorization{Operation: filter.ChangeCreate}, nil
case filter.CheckActionAdopt, filter.CheckActionSelectAdopt:
for _, candidate := range claims.AdoptionCandidates {
if candidate.InstanceKey == adoptInstanceKey && adoptInstanceKey != "" {
locator := candidate.Locator
return firewallRuleCreateAuthorization{Operation: filter.ChangeAdopt, Locator: &locator}, nil
}
}
return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation
default:
return firewallRuleCreateAuthorization{}, filter.ErrRuleOperation
}
return firewallCheckFlagCodec().Authorize(checkFlag, action, adoptInstanceKey, rule, snapshot, managedRevision)
}
func parseFirewallRuleCheck(checkFlag string) (firewallRuleCheckClaims, error) {
parts := strings.Split(strings.TrimSpace(checkFlag), ".")
if len(parts) != 2 || parts[0] == "" || parts[1] == "" {
return firewallRuleCheckClaims{}, filter.ErrRuleCheckRequired
}
payload, err := base64.RawURLEncoding.DecodeString(parts[0])
if err != nil {
return firewallRuleCheckClaims{}, filter.ErrRuleCheckRequired
}
signature, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil || !hmac.Equal(signature, firewallRuleCheckSignature(payload)) {
return firewallRuleCheckClaims{}, filter.ErrRuleCheckRequired
}
var claims firewallRuleCheckClaims
if err := json.Unmarshal(payload, &claims); err != nil {
return firewallRuleCheckClaims{}, filter.ErrRuleCheckRequired
}
return claims, nil
}
func firewallRuleCheckSignature(payload []byte) []byte {
mac := hmac.New(sha256.New, []byte(global.CONF.Base.EncryptKey+"\x00firewall-rule-check-v1"))
_, _ = mac.Write(payload)
return mac.Sum(nil)
}
func firewallRuleDigest(rule filter.FirewallRule) (string, error) {
payload, err := json.Marshal(rule)
if err != nil {
return "", err
}
sum := sha256.Sum256(payload)
return hex.EncodeToString(sum[:]), nil
}
func firewallManagedRevision(rules []model.FirewallRule) (string, error) {
ordered := append([]model.FirewallRule(nil), rules...)
sort.Slice(ordered, func(i, j int) bool {
return ordered[i].UUID < ordered[j].UUID
})
payload, err := json.Marshal(ordered)
if err != nil {
return "", err
}
sum := sha256.Sum256(payload)
return hex.EncodeToString(sum[:]), nil
func firewallCheckFlagCodec() *filter.CheckFlagCodec {
secret := []byte(global.CONF.Base.EncryptKey + "\x00firewall-rule-check-v1")
return filter.NewCheckFlagCodec(secret, constant.FirewallRuleCheckVersion)
}
func containsFirewallCheckAction(actions []filter.CheckAction, expected filter.CheckAction) bool {
@@ -2678,13 +2259,7 @@ func OperateFirewallPort(oldPorts, newPorts []int) error {
}
func containsFirewallPort(ports []firewall.PortWhitelist, target firewall.PortWhitelist) bool {
for _, item := range ports {
familyMatches := item.Family == "" || target.Family == "" || item.Family == target.Family
if familyMatches && item.Port == target.Port && item.Protocol == target.Protocol {
return true
}
}
return false
return firewall.ContainsPort(ports, target)
}
func LoadPanelPort() string {
@@ -2949,21 +2524,7 @@ func systemPorts(ports []firewall.PortWhitelist) []dto.FirewallSystemPort {
}
func excludeFirewallPorts(ports, excluded []firewall.PortWhitelist) []firewall.PortWhitelist {
result := make([]firewall.PortWhitelist, 0, len(ports))
for _, port := range ports {
exists := false
for _, item := range excluded {
familyMatches := item.Family == "" || port.Family == "" || item.Family == port.Family
if familyMatches && item.Port == port.Port && item.Protocol == port.Protocol {
exists = true
break
}
}
if !exists {
result = append(result, port)
}
}
return result
return firewall.ExcludePorts(ports, excluded)
}
const (
@@ -1,41 +0,0 @@
package service
import (
"errors"
"fmt"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/dto"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
func TestDatabaseSyncPlanMatchesOnceAndSummarizesStates(t *testing.T) {
desired := []databaseSyncDesired[int]{
{value: 1, item: dto.FirewallRuleSyncItem{SourceUUID: "existing"}},
{value: 1, item: dto.FirewallRuleSyncItem{SourceUUID: "duplicate"}},
{value: 2, item: dto.FirewallRuleSyncItem{SourceUUID: "blocked"}, err: errors.New("invalid policy")},
}
plan := buildDatabaseSyncPlan(
"test", filter.ProviderNftables, desired, []int{1, 3},
func(value int) string { return fmt.Sprint(value) },
func(value int) dto.FirewallRuleSyncItem {
return dto.FirewallRuleSyncItem{SourceUUID: "actual"}
},
)
preview := plan.preview()
if preview.Total != 3 || preview.Ready != 1 || preview.Existing != 1 ||
preview.Blocked != 1 || preview.Removed != 1 {
t.Fatalf("unexpected preview: %#v", preview)
}
completed := plan.completedResult()
if completed.Total != 3 || completed.Succeeded != 1 || completed.Skipped != 1 ||
completed.Failed != 1 || completed.Removed != 1 || len(completed.Errors) != 1 {
t.Fatalf("unexpected completed result: %#v", completed)
}
failed := plan.failedResult(errors.New("reconcile failed"))
if failed.Succeeded != 0 || failed.Skipped != 1 || failed.Failed != 2 || failed.Removed != 0 ||
len(failed.Errors) != 2 {
t.Fatalf("unexpected failed result: %#v", failed)
}
}
+16 -185
View File
@@ -7,7 +7,6 @@ import (
"fmt"
"net/netip"
"sort"
"strconv"
"strings"
"sync"
@@ -33,16 +32,7 @@ const (
dockerGuardComposeCreatedBy = "createdBy"
)
type dockerGuardRuntime interface {
Initialize([]docker_guard.Policy) error
Bind() error
Reconcile([]docker_guard.Policy) error
Unbind() error
Cleanup() error
Initialized(string) (bool, error)
Status(string) docker_guard.FamilyStatus
ListPolicies() ([]docker_guard.Policy, error)
}
type dockerGuardRuntime = docker_guard.Runtime
type DockerPortGuardService struct {
policies repo.IDockerPortGuardRepo
@@ -52,20 +42,12 @@ type DockerPortGuardService struct {
version func(string) string
}
type normalizedDockerGuardPolicy struct {
Family string
HostIP string
HostPort uint16
Protocol string
Mode string
}
var (
dockerPortGuardServiceMu sync.Mutex
dockerPortGuardSyncMu sync.RWMutex
dockerPortGuardSyncErr error
ErrDockerGuardInvalid = errors.New("invalid Docker port guard request")
ErrDockerUnavailable = errors.New("Docker is unavailable")
ErrDockerGuardInvalid = docker_guard.ErrInvalidPolicy
ErrDockerUnavailable = docker.ErrUnavailable
)
type IDockerPortGuardService interface {
@@ -156,7 +138,7 @@ func matchDockerGuardPolicies(
if !ok {
continue
}
endpoints[i].PolicyUUID, endpoints[i].Mode, endpoints[i].Sources = policy.UUID, policy.Mode, decodeGuardSources(policy.Sources)
endpoints[i].PolicyUUID, endpoints[i].Mode, endpoints[i].Sources = policy.UUID, policy.Mode, docker_guard.DecodeSources(policy.Sources)
endpoints[i].Description = policy.Description
endpoints[i].Effective = (policy.Family == docker_guard.FamilyIPv4 && base.IPv4.Effective) || (policy.Family == docker_guard.FamilyIPv6 && base.IPv6.Effective)
delete(byEndpoint, key)
@@ -165,7 +147,7 @@ func matchDockerGuardPolicies(
for _, policy := range byEndpoint {
orphanPolicies = append(orphanPolicies, dto.DockerPortGuardEndpoint{
Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol,
PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources), Description: policy.Description,
PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources), Description: policy.Description,
})
}
return endpoints, orphanPolicies
@@ -224,7 +206,7 @@ func (s *DockerPortGuardService) Operate(ctx context.Context, request dto.Docker
func (s *DockerPortGuardService) DeletePolicies(ctx context.Context, request dto.DockerPortGuardPolicyBatchDelete) error {
dockerPortGuardServiceMu.Lock()
defer dockerPortGuardServiceMu.Unlock()
uuids, err := normalizeDockerGuardPolicyUUIDs(request.UUIDs)
uuids, err := docker_guard.NormalizePolicyUUIDs(request.UUIDs)
if err != nil {
return err
}
@@ -234,33 +216,16 @@ func (s *DockerPortGuardService) DeletePolicies(ctx context.Context, request dto
return s.reconcileLocked(ctx)
}
func normalizeDockerGuardPolicyUUIDs(values []string) ([]string, error) {
uuids := make([]string, 0, len(values))
seen := make(map[string]struct{}, len(values))
for _, policyUUID := range values {
policyUUID = strings.TrimSpace(policyUUID)
if policyUUID == "" {
return nil, fmt.Errorf("%w: policy UUID cannot be empty", ErrDockerGuardInvalid)
}
if _, exists := seen[policyUUID]; exists {
continue
}
seen[policyUUID] = struct{}{}
uuids = append(uuids, policyUUID)
}
if len(uuids) == 0 {
return nil, fmt.Errorf("%w: policy UUIDs cannot be empty", ErrDockerGuardInvalid)
}
return uuids, nil
}
func (s *DockerPortGuardService) UpsertPolicies(ctx context.Context, request dto.DockerPortGuardPolicyBatch) error {
dockerPortGuardServiceMu.Lock()
defer dockerPortGuardServiceMu.Unlock()
policies := make([]model.DockerPortGuardPolicy, 0, len(request.Endpoints))
seen := make(map[string]struct{}, len(request.Endpoints))
for _, endpoint := range request.Endpoints {
normalized, sources, err := normalizeGuardPolicy(endpoint.Family, endpoint.HostIP, endpoint.HostPort, endpoint.Protocol, request.Mode, request.Sources)
normalized, err := docker_guard.NormalizePolicy(docker_guard.Policy{
Family: endpoint.Family, HostIP: endpoint.HostIP, HostPort: endpoint.HostPort,
Protocol: endpoint.Protocol, Mode: request.Mode, Sources: request.Sources,
})
if err != nil {
return err
}
@@ -269,7 +234,7 @@ func (s *DockerPortGuardService) UpsertPolicies(ctx context.Context, request dto
continue
}
seen[key] = struct{}{}
encoded, _ := json.Marshal(sources)
encoded, _ := json.Marshal(normalized.Sources)
policies = append(policies, model.DockerPortGuardPolicy{
UUID: uuid.NewString(), Family: normalized.Family, HostIP: normalized.HostIP,
HostPort: normalized.HostPort, Protocol: normalized.Protocol, Mode: normalized.Mode,
@@ -348,92 +313,14 @@ func dockerGuardPoliciesFromModels(policies []model.DockerPortGuardPolicy) []doc
func dockerGuardPolicyFromModel(policy model.DockerPortGuardPolicy) docker_guard.Policy {
return docker_guard.Policy{
UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort,
Protocol: policy.Protocol, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources),
Protocol: policy.Protocol, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources),
}
}
func dockerGuardPolicySyncKey(policy docker_guard.Policy) string {
mode := policy.Mode
if mode == docker_guard.ModeAllow && len(policy.Sources) == 0 {
mode = docker_guard.ModeAll
}
sources := append([]string(nil), policy.Sources...)
sort.Strings(sources)
return strings.Join([]string{
policy.UUID, policy.Family, canonicalGuardHost(policy.HostIP), strconv.Itoa(int(policy.HostPort)),
policy.Protocol, mode, strings.Join(sources, ","),
}, "\x00")
}
func canonicalGuardHost(value string) string {
if address, err := netip.ParseAddr(value); err == nil {
return address.String()
}
return value
}
func verifyDockerGuardRuleSync(runtime dockerGuardRuntime, desired []docker_guard.Policy) error {
actual, err := runtime.ListPolicies()
if err != nil {
return fmt.Errorf("verify synchronized Docker firewall policies: %w", err)
}
if !databaseSyncStatesEqual(actual, desired, dockerGuardPolicySyncKey) {
return fmt.Errorf("verify synchronized Docker firewall policies: target policies do not match the database")
}
return nil
}
func reconcileDockerGuardSyncTarget(backend string, policies []docker_guard.Policy, runtime dockerGuardRuntime) error {
families := make(map[string]struct{}, len(policies))
needsInitialize, needsBind := false, false
for _, policy := range policies {
families[policy.Family] = struct{}{}
}
if len(families) == 0 {
initialized := false
for _, family := range []string{docker_guard.FamilyIPv4, docker_guard.FamilyIPv6} {
status := runtime.Status(family)
if status.Reason == docker_guard.ReasonInspectFailed {
return fmt.Errorf("inspect Docker firewall target %s for %s failed", backend, family)
}
initialized = initialized || status.Initialized
}
if initialized {
return runtime.Reconcile(nil)
}
return nil
}
for family := range families {
status := runtime.Status(family)
needsInitialize = needsInitialize || !status.Initialized
needsBind = needsBind || !status.Bound || !status.Effective
}
var err error
if needsInitialize {
err = runtime.Initialize(policies)
} else {
if needsBind {
err = runtime.Bind()
}
if err == nil {
err = runtime.Reconcile(policies)
}
}
if err != nil {
return err
}
for family := range families {
if !runtime.Status(family).Effective {
return fmt.Errorf("Docker firewall target %s is not effective for %s", backend, family)
}
}
return nil
}
func dockerGuardRuleSyncDTO(policy model.DockerPortGuardPolicy) *dto.DockerPortGuardEndpoint {
return &dto.DockerPortGuardEndpoint{
Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol,
PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources), Description: policy.Description,
PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources), Description: policy.Description,
}
}
@@ -551,7 +438,7 @@ func (s *DockerPortGuardService) runtimePolicies(ctx context.Context) ([]docker_
}
policies := make([]docker_guard.Policy, 0, len(stored))
for _, policy := range stored {
policies = append(policies, docker_guard.Policy{UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources)})
policies = append(policies, docker_guard.Policy{UUID: policy.UUID, Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources)})
}
return policies, nil
}
@@ -577,10 +464,7 @@ func (s *DockerPortGuardService) guardRuntime(backend string) dockerGuardRuntime
if s.runtime != nil {
return s.runtime
}
if backend == constant.FirewallProviderNftables {
return docker_guard.NewNftablesManager()
}
return docker_guard.NewManager()
return docker_guard.NewRuntime(backend)
}
func (s *DockerPortGuardService) runtimeForDocker(ctx context.Context) (dockerGuardRuntime, string, error) {
@@ -668,65 +552,12 @@ func discoverDockerEndpoints(ctx context.Context, cli *client.Client) ([]dto.Doc
return endpoints, nil
}
func normalizeGuardPolicy(family, hostIP string, hostPort uint16, protocol, mode string, sources []string) (normalizedDockerGuardPolicy, []string, error) {
family, hostIP, protocol, mode = strings.ToLower(strings.TrimSpace(family)), strings.TrimSpace(hostIP), strings.ToLower(strings.TrimSpace(protocol)), strings.ToLower(strings.TrimSpace(mode))
if hostPort == 0 || (protocol != "tcp" && protocol != "udp") || (family != docker_guard.FamilyIPv4 && family != docker_guard.FamilyIPv6) || (mode != docker_guard.ModeAll && mode != docker_guard.ModeSources && mode != docker_guard.ModeAllow) {
return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: invalid policy fields", ErrDockerGuardInvalid)
}
addr, err := netip.ParseAddr(hostIP)
if err != nil || (family == docker_guard.FamilyIPv4) != addr.Is4() {
return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: host IP does not match address family", ErrDockerGuardInvalid)
}
normalizedSources := make([]string, 0, len(sources))
seen := map[string]struct{}{}
for _, source := range sources {
source = strings.TrimSpace(source)
if source == "" {
continue
}
prefix, err := netip.ParsePrefix(source)
if err != nil {
if sourceAddr, addrErr := netip.ParseAddr(source); addrErr == nil {
bits := 128
if sourceAddr.Is4() {
bits = 32
}
prefix = netip.PrefixFrom(sourceAddr, bits)
} else {
return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: invalid source address %q", ErrDockerGuardInvalid, source)
}
}
if (family == docker_guard.FamilyIPv4) != prefix.Addr().Is4() {
return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: source %q does not match address family", ErrDockerGuardInvalid, source)
}
canonical := prefix.Masked().String()
if _, ok := seen[canonical]; !ok {
seen[canonical] = struct{}{}
normalizedSources = append(normalizedSources, canonical)
}
}
if mode == docker_guard.ModeSources && len(normalizedSources) == 0 {
return normalizedDockerGuardPolicy{}, nil, fmt.Errorf("%w: deny_sources requires at least one source", ErrDockerGuardInvalid)
}
if mode == docker_guard.ModeAll {
normalizedSources = []string{}
}
sort.Strings(normalizedSources)
return normalizedDockerGuardPolicy{Family: family, HostIP: hostIP, HostPort: hostPort, Protocol: protocol, Mode: mode}, normalizedSources, nil
}
func decodeGuardSources(value string) []string {
result := []string{}
_ = json.Unmarshal([]byte(value), &result)
return result
}
func dockerGuardPolicyEndpoints(policies []model.DockerPortGuardPolicy) []dto.DockerPortGuardEndpoint {
endpoints := make([]dto.DockerPortGuardEndpoint, 0, len(policies))
for _, policy := range policies {
endpoints = append(endpoints, dto.DockerPortGuardEndpoint{
Family: policy.Family, HostIP: policy.HostIP, HostPort: policy.HostPort, Protocol: policy.Protocol,
PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: decodeGuardSources(policy.Sources),
PolicyUUID: policy.UUID, Mode: policy.Mode, Sources: docker_guard.DecodeSources(policy.Sources),
Description: policy.Description,
})
}
-572
View File
@@ -1,572 +0,0 @@
package service
import (
"context"
"errors"
"fmt"
"reflect"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/dto"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/global"
agenti18n "github.com/1Panel-dev/1Panel/agent/i18n"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
"github.com/docker/docker/client"
"github.com/glebarez/sqlite"
"github.com/google/uuid"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
type persistentDockerGuardRuntime struct {
initialized bool
initialize int
reconcile int
bind int
unbind int
policies []docker_guard.Policy
statuses map[string]docker_guard.FamilyStatus
}
func (r *persistentDockerGuardRuntime) Initialize(policies []docker_guard.Policy) error {
r.initialize++
r.initialized = true
r.policies = append([]docker_guard.Policy(nil), policies...)
r.markFamiliesEffective(policies)
return nil
}
func (r *persistentDockerGuardRuntime) Bind() error {
r.bind++
r.initialized = true
for family, status := range r.statuses {
status.Initialized, status.Bound, status.Effective = true, true, true
r.statuses[family] = status
}
return nil
}
func (r *persistentDockerGuardRuntime) Reconcile(policies []docker_guard.Policy) error {
r.reconcile++
r.policies = append([]docker_guard.Policy(nil), policies...)
return nil
}
func (r *persistentDockerGuardRuntime) markFamiliesEffective(policies []docker_guard.Policy) {
if r.statuses == nil {
return
}
for _, policy := range policies {
r.statuses[policy.Family] = docker_guard.FamilyStatus{Initialized: true, Bound: true, Effective: true}
}
}
func (r *persistentDockerGuardRuntime) Unbind() error {
r.unbind++
return nil
}
func (r *persistentDockerGuardRuntime) Cleanup() error { return nil }
func (r *persistentDockerGuardRuntime) Initialized(string) (bool, error) {
return r.initialized, nil
}
func (r *persistentDockerGuardRuntime) Status(family string) docker_guard.FamilyStatus {
if r.statuses != nil {
return r.statuses[family]
}
return docker_guard.FamilyStatus{Initialized: r.initialized, Bound: r.initialized, Effective: r.initialized}
}
func (r *persistentDockerGuardRuntime) ListPolicies() ([]docker_guard.Policy, error) {
return append([]docker_guard.Policy(nil), r.policies...), nil
}
type persistentDockerGuardPolicies struct{ items []model.DockerPortGuardPolicy }
func (r *persistentDockerGuardPolicies) List(context.Context) ([]model.DockerPortGuardPolicy, error) {
return append([]model.DockerPortGuardPolicy(nil), r.items...), nil
}
func (r *persistentDockerGuardPolicies) DeleteBatch(context.Context, []string) error { return nil }
func (r *persistentDockerGuardPolicies) UpsertBatch(context.Context, []model.DockerPortGuardPolicy) error {
return nil
}
func TestDockerGuardOverviewLocalizesUnavailableDocker(t *testing.T) {
agenti18n.Init()
service := &DockerPortGuardService{
policies: &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{{
UUID: "orphan-policy", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080,
Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]",
}}},
runtime: &persistentDockerGuardRuntime{},
version: func(string) string { return "1.8.10" },
client: func() (*client.Client, error) {
return nil, errors.New("Cannot connect to the Docker daemon at unix:///var/run/docker.sock")
},
}
overview, err := service.LoadOverview(context.Background())
if err != nil {
t.Fatalf("load overview: %v", err)
}
if overview.Base.Message != agenti18n.Get("ErrDockerFailed") {
t.Fatalf("message = %q, want localized Docker failure", overview.Base.Message)
}
if overview.Base.Version != "1.8.10" {
t.Fatalf("version = %q, want 1.8.10", overview.Base.Version)
}
if len(overview.Containers) != 0 || len(overview.OrphanPolicies) != 1 {
t.Fatalf("overview = %#v, want persisted policy returned separately from Docker containers", overview)
}
if overview.OrphanPolicies[0].PolicyUUID != "orphan-policy" {
t.Fatalf("orphan policy = %#v", overview.OrphanPolicies[0])
}
}
func TestDockerGuardRuntimeStatusAggregatesAvailableFamilies(t *testing.T) {
runtime := &persistentDockerGuardRuntime{statuses: map[string]docker_guard.FamilyStatus{
docker_guard.FamilyIPv6: {Initialized: true, Bound: true, Effective: true},
}}
base := (&DockerPortGuardService{}).runtimeStatus(runtime, constant.FirewallProviderNftables)
if !base.Initialized || !base.Bound || base.IPv4.Initialized || !base.IPv6.Initialized {
t.Fatalf("unexpected aggregate Docker guard status: %#v", base)
}
}
func TestDockerGuardRuntimeStatusReportsMissingBackend(t *testing.T) {
runtime := &persistentDockerGuardRuntime{statuses: map[string]docker_guard.FamilyStatus{
docker_guard.FamilyIPv4: {Reason: docker_guard.ReasonCommandMissing},
docker_guard.FamilyIPv6: {Reason: docker_guard.ReasonCommandMissing},
}}
base := (&DockerPortGuardService{}).runtimeStatus(runtime, constant.FirewallProviderIptables)
if base.IsExist {
t.Fatalf("missing Docker firewall backend reported as installed: %#v", base)
}
}
func TestMatchDockerGuardPoliciesReturnsUnmatchedDatabaseRules(t *testing.T) {
policies := []model.DockerPortGuardPolicy{
{UUID: "matched", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]"},
{UUID: "orphan", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 9090, Protocol: "tcp", Mode: docker_guard.ModeSources, Sources: `["203.0.113.0/24"]`},
}
endpoints := []dto.DockerPortGuardEndpoint{{
Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, Protocol: "tcp", ContainerID: "container-1",
}}
matched, orphanPolicies := matchDockerGuardPolicies(dto.DockerPortGuardBase{}, policies, endpoints)
if len(matched) != 1 || matched[0].PolicyUUID != "matched" {
t.Fatalf("matched endpoints = %#v", matched)
}
if len(orphanPolicies) != 1 || orphanPolicies[0].PolicyUUID != "orphan" || orphanPolicies[0].HostPort != 9090 {
t.Fatalf("orphan policies = %#v", orphanPolicies)
}
}
func TestDockerGuardRuleSyncInitializesTargetWithPersistedPolicies(t *testing.T) {
setupDockerGuardSettingsDB(t)
selectDockerGuardBackend(t, constant.FirewallProviderNftables)
target := &persistentDockerGuardRuntime{}
policies := &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{{
UUID: "policy-1", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080,
Protocol: "tcp", Mode: docker_guard.ModeSources, Sources: `["203.0.113.0/24"]`,
}}}
service := &DockerPortGuardService{
policies: policies,
runtimeForBackend: func(string) dockerGuardRuntime { return target },
}
request := dto.FirewallRuleSyncRequest{Subsystem: "docker", TargetProvider: filter.ProviderNftables}
preview, err := service.previewRuleSync(context.Background(), request)
if err != nil {
t.Fatal(err)
}
if preview.Total != 1 || preview.Ready != 1 || preview.TargetProvider != filter.ProviderNftables || preview.Items[0].DockerRule == nil {
t.Fatalf("unexpected preview: %#v", preview)
}
result, err := service.syncRules(context.Background(), request)
if err != nil {
t.Fatal(err)
}
if result.Succeeded != 1 || result.Failed != 0 || target.initialize != 1 || len(target.policies) != 1 {
t.Fatalf("unexpected sync result=%#v target=%#v", result, target)
}
retry, err := service.previewRuleSync(context.Background(), request)
if err != nil {
t.Fatal(err)
}
if retry.Ready != 0 || retry.Existing != 1 || retry.Removed != 0 {
t.Fatalf("synchronized Docker policy was not recognized: %#v", retry)
}
retryResult, err := service.syncRules(context.Background(), request)
if err != nil {
t.Fatal(err)
}
if retryResult.Succeeded != 0 || retryResult.Skipped != 1 || retryResult.Removed != 0 {
t.Fatalf("existing Docker policy was not counted as skipped: %#v", retryResult)
}
}
func TestDockerGuardRuleSyncReconcilesInitializedEffectiveTarget(t *testing.T) {
setupDockerGuardSettingsDB(t)
selectDockerGuardBackend(t, constant.FirewallProviderNftables)
policy := model.DockerPortGuardPolicy{
UUID: "policy-1", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080,
Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]",
}
target := &persistentDockerGuardRuntime{initialized: true}
service := &DockerPortGuardService{
policies: &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{policy}},
runtimeForBackend: func(string) dockerGuardRuntime { return target },
}
result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "docker", TargetProvider: filter.ProviderNftables,
})
if err != nil {
t.Fatal(err)
}
if result.Succeeded != 1 || result.Skipped != 0 || target.reconcile != 1 || len(target.policies) != 1 {
t.Fatalf("initialized target skipped runtime reconciliation: result=%#v target=%#v", result, target)
}
}
func TestDockerGuardRuleSyncRebindsInitializedIneffectiveTarget(t *testing.T) {
setupDockerGuardSettingsDB(t)
selectDockerGuardBackend(t, constant.FirewallProviderNftables)
policy := model.DockerPortGuardPolicy{
UUID: "policy-1", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080,
Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]",
}
target := &persistentDockerGuardRuntime{
initialized: true,
statuses: map[string]docker_guard.FamilyStatus{
docker_guard.FamilyIPv4: {Initialized: true},
},
}
service := &DockerPortGuardService{
policies: &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{policy}},
runtimeForBackend: func(string) dockerGuardRuntime { return target },
}
result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "docker", TargetProvider: filter.ProviderNftables,
})
if err != nil {
t.Fatal(err)
}
if result.Succeeded != 1 || target.bind != 1 || target.reconcile != 1 ||
!target.Status(docker_guard.FamilyIPv4).Effective {
t.Fatalf("initialized target was not rebound and reconciled: result=%#v target=%#v", result, target)
}
}
func TestDockerGuardRuleSyncClearsInitializedTargetWhenDatabaseIsEmpty(t *testing.T) {
setupDockerGuardSettingsDB(t)
selectDockerGuardBackend(t, constant.FirewallProviderNftables)
target := &persistentDockerGuardRuntime{
initialized: true,
policies: []docker_guard.Policy{{
UUID: "stale-policy", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080,
Protocol: "tcp", Mode: docker_guard.ModeAll,
}},
}
service := &DockerPortGuardService{
policies: &persistentDockerGuardPolicies{},
runtimeForBackend: func(string) dockerGuardRuntime { return target },
}
preview, err := service.previewRuleSync(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "docker", TargetProvider: filter.ProviderNftables,
})
if err != nil {
t.Fatal(err)
}
if preview.Total != 0 || preview.Removed != 1 || len(preview.Items) != 1 || preview.Items[0].Status != "remove" {
t.Fatalf("stale runtime policy was not included in preview: %#v", preview)
}
result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "docker", TargetProvider: filter.ProviderNftables,
})
if err != nil {
t.Fatal(err)
}
if result.Total != 0 || result.Removed != 1 || target.initialize != 0 || target.reconcile != 1 || len(target.policies) != 0 {
t.Fatalf("empty database did not clear initialized target: result=%#v target=%#v", result, target)
}
}
func TestDockerGuardRuleSyncLeavesUninitializedTargetEmpty(t *testing.T) {
setupDockerGuardSettingsDB(t)
selectDockerGuardBackend(t, constant.FirewallProviderNftables)
target := &persistentDockerGuardRuntime{}
service := &DockerPortGuardService{
policies: &persistentDockerGuardPolicies{},
runtimeForBackend: func(string) dockerGuardRuntime { return target },
}
result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "docker", TargetProvider: filter.ProviderNftables,
})
if err != nil {
t.Fatal(err)
}
if result.Total != 0 || target.initialize != 0 || target.reconcile != 0 {
t.Fatalf("empty database initialized an unused target: result=%#v target=%#v", result, target)
}
}
func TestDockerGuardRuleSyncRejectsUnselectedTarget(t *testing.T) {
setupDockerGuardSettingsDB(t)
selectDockerGuardBackend(t, constant.FirewallProviderIptables)
target := &persistentDockerGuardRuntime{}
service := &DockerPortGuardService{
policies: &persistentDockerGuardPolicies{},
runtimeForBackend: func(string) dockerGuardRuntime { return target },
}
request := dto.FirewallRuleSyncRequest{Subsystem: "docker", TargetProvider: filter.ProviderNftables}
if _, err := service.previewRuleSync(context.Background(), request); !errors.Is(err, filter.ErrProviderUnavailable) {
t.Fatalf("preview error = %v, want provider unavailable", err)
}
if _, err := service.syncRules(context.Background(), request); !errors.Is(err, filter.ErrProviderUnavailable) {
t.Fatalf("sync error = %v, want provider unavailable", err)
}
if target.initialize != 0 || target.bind != 0 || target.reconcile != 0 {
t.Fatalf("unselected target was modified: %#v", target)
}
}
func selectDockerGuardBackend(t *testing.T, backend string) {
t.Helper()
if err := settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, backend); err != nil {
t.Fatal(err)
}
}
func setupDockerGuardSettingsDB(t *testing.T) {
t.Helper()
previousDB := global.DB
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString())
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatalf("open settings database: %v", err)
}
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatalf("migrate settings database: %v", err)
}
global.DB = db
t.Cleanup(func() { global.DB = previousDB })
}
func TestNormalizeDockerPortGuardPolicy(t *testing.T) {
policy, sources, err := normalizeGuardPolicy("ipv4", "0.0.0.0", 8080, "TCP", "deny_sources", []string{"203.0.113.10", "192.0.2.7/24", "203.0.113.10/32"})
if err != nil {
t.Fatal(err)
}
if policy.Protocol != "tcp" {
t.Fatalf("protocol was not normalized: %#v", policy)
}
want := []string{"192.0.2.0/24", "203.0.113.10/32"}
if !reflect.DeepEqual(sources, want) {
t.Fatalf("sources = %#v, want %#v", sources, want)
}
}
func TestNormalizeDockerPortGuardPolicyRejectsInvalidSources(t *testing.T) {
for _, test := range []struct {
name string
family string
hostIP string
mode string
sources []string
}{
{name: "empty deny sources", family: "ipv4", hostIP: "0.0.0.0", mode: "deny_sources"},
{name: "mixed source family", family: "ipv4", hostIP: "0.0.0.0", mode: "deny_sources", sources: []string{"2001:db8::/64"}},
{name: "mixed host family", family: "ipv6", hostIP: "0.0.0.0", mode: "deny_all"},
} {
t.Run(test.name, func(t *testing.T) {
if _, _, err := normalizeGuardPolicy(test.family, test.hostIP, 80, "tcp", test.mode, test.sources); !errors.Is(err, ErrDockerGuardInvalid) {
t.Fatalf("expected typed validation error, got %v", err)
}
})
}
}
func TestNormalizeDockerPortGuardPolicyAllowsEmptyAllowList(t *testing.T) {
policy, sources, err := normalizeGuardPolicy("ipv4", "0.0.0.0", 5432, "tcp", "allow_sources", nil)
if err != nil {
t.Fatal(err)
}
if policy.Mode != "allow_sources" || len(sources) != 0 {
t.Fatalf("policy = %#v, sources = %#v", policy, sources)
}
}
func TestDockerGuardEndpointKeyIncludesAddressFamilyAndProtocol(t *testing.T) {
first := guardEndpointKey("ipv4", "0.0.0.0", 53, "udp")
if first == guardEndpointKey("ipv6", "::", 53, "udp") || first == guardEndpointKey("ipv4", "0.0.0.0", 53, "tcp") {
t.Fatal("endpoint identity collapsed distinct endpoint dimensions")
}
}
func TestNormalizeDockerGuardPolicyUUIDs(t *testing.T) {
got, err := normalizeDockerGuardPolicyUUIDs([]string{" first ", "second", "first"})
if err != nil {
t.Fatal(err)
}
if want := []string{"first", "second"}; !reflect.DeepEqual(got, want) {
t.Fatalf("UUIDs = %#v, want %#v", got, want)
}
if _, err := normalizeDockerGuardPolicyUUIDs([]string{""}); err == nil {
t.Fatal("expected empty UUID to be rejected")
}
}
func TestGroupDockerGuardContainersMergesCompatiblePorts(t *testing.T) {
endpoints := []dto.DockerPortGuardEndpoint{
{Family: "ipv4", HostIP: "0.0.0.0", HostPort: 8001, Protocol: "tcp", ContainerID: "container-1", ContainerName: "demo", ContainerPort: 81, Mode: "deny_all", PolicyUUID: "policy-2", Effective: true, Sources: []string{}},
{Family: "ipv4", HostIP: "0.0.0.0", HostPort: 8000, Protocol: "tcp", ContainerID: "container-1", ContainerName: "demo", ContainerPort: 80, Mode: "deny_all", PolicyUUID: "policy-1", Effective: true, Sources: []string{}},
{Family: "ipv4", HostIP: "0.0.0.0", HostPort: 8002, Protocol: "tcp", ContainerID: "container-1", ContainerName: "demo", ContainerPort: 82, Mode: "allow_sources", PolicyUUID: "policy-3", Effective: true, Sources: []string{"192.0.2.1/32"}},
}
containers := groupDockerGuardContainers(endpoints)
if len(containers) != 1 || len(containers[0].PortGroups) != 2 {
t.Fatalf("containers = %#v, want one container with two port groups", containers)
}
foundRange := false
for _, group := range containers[0].PortGroups {
foundRange = foundRange || group.Label == "0.0.0.0:8000-8001/tcp"
}
if !foundRange {
t.Fatalf("port groups = %#v, expected merged range", containers[0].PortGroups)
}
for _, group := range containers[0].PortGroups {
if group.Label == "0.0.0.0:8000-8001/tcp" && len(group.Endpoints) != 2 {
t.Fatalf("merged group endpoints = %#v, want 2 endpoints", group.Endpoints)
}
}
}
func TestMarkDockerGuardReconcileFailureOnlyAffectsFailedFamily(t *testing.T) {
base := dto.DockerPortGuardBase{
IPv4: dto.DockerPortGuardFamilyStatus{State: docker_guard.StatusEffective, Initialized: true, Bound: true, Effective: true},
IPv6: dto.DockerPortGuardFamilyStatus{State: docker_guard.StatusEffective, Initialized: true, Bound: true, Effective: true},
}
markDockerGuardReconcileFailure(&base, &docker_guard.FamilyError{Family: docker_guard.FamilyIPv6, Err: errors.New("restore failed")})
if !base.IPv4.Effective || base.IPv4.State != docker_guard.StatusEffective {
t.Fatalf("IPv4 status changed unexpectedly: %#v", base.IPv4)
}
if base.IPv6.Effective || base.IPv6.State != docker_guard.StatusNotEffective || base.IPv6.Reason != docker_guard.ReasonInspectFailed {
t.Fatalf("IPv6 status = %#v, want not effective", base.IPv6)
}
}
func TestMarkDockerGuardIPv4FailureAlsoMarksUnattemptedIPv6(t *testing.T) {
base := dto.DockerPortGuardBase{
IPv4: dto.DockerPortGuardFamilyStatus{State: docker_guard.StatusEffective, Initialized: true, Bound: true, Effective: true},
IPv6: dto.DockerPortGuardFamilyStatus{State: docker_guard.StatusEffective, Initialized: true, Bound: true, Effective: true},
}
markDockerGuardReconcileFailure(&base, &docker_guard.FamilyError{Family: docker_guard.FamilyIPv4, Err: errors.New("restore failed")})
if base.IPv4.Effective || base.IPv6.Effective {
t.Fatalf("statuses = IPv4 %#v, IPv6 %#v; both must be not effective", base.IPv4, base.IPv6)
}
}
func TestDockerGuardReconcileErrorState(t *testing.T) {
t.Cleanup(func() { recordDockerPortGuardReconcileError(nil) })
want := errors.New("restore failed")
recordDockerPortGuardReconcileError(want)
if got := lastDockerPortGuardReconcileError(); !errors.Is(got, want) {
t.Fatalf("last reconcile error = %v, want %v", got, want)
}
recordDockerPortGuardReconcileError(nil)
if got := lastDockerPortGuardReconcileError(); got != nil {
t.Fatalf("last reconcile error = %v, want nil", got)
}
}
func TestDockerFirewallDisplayName(t *testing.T) {
for backend, want := range map[string]string{
"iptables": "iptables-docker",
"nftables": "nftables-docker",
"": "iptables-docker",
} {
if got := dockerFirewallDisplayName(backend); got != want {
t.Fatalf("dockerFirewallDisplayName(%q) = %q, want %q", backend, got, want)
}
}
}
func TestDockerGuardRuntimeMatchesDockerBackend(t *testing.T) {
service := &DockerPortGuardService{}
if _, ok := service.guardRuntime("iptables").(*docker_guard.Manager); !ok {
t.Fatal("iptables backend did not select the iptables Docker guard runtime")
}
if _, ok := service.guardRuntime("nftables").(*docker_guard.NftablesManager); !ok {
t.Fatal("nftables backend did not select the nftables Docker guard runtime")
}
}
func TestDockerGuardInitializeAndUnbindPersistStatus(t *testing.T) {
setupDockerGuardSettingsDB(t)
runtime := &persistentDockerGuardRuntime{}
service := &DockerPortGuardService{
policies: &persistentDockerGuardPolicies{},
runtime: runtime,
}
if err := service.Operate(context.Background(), dto.DockerPortGuardOperation{Operation: "initialize"}); err != nil {
t.Fatalf("initialize Docker port guard: %v", err)
}
status, err := settingRepo.GetValueByKey(constant.FirewallDockerPortGuardStatusKey)
if err != nil || status != constant.StatusEnable {
t.Fatalf("persisted status = %q, %v; want %q", status, err, constant.StatusEnable)
}
if runtime.initialize != 1 {
t.Fatalf("initialize calls = %d, want 1", runtime.initialize)
}
if err := service.Operate(context.Background(), dto.DockerPortGuardOperation{Operation: "unbind"}); err != nil {
t.Fatalf("unbind Docker port guard: %v", err)
}
status, err = settingRepo.GetValueByKey(constant.FirewallDockerPortGuardStatusKey)
if err != nil || status != constant.StatusDisable {
t.Fatalf("persisted status = %q, %v; want %q", status, err, constant.StatusDisable)
}
}
func TestDockerGuardReconcileRestoresPersistedInitialization(t *testing.T) {
setupDockerGuardSettingsDB(t)
if err := settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusEnable); err != nil {
t.Fatal(err)
}
runtime := &persistentDockerGuardRuntime{}
service := &DockerPortGuardService{
policies: &persistentDockerGuardPolicies{items: []model.DockerPortGuardPolicy{{
UUID: "policy-1", Family: docker_guard.FamilyIPv4, HostIP: "0.0.0.0",
HostPort: 8080, Protocol: "tcp", Mode: docker_guard.ModeAll, Sources: "[]",
}}},
runtime: runtime,
}
if err := service.Reconcile(context.Background()); err != nil {
t.Fatalf("restore Docker port guard: %v", err)
}
if runtime.initialize != 1 || !runtime.initialized {
t.Fatalf("runtime was not initialized: %#v", runtime)
}
if len(runtime.policies) != 1 || runtime.policies[0].UUID != "policy-1" {
t.Fatalf("restored policies = %#v", runtime.policies)
}
}
func TestDockerGuardReconcileLeavesDisabledGuardUninitialized(t *testing.T) {
setupDockerGuardSettingsDB(t)
if err := settingRepo.UpdateOrCreate(constant.FirewallDockerPortGuardStatusKey, constant.StatusDisable); err != nil {
t.Fatal(err)
}
runtime := &persistentDockerGuardRuntime{}
service := &DockerPortGuardService{policies: &persistentDockerGuardPolicies{}, runtime: runtime}
if err := service.Reconcile(context.Background()); err != nil {
t.Fatalf("reconcile disabled Docker port guard: %v", err)
}
if runtime.initialize != 0 || runtime.initialized {
t.Fatalf("disabled guard was initialized: %#v", runtime)
}
}
File diff suppressed because it is too large Load Diff
+10 -16
View File
@@ -2,6 +2,7 @@ package service
import (
"context"
"errors"
"fmt"
"strings"
@@ -24,7 +25,9 @@ type IFirewallSettingService interface {
type FirewallSettingService struct{}
func NewIFirewallSettingService() IFirewallSettingService { return &FirewallSettingService{} }
func NewIFirewallSettingService() IFirewallSettingService {
return &FirewallSettingService{}
}
func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings, error) {
result := dto.FirewallSettings{PingStatus: ping.LoadStatus()}
@@ -132,10 +135,7 @@ func (s *FirewallSettingService) Load(ctx context.Context) (dto.FirewallSettings
Name: name, Installed: installed[name], Supported: dockerInstalled,
Active: dockerInstalled && installed[name] && result.Docker.Selected == name,
}
var guard dockerGuardRuntime = docker_guard.NewManager()
if name == constant.FirewallProviderNftables {
guard = docker_guard.NewNftablesManager()
}
guard := docker_guard.NewRuntime(name)
ipv4, ipv6 := guard.Status(docker_guard.FamilyIPv4), guard.Status(docker_guard.FamilyIPv6)
option.Initialized = ipv4.Initialized || ipv6.Initialized
option.Bound = ipv4.Bound || ipv6.Bound
@@ -196,10 +196,7 @@ func (s *FirewallSettingService) Operate(ctx context.Context, request dto.Firewa
}
func (s *FirewallSettingService) operateDocker(ctx context.Context, request dto.FirewallBackendOperation) error {
var guard dockerGuardRuntime = docker_guard.NewManager()
if request.Backend == constant.FirewallProviderNftables {
guard = docker_guard.NewNftablesManager()
}
guard := docker_guard.NewRuntime(request.Backend)
if request.Operation == "cleanup" {
if err := guard.Cleanup(); err != nil {
return err
@@ -224,7 +221,7 @@ func (s *FirewallSettingService) operateDocker(ctx context.Context, request dto.
return err
}
if request.Operation == "initialize" {
if err := NewIDockerPortGuardService().Operate(ctx, dto.DockerPortGuardOperation{Operation: "initialize"}); err != nil {
if err := newDockerPortGuardService().Operate(ctx, dto.DockerPortGuardOperation{Operation: "initialize"}); err != nil {
_ = settingRepo.UpdateOrCreate(constant.FirewallDockerBackendKey, previous)
return err
}
@@ -233,10 +230,7 @@ func (s *FirewallSettingService) operateDocker(ctx context.Context, request dto.
}
func dockerGuardBackendInitialized(backend string) (bool, error) {
var guard dockerGuardRuntime = docker_guard.NewManager()
if backend == constant.FirewallProviderNftables {
guard = docker_guard.NewNftablesManager()
}
guard := docker_guard.NewRuntime(backend)
for _, family := range []string{docker_guard.FamilyIPv4, docker_guard.FamilyIPv6} {
initialized, err := guard.Initialized(family)
if err != nil {
@@ -372,7 +366,7 @@ func (s *FirewallSettingService) operateForwarding(request dto.FirewallBackendOp
return err
}
if request.Operation == "initialize" {
return NewIForwardingService().Enable()
return newForwardingService().Enable()
}
recordForwardingSyncError(nil)
return nil
@@ -381,7 +375,7 @@ func (s *FirewallSettingService) operateForwarding(request dto.FirewallBackendOp
func forwardingBackendInitialized(backend string) (bool, error) {
manager, err := newForwardingManagerFor(backend)
if err != nil {
if strings.Contains(err.Error(), "is not installed") {
if errors.Is(err, lifecycle.ErrNotInstalled) {
return false, nil
}
return false, err
-174
View File
@@ -1,174 +0,0 @@
package service
import (
"context"
"fmt"
"os"
"path/filepath"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/dto"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/global"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard"
"github.com/glebarez/sqlite"
"github.com/google/uuid"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func TestFirewallSettingServiceKeepsSelectionsWhenBackendsAreMissing(t *testing.T) {
previousDB := global.DB
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString())
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatalf("open settings database: %v", err)
}
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatalf("migrate settings database: %v", err)
}
global.DB = db
t.Cleanup(func() { global.DB = previousDB })
binDir := t.TempDir()
if err := os.WriteFile(filepath.Join(binDir, "docker"), []byte("#!/bin/sh\nexit 0\n"), 0o755); err != nil {
t.Fatalf("create fake Docker executable: %v", err)
}
t.Setenv("PATH", binDir)
for key, backend := range map[string]string{
constant.FirewallSystemBackendKey: constant.FirewallProviderNftables,
constant.FirewallForwardingBackendKey: constant.FirewallProviderNftables,
constant.FirewallDockerBackendKey: constant.FirewallProviderNftables,
} {
if err := settingRepo.UpdateOrCreate(key, backend); err != nil {
t.Fatalf("save selected backend %s: %v", key, err)
}
}
settings, err := (&FirewallSettingService{}).Load(context.Background())
if err != nil {
t.Fatalf("load firewall settings: %v", err)
}
for subsystem, selected := range map[string]string{
"system": settings.System.Selected,
"forwarding": settings.Forwarding.Selected,
"docker": settings.Docker.Selected,
} {
if selected != constant.FirewallProviderNftables {
t.Fatalf("%s selected backend = %q, want persisted nftables", subsystem, selected)
}
}
for _, option := range settings.Docker.Options {
if option.Installed || option.Active || !option.Supported {
t.Fatalf("unexpected Docker backend availability without firewall commands: %#v", option)
}
}
}
func TestFirewallSettingServiceSeparatesDockerAndFirewallAvailability(t *testing.T) {
previousDB := global.DB
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString())
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatalf("open settings database: %v", err)
}
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatalf("migrate settings database: %v", err)
}
global.DB = db
t.Cleanup(func() { global.DB = previousDB })
binDir := t.TempDir()
for _, name := range []string{"iptables", "iptables-restore"} {
if err := os.WriteFile(filepath.Join(binDir, name), []byte("#!/bin/sh\nexit 0\n"), 0o755); err != nil {
t.Fatalf("create fake %s executable: %v", name, err)
}
}
t.Setenv("PATH", binDir)
settings, err := (&FirewallSettingService{}).Load(context.Background())
if err != nil {
t.Fatalf("load firewall settings: %v", err)
}
for _, option := range settings.Docker.Options {
if option.Name != constant.FirewallProviderIptables {
continue
}
if !option.Installed || option.Supported || option.Active {
t.Fatalf("Docker and firewall availability were not separated: %#v", option)
}
return
}
t.Fatal("iptables Docker backend option was not returned")
}
func TestFirewallSettingServiceUsesExplicitEmptyDockerAndIptablesForwardingDefaults(t *testing.T) {
previousDB := global.DB
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString())
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatalf("open settings database: %v", err)
}
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatalf("migrate settings database: %v", err)
}
global.DB = db
t.Cleanup(func() { global.DB = previousDB })
t.Setenv("PATH", t.TempDir())
settings, err := (&FirewallSettingService{}).Load(context.Background())
if err != nil {
t.Fatalf("load firewall settings: %v", err)
}
if settings.Docker.Selected != "" {
t.Fatalf("Docker selected backend = %q, want empty", settings.Docker.Selected)
}
if settings.Forwarding.Selected != constant.FirewallProviderIptables {
t.Fatalf("forwarding selected backend = %q, want iptables", settings.Forwarding.Selected)
}
}
func TestFirewallSettingServiceSelectsGuardBackendWithoutDockerRestart(t *testing.T) {
previousDB := global.DB
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", uuid.NewString())
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatalf("open settings database: %v", err)
}
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatalf("migrate settings database: %v", err)
}
global.DB = db
t.Cleanup(func() { global.DB = previousDB })
err = (&FirewallSettingService{}).Operate(context.Background(), dto.FirewallBackendOperation{
Subsystem: "docker",
Backend: "nftables",
Operation: "select",
})
if err != nil {
t.Fatalf("select Docker guard backend: %v", err)
}
if got := selectedDockerFirewallBackend("iptables"); got != "nftables" {
t.Fatalf("selected Docker guard backend = %q, want nftables", got)
}
if _, ok := (&DockerPortGuardService{}).guardRuntime(selectedDockerFirewallBackend("iptables")).(*docker_guard.NftablesManager); !ok {
t.Fatal("Docker port guard did not use the selected nftables runtime")
}
}
func TestFirewallSettingServiceRejectsServiceBackendInitialization(t *testing.T) {
for _, backend := range []string{"firewalld", "ufw"} {
for _, operation := range []string{"initialize", "cleanup"} {
err := (&FirewallSettingService{}).Operate(context.Background(), dto.FirewallBackendOperation{
Subsystem: "system",
Backend: backend,
Operation: operation,
})
if err == nil {
t.Fatalf("%s %s unexpectedly succeeded", backend, operation)
}
}
}
}
+60 -343
View File
@@ -4,7 +4,6 @@ import (
"context"
"errors"
"fmt"
"slices"
"sort"
"strings"
"sync"
@@ -20,6 +19,7 @@ import (
"github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding"
firewallsync "github.com/1Panel-dev/1Panel/agent/utils/firewall/sync"
"gorm.io/gorm"
)
@@ -29,9 +29,9 @@ var (
)
const (
firewallRuleSyncReady = "ready"
firewallRuleSyncExisting = "existing"
firewallRuleSyncBlocked = "blocked"
firewallRuleSyncReady = firewallsync.StatusReady
firewallRuleSyncExisting = firewallsync.StatusExisting
firewallRuleSyncBlocked = firewallsync.StatusBlocked
)
func firewallSyncSubsystem(value string) string {
@@ -42,13 +42,13 @@ func firewallSyncSubsystem(value string) string {
return value
}
type firewallRuleSyncOutcome string
type firewallRuleSyncOutcome = firewallsync.Outcome
const (
firewallRuleSyncApplied firewallRuleSyncOutcome = "applied"
firewallRuleSyncSkipped firewallRuleSyncOutcome = "skipped"
firewallRuleSyncRemoved firewallRuleSyncOutcome = "removed"
firewallRuleSyncFailed firewallRuleSyncOutcome = "failed"
firewallRuleSyncApplied = firewallsync.OutcomeApplied
firewallRuleSyncSkipped = firewallsync.OutcomeSkipped
firewallRuleSyncRemoved = firewallsync.OutcomeRemoved
firewallRuleSyncFailed = firewallsync.OutcomeFailed
)
type firewallRuleSyncEntry struct {
@@ -146,7 +146,7 @@ func (s *FirewallService) loadFirewallRuleSyncPlan(
expectedMarkers[scopeKey+"\x00"+entry.desired.Marker] = struct{}{}
}
scopes := firewallRuleSyncScopes(target)
scopes := filter.ManagedInputScopes(target)
if hasCompileErrors {
scopes = scopesWithFirewallSyncCandidates(scopes, entriesByScope)
}
@@ -203,14 +203,18 @@ func (s *FirewallService) loadFirewallRuleSyncPlan(
source: model.FirewallRule{UUID: strings.TrimPrefix(observed.Marker, "1panel-rule:")},
rule: observed.Rule, remove: &copy,
}
status, reason := "remove", "managed rule exists only in target backend"
status := firewallsync.StatusRemove
reasonCode := firewallsync.ReasonManagedOnlyInTarget
reason := firewallsync.ReasonMessage(reasonCode)
if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque {
status, reason = firewallRuleSyncBlocked, "managed runtime rule cannot be safely removed"
status = firewallRuleSyncBlocked
reasonCode = firewallsync.ReasonUnsafeRemoval
reason = firewallsync.ReasonMessage(reasonCode)
entry.err = filter.ErrProtectedRule
}
entry.item = dto.FirewallRuleSyncItem{
SourceUUID: entry.source.UUID, Rule: &entry.rule,
Status: status, Reason: reason,
Status: status, ReasonCode: reasonCode, Reason: reason,
}
entries = append(entries, entry)
scopeEntries = append(scopeEntries, entry)
@@ -232,76 +236,20 @@ func (s *FirewallService) loadFirewallRuleSyncPlan(
}, nil
}
func firewallProviderHasOrderedManagedRules(provider filter.Provider) bool {
return provider == filter.ProviderIptables || provider == filter.ProviderNftables || provider == filter.ProviderUFW
}
func planFirewallManagedOrder(
snapshot filter.Snapshot,
entries []*firewallRuleSyncEntry,
) {
if !firewallProviderHasOrderedManagedRules(snapshot.Scope.Provider) {
return
}
byMarker := make(map[string]*firewallRuleSyncEntry, len(entries))
desiredMarkers := make([]string, 0, len(entries))
for _, entry := range entries {
if entry.remove == nil && entry.err == nil && entry.desired.Marker != "" {
byMarker[entry.desired.Marker] = entry
desiredMarkers = append(desiredMarkers, entry.desired.Marker)
}
}
if len(byMarker) < 2 {
return
}
actual := make([]string, 0, len(byMarker))
segments := make(map[string]int, len(byMarker))
segment := 0
for _, observed := range snapshot.Rules {
_, expected := byMarker[observed.Marker]
if expected {
actual = append(actual, observed.Marker)
if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque {
segment++
segments[observed.Marker] = segment
segment++
} else {
segments[observed.Marker] = segment
}
continue
}
if strings.HasPrefix(observed.Marker, "1panel-rule:") &&
!observed.Protected && observed.ParseStatus != filter.ParseStatusOpaque {
continue
}
segment++
}
desired := make([]string, 0, len(actual))
for _, entry := range entries {
marker := entry.desired.Marker
if _, exists := segments[marker]; exists {
desired = append(desired, marker)
}
}
if slices.Equal(actual, desired) {
return
}
drifted := make(map[string]struct{}, len(desired))
for index := range desired {
if actual[index] != desired[index] {
drifted[actual[index]] = struct{}{}
drifted[desired[index]] = struct{}{}
}
}
feasible, previousSegment := true, -1
for _, marker := range desired {
if segments[marker] < previousSegment {
feasible = false
break
}
previousSegment = segments[marker]
}
for _, marker := range desired {
drifted, feasible := firewallsync.ManagedOrderDrift(snapshot, desiredMarkers)
for _, marker := range desiredMarkers {
if _, exists := drifted[marker]; !exists {
continue
}
@@ -529,7 +477,7 @@ func (s *FirewallService) loadStoredFirewallRuleSyncCandidates(
if err != nil {
return "", nil, false, err
}
sortFirewallPolicies(stored, selected)
model.SortFirewallRules(stored, selected)
entries := make([]*firewallRuleSyncEntry, 0, len(stored))
hasCompileErrors := false
for _, record := range stored {
@@ -558,137 +506,12 @@ func scopesWithFirewallSyncCandidates(scopes []filter.Scope, candidates map[stri
return result
}
func firewallRuleSyncScopes(provider filter.Provider) []filter.Scope {
base := filter.Scope{Provider: provider, Direction: filter.DirectionInput}
switch provider {
case filter.ProviderIptables, filter.ProviderNftables:
result := make([]filter.Scope, 0, 6)
for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} {
for _, chain := range []string{filter.BasicBeforeChain, filter.IptablesInputChain, filter.BasicAfterChain} {
scope := base
scope.Family, scope.Table, scope.Chain = family, "filter", chain
result = append(result, scope)
}
}
return result
case filter.ProviderFirewalld:
base.Family, base.Zone = filter.FamilyInet, filter.FirewalldInputZone
return []filter.Scope{base}
case filter.ProviderUFW:
result := make([]filter.Scope, 0, 2)
for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} {
scope := base
scope.Family, scope.Chain = family, filter.UFWInputChain
result = append(result, scope)
}
return result
default:
return nil
}
}
func firewallPolicyRulesForProvider(stored model.FirewallRule, provider filter.Provider) ([]filter.FirewallRule, error) {
if stored.CompatibilityError != "" {
return nil, fmt.Errorf("%w: %s", filter.ErrUnsupportedScope, stored.CompatibilityError)
}
connectionStates := make([]string, 0)
if stored.ConnectionStates != "" {
connectionStates = strings.Split(stored.ConnectionStates, ",")
}
base := filter.FirewallRule{
Protocol: stored.Protocol, SourceAddress: stored.SourceAddress, SourcePort: stored.SourcePort,
DestinationAddress: stored.DestinationAddress, DestinationPort: stored.DestinationPort,
Interface: stored.Interface, ConnectionStates: connectionStates,
Action: filter.Action(stored.Action), Description: stored.Description,
}
if provider == filter.ProviderFirewalld {
base.Priority = stored.Priority
}
families := []filter.Family{filter.Family(stored.Family)}
if provider != filter.ProviderFirewalld && len(families) == 1 && families[0] == filter.FamilyInet {
hasIPv4, hasIPv6 := firewallRuleAddressFamilies(base)
switch {
case hasIPv4 && hasIPv6:
return nil, fmt.Errorf("%w: inet policy contains both IPv4 and IPv6 addresses", filter.ErrUnsupportedScope)
case hasIPv6 || strings.EqualFold(base.Protocol, "icmpv6"):
families = []filter.Family{filter.FamilyIPv6}
case hasIPv4:
families = []filter.Family{filter.FamilyIPv4}
default:
families = []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6}
}
}
result := make([]filter.FirewallRule, 0, len(families))
for _, family := range families {
rule := base
rule.Scope = filter.Scope{Provider: provider, Family: family, Direction: filter.DirectionInput}
switch provider {
case filter.ProviderIptables, filter.ProviderNftables:
rule.Scope.Table, rule.Scope.Chain = "filter", filter.IptablesInputChain
case filter.ProviderFirewalld:
rule.Scope.Zone = filter.FirewalldInputZone
case filter.ProviderUFW:
rule.Scope.Chain = filter.UFWInputChain
default:
return nil, fmt.Errorf("%w: unsupported firewall provider %q", filter.ErrProviderUnavailable, provider)
}
expanded, err := filter.ExpandAtomicRules(rule)
if err != nil {
return nil, err
}
result = append(result, expanded...)
}
return result, nil
}
func sortFirewallPolicies(stored []model.FirewallRule, provider filter.Provider) {
sort.SliceStable(stored, func(i, j int) bool {
left, right := stored[i], stored[j]
if provider == filter.ProviderFirewalld {
switch {
case left.Priority == nil && right.Priority != nil:
return false
case left.Priority != nil && right.Priority == nil:
return true
case left.Priority != nil && right.Priority != nil && *left.Priority != *right.Priority:
return *left.Priority < *right.Priority
}
} else {
switch {
case left.Sequence == nil && right.Sequence != nil:
return false
case left.Sequence != nil && right.Sequence == nil:
return true
case left.Sequence != nil && right.Sequence != nil && *left.Sequence != *right.Sequence:
return *left.Sequence < *right.Sequence
}
}
return left.UUID < right.UUID
})
}
func firewallRuleAddressFamilies(rule filter.FirewallRule) (bool, bool) {
hasIPv4, hasIPv6 := false, false
for _, address := range []string{rule.SourceAddress, rule.DestinationAddress} {
address = strings.TrimSpace(address)
if address == "" {
continue
}
if strings.Contains(address, ":") {
hasIPv6 = true
} else {
hasIPv4 = true
}
}
return hasIPv4, hasIPv6
}
func (s *FirewallService) classifyFirewallRuleSyncCandidate(
clientIP string,
snapshot filter.Snapshot,
entry *firewallRuleSyncEntry,
item filter.InventoryItem,
) (string, string) {
) (firewallsync.Status, string) {
switch item.Match {
case filter.InventoryMatchExact:
return firewallRuleSyncExisting, "rule already matches database policy"
@@ -818,64 +641,24 @@ func buildDatabaseSyncPlan[T any](
key func(T) string,
actualItem func(T) dto.FirewallRuleSyncItem,
) databaseSyncPlan {
items := make([]dto.FirewallRuleSyncItem, 0, len(desired)+len(actual))
actualByKey := make(map[string][]int, len(actual))
for index, value := range actual {
actualByKey[key(value)] = append(actualByKey[key(value)], index)
}
matched := make([]bool, len(actual))
candidates := make([]firewallsync.Desired[T, dto.FirewallRuleSyncItem], 0, len(desired))
for _, candidate := range desired {
item := candidate.item
switch {
case candidate.err != nil:
item.Status, item.Reason = firewallRuleSyncBlocked, candidate.err.Error()
default:
match := unmatchedDatabaseSyncIndex(actualByKey[key(candidate.value)], matched)
if match >= 0 {
matched[match] = true
item.Status, item.Reason = firewallRuleSyncExisting, "rule already exists in target backend"
} else {
item.Status = firewallRuleSyncReady
}
}
items = append(items, item)
candidates = append(candidates, firewallsync.Desired[T, dto.FirewallRuleSyncItem]{
Value: candidate.value, Payload: candidate.item, Err: candidate.err,
})
}
for index, value := range actual {
if matched[index] {
continue
}
item := actualItem(value)
item.Status, item.Reason = "remove", "rule exists only in target backend"
diff := firewallsync.Diff(candidates, actual, key, actualItem)
items := make([]dto.FirewallRuleSyncItem, 0, len(diff))
for _, diffItem := range diff {
item := diffItem.Payload
item.Status, item.ReasonCode, item.Reason = diffItem.Status, diffItem.ReasonCode, diffItem.Reason
items = append(items, item)
}
return databaseSyncPlan{subsystem: subsystem, target: target, items: items}
}
func unmatchedDatabaseSyncIndex(indices []int, matched []bool) int {
for _, index := range indices {
if !matched[index] {
return index
}
}
return -1
}
func databaseSyncStatesEqual[T any](left, right []T, key func(T) string) bool {
if len(left) != len(right) {
return false
}
counts := make(map[string]int, len(left))
for _, value := range left {
counts[key(value)]++
}
for _, value := range right {
valueKey := key(value)
if counts[valueKey] == 0 {
return false
}
counts[valueKey]--
}
return true
return firewallsync.StatesEqual(left, right, key)
}
func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview {
@@ -894,7 +677,7 @@ func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview {
case firewallRuleSyncBlocked:
result.Blocked++
result.Total++
case "remove":
case firewallsync.StatusRemove:
result.Removed++
}
}
@@ -904,7 +687,7 @@ func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview {
func (p databaseSyncPlan) baseResult() dto.FirewallRuleSyncResult {
result := dto.FirewallRuleSyncResult{Subsystem: p.subsystem, TargetProvider: p.target}
for _, item := range p.items {
if item.Status != "remove" {
if item.Status != firewallsync.StatusRemove {
result.Total++
}
}
@@ -921,7 +704,7 @@ func (p databaseSyncPlan) completedResult() dto.FirewallRuleSyncResult {
result.Skipped++
case firewallRuleSyncBlocked:
appendDatabaseSyncFailure(&result, item, errors.New(item.Reason))
case "remove":
case firewallsync.StatusRemove:
result.Removed++
}
}
@@ -1116,10 +899,10 @@ func (s *DockerPortGuardService) syncRules(
}
plan := buildDockerDatabaseSyncPlan(filter.Provider(target), policies, targetPolicies)
result, reconcileErr := plan.reconcile(func() error {
if err := reconcileDockerGuardSyncTarget(target, runtimePolicies, targetRuntime); err != nil {
if err := docker_guard.ReconcileTarget(target, runtimePolicies, targetRuntime); err != nil {
return err
}
if err := verifyDockerGuardRuleSync(targetRuntime, runtimePolicies); err != nil {
if err := docker_guard.Verify(targetRuntime, runtimePolicies); err != nil {
return err
}
if len(policies) == 0 {
@@ -1153,7 +936,7 @@ func buildDockerDatabaseSyncPlan(
})
}
return buildDatabaseSyncPlan(
"docker", target, desired, actual, dockerGuardPolicySyncKey,
"docker", target, desired, actual, docker_guard.PolicySyncKey,
func(policy docker_guard.Policy) dto.FirewallRuleSyncItem {
return dto.FirewallRuleSyncItem{SourceUUID: policy.UUID, DockerRule: dockerGuardRuntimeRuleSyncDTO(policy)}
},
@@ -1178,7 +961,7 @@ func (r *firewallScopeReconciler) reconcile() {
continue
}
switch entry.item.Status {
case "remove":
case firewallsync.StatusRemove:
removes = append(removes, entry)
case firewallRuleSyncReady:
switch entry.match {
@@ -1217,7 +1000,7 @@ func (r *firewallScopeReconciler) applyGroups(entries []*firewallRuleSyncEntry)
if len(entries) == 0 {
return
}
batch := r.runtime.adapter.Provider() == filter.ProviderIptables || r.runtime.adapter.Provider() == filter.ProviderNftables
batch := r.runtime.Provider() == filter.ProviderIptables || r.runtime.Provider() == filter.ProviderNftables
if batch {
r.apply(entries)
return
@@ -1237,7 +1020,7 @@ func (r *firewallScopeReconciler) apply(entries []*firewallRuleSyncEntry) {
continue
}
if !changed {
if entry.item.Status == "remove" {
if entry.item.Status == firewallsync.StatusRemove {
entry.outcome = firewallRuleSyncRemoved
} else {
entry.outcome = firewallRuleSyncSkipped
@@ -1261,7 +1044,7 @@ func (r *firewallScopeReconciler) apply(entries []*firewallRuleSyncEntry) {
return
}
for _, entry := range active {
if entry.item.Status == "remove" {
if entry.item.Status == firewallsync.StatusRemove {
entry.outcome = firewallRuleSyncRemoved
} else {
entry.outcome = firewallRuleSyncApplied
@@ -1294,7 +1077,7 @@ func (r *firewallScopeReconciler) restoreOrder() {
changed := false
for step := 0; step < len(desiredMarkers); step++ {
marker, position, converged, err := nextFirewallManagedOrderChange(r.snapshot, desiredMarkers)
marker, position, converged, err := firewallsync.NextManagedOrderChange(r.snapshot, desiredMarkers)
if err != nil {
r.failOrder(reorderEntries, err)
return
@@ -1309,18 +1092,18 @@ func (r *firewallScopeReconciler) restoreOrder() {
}
return
}
observed, _, exists := firewallRuleSyncObservedByMarker(r.snapshot, marker)
observed, _, exists := firewallsync.ObservedByMarker(r.snapshot, marker)
if !exists {
r.failOrder(reorderEntries, filter.ErrRuleStale)
return
}
after := firewallRuleSyncObservedRule(observed)
after := firewallsync.ObservedRule(observed)
target := int64(position)
after.OrderIndex = &target
before := firewallRuleSyncObservedRule(observed)
before := firewallsync.ObservedRule(observed)
locator := observed.Locator
operation := filter.ChangeReorder
if r.runtime.adapter.Provider() == filter.ProviderUFW {
if r.runtime.Provider() == filter.ProviderUFW {
operation = filter.ChangeUpdate
}
_, verification, executeErr := r.runtime.Execute(r.ctx, r.snapshot, []filter.DesiredChange{{
@@ -1339,59 +1122,6 @@ func (r *firewallScopeReconciler) restoreOrder() {
r.failOrder(reorderEntries, fmt.Errorf("%w: managed rule order did not converge", filter.ErrVerificationFailed))
}
func nextFirewallManagedOrderChange(
snapshot filter.Snapshot,
desiredMarkers []string,
) (string, int, bool, error) {
expected := make(map[string]struct{}, len(desiredMarkers))
for _, marker := range desiredMarkers {
expected[marker] = struct{}{}
}
actual := make([]string, 0, len(desiredMarkers))
positions := make([]int, 0, len(desiredMarkers))
for index, observed := range snapshot.Rules {
if _, exists := expected[observed.Marker]; !exists {
continue
}
actual = append(actual, observed.Marker)
position := index + 1
if observed.Locator.Position != nil {
position = *observed.Locator.Position
}
positions = append(positions, position)
}
if len(actual) != len(desiredMarkers) {
return "", 0, false, filter.ErrRuleStale
}
for index := range desiredMarkers {
if actual[index] != desiredMarkers[index] {
return desiredMarkers[index], positions[index], false, nil
}
}
return "", 0, true, nil
}
func firewallRuleSyncObservedByMarker(snapshot filter.Snapshot, marker string) (filter.ObservedRule, int, bool) {
for index, observed := range snapshot.Rules {
if observed.Marker == marker {
position := index + 1
if observed.Locator.Position != nil {
position = *observed.Locator.Position
}
return observed, position, true
}
}
return filter.ObservedRule{}, 0, false
}
func firewallRuleSyncObservedRule(observed filter.ObservedRule) filter.FirewallRule {
rule := observed.Rule
if rule.UUID == "" && strings.HasPrefix(observed.Marker, "1panel-rule:") {
rule.UUID = strings.TrimSpace(strings.TrimPrefix(observed.Marker, "1panel-rule:"))
}
return rule
}
func (r *firewallScopeReconciler) failOrder(entries []*firewallRuleSyncEntry, err error) {
for _, entry := range entries {
if entry.outcome != firewallRuleSyncFailed {
@@ -1416,7 +1146,7 @@ func firewallRuleSyncChange(
if observed.Protected {
return filter.DesiredChange{}, false, filter.ErrProtectedRule
}
before := firewallRuleSyncObservedRule(observed)
before := firewallsync.ObservedRule(observed)
locator := observed.Locator
return filter.DesiredChange{
Operation: filter.ChangeDelete, Before: &before, Locator: &locator,
@@ -1452,7 +1182,7 @@ func firewallRuleSyncChange(
if err := filter.GuardMutation(snapshot, *item.Observed, entry.rule, clientIP, protectedPorts...); err != nil {
return filter.DesiredChange{}, false, err
}
before := firewallRuleSyncObservedRule(*item.Observed)
before := firewallsync.ObservedRule(*item.Observed)
locator := item.Observed.Locator
change.Operation = filter.ChangeUpdate
change.Before = &before
@@ -1468,29 +1198,16 @@ func firewallRuleSyncInsertionPosition(
entries []*firewallRuleSyncEntry,
target *firewallRuleSyncEntry,
) *int64 {
if !firewallProviderHasOrderedManagedRules(snapshot.Scope.Provider) {
desiredMarkers := make([]string, 0, len(entries))
for _, entry := range entries {
if entry.remove != nil || entry.err != nil {
continue
}
desiredMarkers = append(desiredMarkers, entry.desired.Marker)
}
position, exists := firewallsync.InsertionPosition(snapshot, desiredMarkers, target.desired.Marker)
if !exists {
return nil
}
targetIndex := slices.Index(entries, target)
for index := targetIndex - 1; index >= 0; index-- {
entry := entries[index]
if entry.remove != nil || entry.err != nil {
continue
}
if _, position, exists := firewallRuleSyncObservedByMarker(snapshot, entry.desired.Marker); exists {
value := int64(position + 1)
return &value
}
}
for index := targetIndex + 1; index < len(entries); index++ {
entry := entries[index]
if entry.remove != nil || entry.err != nil {
continue
}
if _, position, exists := firewallRuleSyncObservedByMarker(snapshot, entry.desired.Marker); exists {
value := int64(position)
return &value
}
}
return nil
return &position
}
-683
View File
@@ -1,683 +0,0 @@
package service
import (
"context"
"errors"
"path/filepath"
"testing"
"time"
"github.com/1Panel-dev/1Panel/agent/app/dto"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/app/repo"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/global"
agenti18n "github.com/1Panel-dev/1Panel/agent/i18n"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
"github.com/glebarez/sqlite"
"github.com/go-playground/validator/v10"
"gorm.io/gorm"
)
type fakeFirewallDatabaseSyncAdapter struct {
previewResult dto.FirewallRuleSyncPreview
syncResult dto.FirewallRuleSyncResult
previewErr error
syncErr error
previewCalls int
syncCalls int
}
func (f *fakeFirewallDatabaseSyncAdapter) previewRuleSync(
context.Context,
dto.FirewallRuleSyncRequest,
) (dto.FirewallRuleSyncPreview, error) {
f.previewCalls++
return f.previewResult, f.previewErr
}
func (f *fakeFirewallDatabaseSyncAdapter) syncRules(
context.Context,
dto.FirewallRuleSyncRequest,
) (dto.FirewallRuleSyncResult, error) {
f.syncCalls++
return f.syncResult, f.syncErr
}
func TestFirewallRuleSyncCoordinatorDispatchesSubsystemAdapters(t *testing.T) {
forwardingErr := errors.New("forwarding preview")
dockerErr := errors.New("docker sync")
forwarding := &fakeFirewallDatabaseSyncAdapter{previewErr: forwardingErr}
docker := &fakeFirewallDatabaseSyncAdapter{syncErr: dockerErr}
service := &FirewallService{forwardingSync: forwarding, dockerSync: docker}
_, err := service.PreviewRuleSync(context.Background(), "client-ip", dto.FirewallRuleSyncRequest{
Subsystem: "forwarding", TargetProvider: filter.ProviderNftables,
})
if !errors.Is(err, forwardingErr) {
t.Fatalf("forwarding preview was not dispatched to its adapter: %v", err)
}
_, err = service.SyncRules(context.Background(), "client-ip", dto.FirewallRuleSyncRequest{
Subsystem: "docker", TargetProvider: filter.ProviderNftables,
})
if !errors.Is(err, dockerErr) {
t.Fatalf("Docker synchronization was not dispatched to its adapter: %v", err)
}
if forwarding.previewCalls != 1 || forwarding.syncCalls != 0 || docker.previewCalls != 0 || docker.syncCalls != 1 {
t.Fatalf("unexpected adapter calls: forwarding=%#v docker=%#v", forwarding, docker)
}
}
func TestFirewallRuleSyncRequestValidationAllowsDatabaseSource(t *testing.T) {
validate := validator.New()
for _, subsystem := range []string{"forwarding", "docker"} {
request := dto.FirewallRuleSyncRequest{Subsystem: subsystem, TargetProvider: filter.ProviderNftables}
if err := validate.Struct(request); err != nil {
t.Fatalf("%s database synchronization rejected missing source provider: %v", subsystem, err)
}
}
if err := validate.Struct(dto.FirewallRuleSyncRequest{Subsystem: "system", SourceProvider: filter.ProviderIptables}); err == nil {
t.Fatal("synchronization request without target provider was accepted")
}
}
func TestFirewallRuleSyncFailureMessagesIncludeDetailsAndGroupDuplicates(t *testing.T) {
messages := firewallRuleSyncFailureMessages([]dto.FirewallRuleSyncFailure{
{SourceUUID: "rule-1", Error: "iptables-restore failed: invalid port"},
{SourceUUID: "rule-2", Error: "iptables-restore failed: invalid port"},
{SourceUUID: "rule-3", Error: "permission denied"},
})
want := []string{
"UUID [rule-1, rule-2]: iptables-restore failed: invalid port",
"UUID [rule-3]: permission denied",
}
if len(messages) != len(want) {
t.Fatalf("failure messages = %#v, want %#v", messages, want)
}
for index := range want {
if messages[index] != want[index] {
t.Fatalf("failure message %d = %q, want %q", index, messages[index], want[index])
}
}
}
func TestFirewallRuleSyncPreviewApplyAndRetry(t *testing.T) {
ctx := context.Background()
db := newFirewallRuleTestDB(t)
ruleRepo := repo.NewFirewallRuleRepo(db)
sourceRule := filter.FirewallRule{
Scope: filter.Scope{
Provider: filter.ProviderIptables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
},
Protocol: "tcp", DestinationPort: "8443", Action: filter.ActionAccept, Description: "sync-test",
}
sourceRecord, err := model.FirewallRuleFromDomain(sourceRule)
if err != nil {
t.Fatalf("encode source rule: %v", err)
}
sourceRecord.Origin = constant.FirewallRuleOriginCreated
sourceRecord.Owner = model.FirewallRuleOwner(constant.FirewallRuleSourceApp, "test-app")
if err := ruleRepo.Create(ctx, &sourceRecord); err != nil {
t.Fatalf("persist source rule: %v", err)
}
targetScope := filter.Scope{
Provider: filter.ProviderNftables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
}
adapter := newFakeFilterAdapter(t, targetScope, nil)
service := &FirewallService{
rules: ruleRepo,
adapters: firewallRuleRuntimeRegistry{
filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil),
},
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil },
}
request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables}
preview, err := service.PreviewRuleSync(ctx, "", request)
if err != nil {
t.Fatalf("preview rule sync: %v", err)
}
if preview.Total != 1 || preview.Ready != 1 || preview.Existing != 0 || preview.Blocked != 0 {
t.Fatalf("unexpected preview: %#v", preview)
}
result, err := service.syncRules(ctx, "", request)
if err != nil {
t.Fatalf("apply rule sync: %v", err)
}
if result.Total != 1 || result.Succeeded != 1 || result.Skipped != 0 || result.Failed != 0 {
t.Fatalf("unexpected sync result: %#v", result)
}
targetRecords, err := ruleRepo.List(ctx)
if err != nil {
t.Fatalf("list target records: %v", err)
}
if len(targetRecords) != 1 || targetRecords[0].Owner != sourceRecord.Owner {
t.Fatalf("target ownership was not preserved: %#v", targetRecords)
}
retry, err := service.syncRules(ctx, "", request)
if err != nil {
t.Fatalf("retry rule sync: %v", err)
}
if retry.Succeeded != 0 || retry.Skipped != 1 || retry.Failed != 0 {
t.Fatalf("rule sync is not idempotent: %#v", retry)
}
setupFirewallTaskTestDB(t)
taskResult, err := service.SyncRules(ctx, "", dto.FirewallRuleSyncRequest{
Subsystem: "system", TargetProvider: filter.ProviderNftables,
TaskID: "firewall-sync-task",
})
if err != nil {
t.Fatal(err)
}
if !taskResult.Queued || taskResult.TaskID != "firewall-sync-task" {
t.Fatalf("plain synchronization was not queued as a task: %#v", taskResult)
}
waitFirewallSyncTask(t, taskResult.TaskID)
sourceRecords, err := ruleRepo.List(ctx)
if err != nil || len(sourceRecords) != 1 {
t.Fatalf("plain synchronization reset the source backend: %#v err=%v", sourceRecords, err)
}
}
func TestFirewallRuleSyncKeepsExpandedRulesInTheSameScope(t *testing.T) {
ctx := context.Background()
ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t))
sourceRule := filter.FirewallRule{
Scope: filter.Scope{
Provider: filter.ProviderIptables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
},
Protocol: "tcp", DestinationPort: "80,443", Action: filter.ActionAccept,
}
sourceRecord, err := model.FirewallRuleFromDomain(sourceRule)
if err != nil {
t.Fatal(err)
}
sourceRecord.Origin = constant.FirewallRuleOriginCreated
sourceRecord.Owner = constant.FirewallRuleSourceUser
if err := ruleRepo.Create(ctx, &sourceRecord); err != nil {
t.Fatal(err)
}
targetScope := filter.Scope{
Provider: filter.ProviderNftables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
}
adapter := newFakeFilterAdapter(t, targetScope, nil)
service := &FirewallService{
rules: ruleRepo,
adapters: firewallRuleRuntimeRegistry{
filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil),
},
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil },
}
request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables}
result, err := service.syncRules(ctx, "", request)
if err != nil || result.Succeeded != 2 || result.Failed != 0 {
t.Fatalf("expanded rule synchronization failed: result=%#v err=%v", result, err)
}
if len(adapter.snapshot.Rules) != 2 {
t.Fatalf("expanded rules overwrote each other: %#v", adapter.snapshot.Rules)
}
if adapter.applyCount != 1 {
t.Fatalf("same-scope rules were applied in %d calls, want one batch", adapter.applyCount)
}
if adapter.observeCount != len(firewallRuleSyncScopes(filter.ProviderNftables)) {
t.Fatalf("synchronization observed scopes %d times, want one read per scope", adapter.observeCount)
}
ports := make(map[string]struct{}, len(adapter.snapshot.Rules))
markers := make(map[string]struct{}, len(adapter.snapshot.Rules))
for _, observed := range adapter.snapshot.Rules {
ports[observed.Rule.DestinationPort] = struct{}{}
markers[observed.Marker] = struct{}{}
}
if _, exists := ports["80"]; !exists {
t.Fatalf("expanded port 80 is missing: %#v", adapter.snapshot.Rules)
}
if _, exists := ports["443"]; !exists {
t.Fatalf("expanded port 443 is missing: %#v", adapter.snapshot.Rules)
}
if len(markers) != 2 {
t.Fatalf("expanded rules reused one runtime marker: %#v", adapter.snapshot.Rules)
}
retry, err := service.syncRules(ctx, "", request)
if err != nil || retry.Succeeded != 0 || retry.Skipped != 2 || retry.Failed != 0 {
t.Fatalf("expanded rule synchronization is not idempotent: result=%#v err=%v", retry, err)
}
}
func TestFirewallRuleSyncRestoresManagedOrderFromDatabase(t *testing.T) {
ctx := context.Background()
ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t))
targetScope := filter.Scope{
Provider: filter.ProviderNftables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
}
adapter := newFakeFilterAdapter(t, targetScope, nil)
service := &FirewallService{
rules: ruleRepo,
adapters: firewallRuleRuntimeRegistry{
filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil),
},
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil },
}
records := make([]model.FirewallRule, 0, 2)
for index, port := range []string{"8080", "8081"} {
rule := filter.FirewallRule{
Scope: filter.Scope{
Provider: filter.ProviderIptables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
},
Protocol: "tcp", DestinationPort: port, Action: filter.ActionAccept,
}
record, err := model.FirewallRuleFromDomain(rule)
if err != nil {
t.Fatal(err)
}
record.UUID = []string{"first", "second"}[index]
record.Origin = constant.FirewallRuleOriginCreated
record.Owner = constant.FirewallRuleSourceUser
sequence := int64(index+1) * model.FirewallRuleSequenceStep
record.Sequence = &sequence
if err := ruleRepo.Create(ctx, &record); err != nil {
t.Fatal(err)
}
records = append(records, record)
}
compiledFirst, err := service.compileStoredFirewallRules(ctx, records[0], filter.ProviderNftables)
if err != nil {
t.Fatal(err)
}
compiledSecond, err := service.compileStoredFirewallRules(ctx, records[1], filter.ProviderNftables)
if err != nil {
t.Fatal(err)
}
observedSecond := executorObservedRule(compiledSecond[0].Rule, compiledSecond[0].Marker, 1)
observedFirst := executorObservedRule(compiledFirst[0].Rule, compiledFirst[0].Marker, 2)
// Native inventory identifies managed rules through their marker and does
// not populate the domain rule UUID.
observedSecond.Rule.UUID = ""
observedFirst.Rule.UUID = ""
adapter.snapshot, err = filter.NewSnapshot(targetScope, []filter.ObservedRule{observedSecond, observedFirst})
if err != nil {
t.Fatal(err)
}
request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables}
preview, err := service.PreviewRuleSync(ctx, "", request)
if err != nil {
t.Fatal(err)
}
if preview.Ready != 2 || preview.Existing != 0 || preview.Blocked != 0 {
t.Fatalf("order drift was not included in preview: %#v", preview)
}
result, err := service.syncRules(ctx, "", request)
if err != nil {
t.Fatal(err)
}
if result.Succeeded != 2 || result.Failed != 0 || adapter.applyCount != 1 {
t.Fatalf("managed order was not synchronized: result=%#v applies=%d", result, adapter.applyCount)
}
if adapter.snapshot.Rules[0].Marker != compiledFirst[0].Marker || adapter.snapshot.Rules[1].Marker != compiledSecond[0].Marker {
t.Fatalf("runtime order does not match database order: %#v", adapter.snapshot.Rules)
}
for index, uuid := range []string{"first", "second"} {
stored, loadErr := ruleRepo.GetByUUID(ctx, uuid)
want := int64(index+1) * model.FirewallRuleSequenceStep
if loadErr != nil || stored.Sequence == nil || *stored.Sequence != want {
t.Fatalf("synchronization rewrote database sequence for %s: record=%#v err=%v", uuid, stored, loadErr)
}
}
}
func TestFirewallRuleSyncBlocksManagedOrderAcrossExternalRule(t *testing.T) {
ctx := context.Background()
ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t))
targetScope := filter.Scope{
Provider: filter.ProviderNftables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
}
adapter := newFakeFilterAdapter(t, targetScope, nil)
service := &FirewallService{
rules: ruleRepo,
adapters: firewallRuleRuntimeRegistry{
filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil),
},
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil },
}
compiled := make([]filter.DesiredRule, 0, 2)
for index, port := range []string{"8080", "8081"} {
rule := filter.FirewallRule{
Scope: filter.Scope{
Provider: filter.ProviderIptables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
},
Protocol: "tcp", DestinationPort: port, Action: filter.ActionAccept,
}
record, err := model.FirewallRuleFromDomain(rule)
if err != nil {
t.Fatal(err)
}
record.UUID = []string{"first", "second"}[index]
record.Origin = constant.FirewallRuleOriginCreated
record.Owner = constant.FirewallRuleSourceUser
sequence := int64(index+1) * model.FirewallRuleSequenceStep
record.Sequence = &sequence
if err := ruleRepo.Create(ctx, &record); err != nil {
t.Fatal(err)
}
rules, compileErr := service.compileStoredFirewallRules(ctx, record, filter.ProviderNftables)
if compileErr != nil {
t.Fatal(compileErr)
}
compiled = append(compiled, rules[0])
}
external := compiled[0].Rule
external.UUID = ""
external.DestinationPort = "9090"
adapter.snapshot, _ = filter.NewSnapshot(targetScope, []filter.ObservedRule{
executorObservedRule(compiled[1].Rule, compiled[1].Marker, 1),
executorObservedRule(external, "", 2),
executorObservedRule(compiled[0].Rule, compiled[0].Marker, 3),
})
request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables}
preview, err := service.PreviewRuleSync(ctx, "", request)
if err != nil {
t.Fatal(err)
}
if preview.Blocked != 2 || preview.Ready != 0 {
t.Fatalf("unsafe order change was not blocked: %#v", preview)
}
result, err := service.syncRules(ctx, "", request)
if err != nil {
t.Fatal(err)
}
if result.Failed != 2 || adapter.applyCount != 0 || adapter.snapshot.Rules[1].Marker != "" {
t.Fatalf("blocked order change mutated runtime rules: result=%#v applies=%d rules=%#v", result, adapter.applyCount, adapter.snapshot.Rules)
}
}
func TestFirewallRuleSyncRejectsProviderSourceAndKeepsDatabasePolicy(t *testing.T) {
ctx := context.Background()
db := newFirewallRuleTestDB(t)
ruleRepo := repo.NewFirewallRuleRepo(db)
sourceRule := filter.FirewallRule{
Scope: filter.Scope{
Provider: filter.ProviderIptables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
},
Protocol: "tcp", DestinationPort: "8443", Action: filter.ActionAccept,
}
sourceRecord, err := model.FirewallRuleFromDomain(sourceRule)
if err != nil {
t.Fatal(err)
}
sourceRecord.Origin = constant.FirewallRuleOriginCreated
sourceRecord.Owner = constant.FirewallRuleSourceUser
if err := ruleRepo.Create(ctx, &sourceRecord); err != nil {
t.Fatal(err)
}
targetScope := filter.Scope{
Provider: filter.ProviderNftables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
}
target := newFakeFilterAdapter(t, targetScope, nil)
service := &FirewallService{
rules: ruleRepo,
adapters: firewallRuleRuntimeRegistry{
filter.ProviderNftables: newFirewallRuleRuntime(target, nil),
},
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil },
}
_, err = service.PreviewRuleSync(ctx, "", dto.FirewallRuleSyncRequest{
Subsystem: "system", SourceProvider: filter.ProviderIptables, TargetProvider: filter.ProviderNftables,
})
if err == nil {
t.Fatal("system database synchronization accepted a source provider")
}
result, err := service.syncRules(ctx, "", dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables})
if err != nil || result.Succeeded != 1 {
t.Fatalf("database synchronization failed: result=%#v err=%v", result, err)
}
records, err := ruleRepo.List(ctx)
if err != nil || len(records) != 1 || records[0].UUID != sourceRecord.UUID {
t.Fatalf("database policy was duplicated or removed: %#v err=%v", records, err)
}
}
func TestFirewallRuleSyncRemovesManagedRuntimeRuleMissingFromDatabase(t *testing.T) {
ctx := context.Background()
ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t))
scope := filter.Scope{
Provider: filter.ProviderFirewalld, Family: filter.FamilyInet,
Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput,
}
rule := filter.FirewallRule{
UUID: "orphan", Scope: scope, NativeKind: filter.NativeKindRichRule,
Protocol: "tcp", DestinationPort: "9443", Action: filter.ActionAccept,
}
observed := executorObservedRule(rule, "1panel-rule:orphan", 1)
observed.Rule.UUID = ""
adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{observed})
service := &FirewallService{
rules: ruleRepo,
adapters: firewallRuleRuntimeRegistry{
filter.ProviderFirewalld: newFirewallRuleRuntime(adapter, nil),
},
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderFirewalld, nil },
}
request := dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderFirewalld}
preview, err := service.PreviewRuleSync(ctx, "", request)
if err != nil {
t.Fatal(err)
}
if preview.Total != 0 || preview.Removed != 1 || len(preview.Items) != 1 || preview.Items[0].Status != "remove" {
t.Fatalf("unexpected orphan preview: %#v", preview)
}
result, err := service.syncRules(ctx, "", request)
if err != nil {
t.Fatal(err)
}
if result.Total != 0 || result.Removed != 1 || result.Failed != 0 || len(adapter.snapshot.Rules) != 0 {
t.Fatalf("orphan runtime rule was not removed: result=%#v rules=%#v", result, adapter.snapshot.Rules)
}
stored, err := ruleRepo.List(ctx)
if err != nil || len(stored) != 0 {
t.Fatalf("runtime cleanup changed database policies: %#v err=%v", stored, err)
}
}
func TestFirewallRuleSyncCompileFailureDoesNotCreateOrphanRemoval(t *testing.T) {
ctx := context.Background()
ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t))
scope := filter.Scope{
Provider: filter.ProviderFirewalld, Family: filter.FamilyInet,
Zone: filter.FirewalldInputZone, Direction: filter.DirectionInput,
}
rule := filter.FirewallRule{
UUID: "incompatible", Scope: scope, NativeKind: filter.NativeKindRichRule,
Protocol: "tcp", DestinationPort: "9443", Action: filter.ActionAccept,
}
record, err := model.FirewallRuleFromDomain(rule)
if err != nil {
t.Fatal(err)
}
record.UUID = rule.UUID
record.Origin = constant.FirewallRuleOriginCreated
record.Owner = constant.FirewallRuleSourceUser
record.CompatibilityError = "manual recreation required"
if err := ruleRepo.Create(ctx, &record); err != nil {
t.Fatal(err)
}
adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{
executorObservedRule(rule, "1panel-rule:"+rule.UUID, 1),
})
service := &FirewallService{
rules: ruleRepo,
adapters: firewallRuleRuntimeRegistry{
filter.ProviderFirewalld: newFirewallRuleRuntime(adapter, nil),
},
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderFirewalld, nil },
}
preview, err := service.PreviewRuleSync(ctx, "", dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderFirewalld})
if err != nil {
t.Fatal(err)
}
if preview.Blocked != 1 || preview.Removed != 0 || len(preview.Items) != 1 {
t.Fatalf("compile failure classified its runtime rule as an orphan: %#v", preview)
}
}
func TestFirewallRuleSyncBlockedPlanDoesNotRemoveOrphans(t *testing.T) {
ctx := context.Background()
ruleRepo := repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t))
scope := filter.Scope{
Provider: filter.ProviderNftables, Family: filter.FamilyIPv4,
Table: "filter", Chain: filter.IptablesInputChain, Direction: filter.DirectionInput,
}
blockedRule := filter.FirewallRule{
Scope: scope, Protocol: "all", Action: filter.ActionDrop,
}
record, err := model.FirewallRuleFromDomain(blockedRule)
if err != nil {
t.Fatal(err)
}
record.Origin = constant.FirewallRuleOriginCreated
record.Owner = constant.FirewallRuleSourceUser
if err := ruleRepo.Create(ctx, &record); err != nil {
t.Fatal(err)
}
orphanRule := filter.FirewallRule{
UUID: "orphan", Scope: scope, Protocol: "tcp", DestinationPort: "9443", Action: filter.ActionAccept,
}
adapter := newFakeFilterAdapter(t, scope, []filter.ObservedRule{
executorObservedRule(orphanRule, "1panel-rule:orphan", 1),
})
service := &FirewallService{
rules: ruleRepo,
adapters: firewallRuleRuntimeRegistry{
filter.ProviderNftables: newFirewallRuleRuntime(adapter, nil),
},
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderNftables, nil },
}
result, err := service.syncRules(ctx, "203.0.113.10", dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables})
if err != nil {
t.Fatal(err)
}
if result.Failed != 1 || result.Removed != 0 || adapter.applyCount != 0 || len(adapter.snapshot.Rules) != 1 {
t.Fatalf("blocked synchronization mutated target rules: result=%#v rules=%#v applies=%d", result, adapter.snapshot.Rules, adapter.applyCount)
}
}
func waitFirewallSyncTask(t *testing.T, taskID string) {
t.Helper()
taskRepo := repo.NewITaskRepo()
deadline := time.Now().Add(5 * time.Second)
for {
completed := false
record, loadErr := taskRepo.GetFirst(taskRepo.WithByID(taskID))
if loadErr == nil && record.Status != constant.StatusExecuting {
if record.Status != constant.StatusSuccess {
t.Fatalf("firewall synchronization task failed: %#v", record)
}
completed = true
}
firewallRuleSyncTaskMu.Lock()
activeTaskID := firewallRuleSyncTaskID
firewallRuleSyncTaskMu.Unlock()
if completed && activeTaskID == "" {
return
}
if time.Now().After(deadline) {
t.Fatal("timed out waiting for firewall synchronization task")
}
time.Sleep(20 * time.Millisecond)
}
}
func setupFirewallTaskTestDB(t *testing.T) {
t.Helper()
agenti18n.Init()
previousDB := global.TaskDB
previousDir := global.Dir.TaskDir
taskDir := t.TempDir()
db, err := gorm.Open(sqlite.Open(filepath.Join(taskDir, "task.db")), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.Task{}); err != nil {
t.Fatal(err)
}
global.TaskDB = db
global.Dir.TaskDir = taskDir
t.Cleanup(func() {
global.TaskDB = previousDB
global.Dir.TaskDir = previousDir
})
}
func TestSortFirewallPoliciesUsesProviderPlacement(t *testing.T) {
sequenceOne, sequenceTwo := model.FirewallRuleSequenceStep, 2*model.FirewallRuleSequenceStep
priorityLow, priorityHigh := -100, 100
policies := []model.FirewallRule{
{UUID: "high", Priority: &priorityHigh, Sequence: &sequenceOne},
{UUID: "none"},
{UUID: "low", Priority: &priorityLow, Sequence: &sequenceTwo},
}
positional := append([]model.FirewallRule(nil), policies...)
sortFirewallPolicies(positional, filter.ProviderUFW)
if positional[0].UUID != "high" || positional[1].UUID != "low" || positional[2].UUID != "none" {
t.Fatalf("positional policies were not sorted by sequence: %#v", positional)
}
weighted := append([]model.FirewallRule(nil), policies...)
sortFirewallPolicies(weighted, filter.ProviderFirewalld)
if weighted[0].UUID != "low" || weighted[1].UUID != "high" || weighted[2].UUID != "none" {
t.Fatalf("firewalld policies were not sorted by priority: %#v", weighted)
}
}
func TestFirewallRuleSyncRejectsCurrentProviderAsSource(t *testing.T) {
service := &FirewallService{
rules: repo.NewFirewallRuleRepo(newFirewallRuleTestDB(t)),
selectedProvider: func(context.Context) (filter.Provider, error) { return filter.ProviderIptables, nil },
}
_, err := service.PreviewRuleSync(context.Background(), "", dto.FirewallRuleSyncRequest{
SourceProvider: filter.ProviderIptables,
TargetProvider: filter.ProviderIptables,
})
if err == nil {
t.Fatal("expected identical source and target providers to be rejected")
}
}
func TestDatabaseRuleSyncTargetUsesExplicitTargetOnly(t *testing.T) {
target, err := databaseRuleSyncTarget(dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderNftables}, "Docker")
if err != nil || target != filter.ProviderNftables {
t.Fatalf("target = %q, err = %v", target, err)
}
if _, err := databaseRuleSyncTarget(dto.FirewallRuleSyncRequest{
SourceProvider: filter.ProviderIptables,
TargetProvider: filter.ProviderNftables,
}, "Docker"); err == nil {
t.Fatal("expected database-backed synchronization to reject a source provider")
}
if _, err := databaseRuleSyncTarget(dto.FirewallRuleSyncRequest{TargetProvider: filter.ProviderUFW}, "Docker"); err == nil {
t.Fatal("expected database-backed synchronization to reject a non-netfilter target")
}
}
@@ -319,7 +319,7 @@ func (s *ForwardingService) loadRuleSyncCandidates(
Family: record.Family, Protocol: record.Protocol, Port: record.Port, TargetIP: record.TargetIP,
TargetPort: record.TargetPort, Interface: record.Interface,
}
normalized, normalizeErr := forwardingproviders.NormalizeRule(rule)
normalized, normalizeErr := forwarding.NormalizeRule(rule)
candidates = append(candidates, forwardingRuleSyncCandidate{rule: normalized, err: normalizeErr})
}
targetStatus, err := target.Status()
@@ -358,7 +358,7 @@ func verifyForwardingRuleSync(target *forwarding.Manager, desired []forwarding.R
func normalizeForwardingRuntimeRules(rules []forwarding.Rule) ([]forwarding.Rule, error) {
normalized := make([]forwarding.Rule, 0, len(rules))
for _, rule := range rules {
item, err := forwardingproviders.NormalizeRule(rule)
item, err := forwarding.NormalizeRule(rule)
if err != nil {
return nil, fmt.Errorf("normalize target forwarding rule %s: %w", rule.Identity(), err)
}
@@ -473,7 +473,7 @@ func mergeForwardingInventory(
items := make([]forwardingInventoryItem, 0, len(stored)+len(runtime))
byIdentity := make(map[string]int, len(stored)+len(runtime))
for _, record := range stored {
rule, err := forwardingproviders.NormalizeRule(forwarding.Rule{
rule, err := forwarding.NormalizeRule(forwarding.Rule{
Family: record.Family, Protocol: record.Protocol, Port: record.Port, TargetIP: record.TargetIP,
TargetPort: record.TargetPort, Interface: record.Interface,
})
@@ -485,7 +485,7 @@ func mergeForwardingInventory(
items = append(items, forwardingInventoryItem{ID: record.ID, Rule: rule, IsDesired: true})
}
for _, observed := range runtime {
rule, err := forwardingproviders.NormalizeRule(observed)
rule, err := forwarding.NormalizeRule(observed)
if err != nil {
return nil, fmt.Errorf("normalize runtime forwarding rule: %w", err)
}
@@ -518,7 +518,7 @@ func lastForwardingSyncError() string {
func applyForwardingOperations(current []forwarding.Rule, requested []dto.ForwardRuleOperation) ([]forwarding.Rule, error) {
desired := make([]forwarding.Rule, 0, len(current)+len(requested))
for _, rule := range current {
normalized, err := forwardingproviders.NormalizeRule(rule)
normalized, err := forwarding.NormalizeRule(rule)
if err != nil {
return nil, fmt.Errorf("normalize persisted forwarding rule: %w", err)
}
@@ -526,7 +526,7 @@ func applyForwardingOperations(current []forwarding.Rule, requested []dto.Forwar
}
for _, operation := range requested {
for _, protocol := range strings.Split(operation.Protocol, "/") {
rule, err := forwardingproviders.NormalizeRule(forwarding.Rule{
rule, err := forwarding.NormalizeRule(forwarding.Rule{
Family: operation.Family, Protocol: protocol, Port: operation.Port, TargetIP: operation.TargetIP,
TargetPort: operation.TargetPort, Interface: operation.Interface,
})
@@ -605,7 +605,10 @@ func newForwardingManagerFor(backend string) (*forwarding.Manager, error) {
return forwarding.NewManager(candidate.adapter, candidate.runtime), nil
}
}
return nil, fmt.Errorf("%w: selected forwarding backend %s is not installed", errForwardingBackendUnavailable, backend)
return nil, fmt.Errorf(
"%w: selected forwarding backend %s %w",
errForwardingBackendUnavailable, backend, lifecycle.ErrNotInstalled,
)
}
return selectForwardingManager(candidates)
}
@@ -1,686 +0,0 @@
package service
import (
"context"
"encoding/json"
"errors"
"fmt"
"reflect"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/dto"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
forwardClient "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle"
"github.com/go-playground/validator/v10"
)
type fakeForwardingAdapter struct {
name string
rules []forwardClient.Rule
listErr error
operateErr error
enableErr error
init bool
initErr error
familyInit map[string]bool
reconciled []forwardClient.Rule
reconciles int
}
func (f *fakeForwardingAdapter) Name() string { return f.name }
func (f *fakeForwardingAdapter) List() ([]forwardClient.Rule, error) {
return append([]forwardClient.Rule(nil), f.rules...), f.listErr
}
func (f *fakeForwardingAdapter) Reconcile(rules []forwardClient.Rule) error {
f.reconciles++
f.reconciled = append([]forwardClient.Rule(nil), rules...)
if f.operateErr == nil {
f.rules = append([]forwardClient.Rule(nil), rules...)
}
return f.operateErr
}
func (f *fakeForwardingAdapter) Enable() error { return f.enableErr }
func (f *fakeForwardingAdapter) Cleanup() error { return nil }
func (f *fakeForwardingAdapter) FamilyStatus(family string) (bool, bool, error) {
if f.familyInit != nil {
initialized := f.familyInit[family]
return initialized, initialized, f.initErr
}
return f.init, f.init, f.initErr
}
func (f *fakeForwardingAdapter) InitStatus() (bool, bool, error) {
return f.init, f.init, f.initErr
}
func (f *fakeForwardingAdapter) Replay() error { return nil }
type fakeForwardingRuleRepo struct {
rules []model.ForwardingRule
listErr error
}
func (r *fakeForwardingRuleRepo) List(context.Context) ([]model.ForwardingRule, error) {
return append([]model.ForwardingRule(nil), r.rules...), r.listErr
}
func (r *fakeForwardingRuleRepo) ReplaceAll(_ context.Context, rules []model.ForwardingRule) error {
r.rules = append([]model.ForwardingRule(nil), rules...)
return nil
}
func forwardingServiceWithAdapter(adapter forwardClient.Adapter) *ForwardingService {
rules, listErr := adapter.List()
return &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return forwardClient.NewManager(adapter, nil), nil
},
rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels(rules), listErr: listErr},
enabled: func() (bool, error) { return true, nil },
persistBackend: func(string) error { return nil },
}
}
func TestForwardingAndFilterInterfacesAreSeparated(t *testing.T) {
filterType := reflect.TypeOf((*lifecycle.Client)(nil)).Elem()
for _, method := range []string{"ListForward", "PortForward", "EnableForward"} {
if _, ok := filterType.MethodByName(method); ok {
t.Fatalf("filter interface still exposes %s", method)
}
}
firewallServiceType := reflect.TypeOf((*IFirewallService)(nil)).Elem()
if _, ok := firewallServiceType.MethodByName("OperateForwardRule"); ok {
t.Fatal("firewall service still owns forwarding writes")
}
for _, method := range []string{"PreviewRuleSync", "SyncRules"} {
if _, ok := firewallServiceType.MethodByName(method); !ok {
t.Fatalf("firewall service missing centralized %s", method)
}
}
forwardingServiceType := reflect.TypeOf((*IForwardingService)(nil)).Elem()
for _, method := range []string{"LoadBaseInfo", "SearchRules", "OperateRules", "Enable", "Restore"} {
if _, ok := forwardingServiceType.MethodByName(method); !ok {
t.Fatalf("forwarding service missing %s", method)
}
}
for _, serviceType := range []reflect.Type{
forwardingServiceType,
reflect.TypeOf((*IDockerPortGuardService)(nil)).Elem(),
} {
for _, method := range []string{"PreviewRuleSync", "SyncRules"} {
if _, ok := serviceType.MethodByName(method); ok {
t.Fatalf("subsystem service still exposes centralized %s", method)
}
}
}
}
func TestForwardingInitNoLongerUsesFilterRequest(t *testing.T) {
request := dto.FilterChainOperation{Name: "1PANEL_FORWARD", Operate: "init-forward"}
if err := validator.New().Struct(request); err == nil {
t.Fatal("forwarding initialization must use the dedicated forwarding endpoint")
}
}
func TestForwardingRestoreHonorsPersistedState(t *testing.T) {
adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80",
}}}
service := forwardingServiceWithAdapter(adapter)
service.enabled = func() (bool, error) { return false, nil }
if err := service.Restore(context.Background()); err != nil {
t.Fatal(err)
}
if adapter.reconciled != nil {
t.Fatalf("disabled forwarding was restored: %#v", adapter.reconciled)
}
service.enabled = func() (bool, error) { return true, nil }
if err := service.Restore(context.Background()); err != nil {
t.Fatal(err)
}
if len(adapter.reconciled) != 1 || adapter.reconciled[0].Port != "8080" {
t.Fatalf("unexpected restored rules: %#v", adapter.reconciled)
}
}
func TestForwardingBackendSelection(t *testing.T) {
nft := &fakeForwardingAdapter{name: "nftables"}
iptables := &fakeForwardingAdapter{name: "iptables"}
manager, err := selectForwardingManager([]forwardingCandidate{{adapter: nft}, {adapter: iptables}})
if err != nil {
t.Fatal(err)
}
if manager.Name() != "nftables" {
t.Fatalf("uninitialized backends selected %q, want nftables", manager.Name())
}
iptables.init = true
manager, err = selectForwardingManager([]forwardingCandidate{{adapter: nft}, {adapter: iptables}})
if err != nil {
t.Fatal(err)
}
if manager.Name() != "iptables" {
t.Fatalf("initialized backend selected %q, want iptables", manager.Name())
}
nft.init = true
if _, err := selectForwardingManager([]forwardingCandidate{{adapter: nft}, {adapter: iptables}}); !errors.Is(err, errForwardingBackendConflict) {
t.Fatalf("both initialized backends returned %v, want conflict", err)
}
}
func TestForwardingBackendSelectionRejectsSplitFamilies(t *testing.T) {
nft := &fakeForwardingAdapter{name: "nftables", familyInit: map[string]bool{forwardClient.FamilyIPv4: true}}
iptables := &fakeForwardingAdapter{name: "iptables", familyInit: map[string]bool{forwardClient.FamilyIPv6: true}}
if _, err := selectForwardingManager([]forwardingCandidate{{adapter: nft}, {adapter: iptables}}); !errors.Is(err, errForwardingBackendConflict) {
t.Fatalf("split-family backends returned %v, want conflict", err)
}
}
func TestForwardingBackendSelectionReturnsStatusError(t *testing.T) {
wantErr := errors.New("status failed")
_, err := selectForwardingManager([]forwardingCandidate{{adapter: &fakeForwardingAdapter{name: "nftables", initErr: wantErr}}})
if !errors.Is(err, wantErr) {
t.Fatalf("got %v want %v", err, wantErr)
}
}
func TestForwardingRuleSyncReplaysPersistedRulesIntoTarget(t *testing.T) {
sourceRule := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
target := &fakeForwardingAdapter{name: "nftables"}
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return forwardClient.NewManager(target, nil), nil
},
rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{sourceRule})},
enabled: func() (bool, error) { return true, nil },
persistBackend: func(string) error { return nil },
markEnabled: func() error { return nil },
}
request := dto.FirewallRuleSyncRequest{Subsystem: "forwarding", TargetProvider: "nftables"}
preview, err := service.previewRuleSync(context.Background(), request)
if err != nil {
t.Fatal(err)
}
if preview.Total != 1 || preview.Ready != 1 || preview.TargetProvider != "nftables" || preview.Items[0].ForwardRule == nil {
t.Fatalf("unexpected preview: %#v", preview)
}
result, err := service.syncRules(context.Background(), request)
if err != nil {
t.Fatal(err)
}
if result.Succeeded != 1 || result.Failed != 0 || len(target.reconciled) != 1 || target.reconciled[0].Identity() != sourceRule.Identity() {
t.Fatalf("unexpected sync result=%#v target=%#v", result, target.reconciled)
}
target.init = true
retry, err := service.previewRuleSync(context.Background(), request)
if err != nil {
t.Fatal(err)
}
if retry.Ready != 0 || retry.Existing != 1 {
t.Fatalf("synchronized forwarding rule was not recognized: %#v", retry)
}
}
func TestForwardingRuleSyncReportsActivationFailure(t *testing.T) {
wantErr := errors.New("enable forwarding failed")
sourceRule := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
target := &fakeForwardingAdapter{name: "nftables", enableErr: wantErr}
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return forwardClient.NewManager(target, nil), nil
},
rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{sourceRule})},
persistBackend: func(string) error { return nil },
markEnabled: func() error { return nil },
}
result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "forwarding", TargetProvider: "nftables",
})
if err != nil {
t.Fatal(err)
}
if result.Succeeded != 0 || result.Failed != 1 || len(result.Errors) != 1 ||
result.Errors[0].Error != wantErr.Error() {
t.Fatalf("activation failure was not reported: %#v", result)
}
if target.reconciles != 0 {
t.Fatalf("failed activation reconciled target %d times", target.reconciles)
}
}
func TestForwardingRuleSyncTreatsWildcardInterfaceAsDatabaseDefault(t *testing.T) {
databaseRule := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
runtimeRule := databaseRule
runtimeRule.Interface = "*"
target := &fakeForwardingAdapter{name: "iptables", init: true, rules: []forwardClient.Rule{runtimeRule}}
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return forwardClient.NewManager(target, nil), nil
},
rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{databaseRule})},
enabled: func() (bool, error) { return true, nil },
persistBackend: func(string) error { return nil },
markEnabled: func() error { return nil },
}
preview, err := service.previewRuleSync(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "forwarding", TargetProvider: "iptables",
})
if err != nil {
t.Fatal(err)
}
if preview.Total != 1 || preview.Existing != 1 || preview.Ready != 0 || preview.Removed != 0 || len(preview.Items) != 1 {
t.Fatalf("wildcard interface produced duplicate sync actions: %#v", preview)
}
}
func TestForwardingRuleSyncReconcilesTargetToDatabaseState(t *testing.T) {
existing := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
missing := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8443", TargetIP: "10.0.0.3", TargetPort: "443",
}
extra := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "udp", Port: "5353", TargetIP: "10.0.0.4", TargetPort: "53",
}
target := &fakeForwardingAdapter{name: "nftables", init: true, rules: []forwardClient.Rule{existing, extra}}
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return forwardClient.NewManager(target, nil), nil
},
rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{existing, missing})},
enabled: func() (bool, error) { return true, nil },
persistBackend: func(string) error { return nil },
markEnabled: func() error { return nil },
}
preview, err := service.previewRuleSync(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "forwarding", TargetProvider: "nftables",
})
if err != nil {
t.Fatal(err)
}
if preview.Total != 2 || preview.Removed != 1 || len(preview.Items) != 3 || preview.Items[2].Status != "remove" {
t.Fatalf("extra target rule was not included in preview: %#v", preview)
}
result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "forwarding", TargetProvider: "nftables",
})
if err != nil {
t.Fatal(err)
}
if result.Succeeded != 1 || result.Skipped != 1 || result.Removed != 1 || result.Failed != 0 || target.reconciles != 1 {
t.Fatalf("unexpected sync result=%#v reconciles=%d", result, target.reconciles)
}
if len(target.rules) != 2 || forwardingRuleIndex(target.rules, existing) < 0 ||
forwardingRuleIndex(target.rules, missing) < 0 || forwardingRuleIndex(target.rules, extra) >= 0 {
t.Fatalf("target did not converge to database state: %#v", target.rules)
}
}
func TestForwardingRuleSyncRemovesExtraRulesWhenDatabaseRulesAlreadyExist(t *testing.T) {
existing := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
extra := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "udp", Port: "5353", TargetIP: "10.0.0.4", TargetPort: "53",
}
target := &fakeForwardingAdapter{name: "nftables", init: true, rules: []forwardClient.Rule{existing, extra}}
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return forwardClient.NewManager(target, nil), nil
},
rules: &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{existing})},
enabled: func() (bool, error) { return true, nil },
persistBackend: func(string) error { return nil },
markEnabled: func() error { return nil },
}
result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "forwarding", TargetProvider: "nftables",
})
if err != nil {
t.Fatal(err)
}
if result.Succeeded != 0 || result.Skipped != 1 || result.Failed != 0 || target.reconciles != 1 ||
len(target.rules) != 1 || forwardingRuleIndex(target.rules, existing) < 0 {
t.Fatalf("extra target rule was not removed: result=%#v target=%#v", result, target)
}
}
func TestForwardingRuleSyncClearsInitializedTargetWhenDatabaseIsEmpty(t *testing.T) {
extra := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
target := &fakeForwardingAdapter{name: "nftables", init: true, rules: []forwardClient.Rule{extra}}
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return forwardClient.NewManager(target, nil), nil
},
rules: &fakeForwardingRuleRepo{},
}
result, err := service.syncRules(context.Background(), dto.FirewallRuleSyncRequest{
Subsystem: "forwarding", TargetProvider: "nftables",
})
if err != nil {
t.Fatal(err)
}
if result.Total != 0 || result.Removed != 1 || target.reconciles != 1 || len(target.rules) != 0 {
t.Fatalf("empty database did not clear initialized target: result=%#v target=%#v", result, target)
}
}
func TestForwardingRuleSyncRejectsUnselectedTarget(t *testing.T) {
current := &fakeForwardingAdapter{name: "iptables"}
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return forwardClient.NewManager(current, nil), nil
},
rules: &fakeForwardingRuleRepo{},
}
request := dto.FirewallRuleSyncRequest{Subsystem: "forwarding", TargetProvider: filter.ProviderNftables}
if _, err := service.previewRuleSync(context.Background(), request); !errors.Is(err, filter.ErrProviderUnavailable) {
t.Fatalf("preview error = %v, want provider unavailable", err)
}
if _, err := service.syncRules(context.Background(), request); !errors.Is(err, filter.ErrProviderUnavailable) {
t.Fatalf("sync error = %v, want provider unavailable", err)
}
if current.reconciles != 0 {
t.Fatalf("unselected target request modified the current backend: %#v", current)
}
}
func TestForwardingDisplayName(t *testing.T) {
for backend, want := range map[string]string{
"iptables": "iptables-forward",
"nftables": "nftables-forward",
"unknown": "unknown",
} {
if got := forwardingDisplayName(backend); got != want {
t.Fatalf("forwardingDisplayName(%q) = %q, want %q", backend, got, want)
}
}
}
func TestForwardingBaseInfoIncludesFamilyStatus(t *testing.T) {
adapter := &fakeForwardingAdapter{
name: "nftables",
familyInit: map[string]bool{forwardClient.FamilyIPv4: true},
}
base, err := forwardingServiceWithAdapter(adapter).LoadBaseInfo()
if err != nil {
t.Fatal(err)
}
if !base.IPv4.Available || !base.IPv4.Initialized || !base.IPv4.Bound {
t.Fatalf("unexpected IPv4 status: %#v", base.IPv4)
}
if !base.IPv6.Available || base.IPv6.Initialized || base.IPv6.Bound {
t.Fatalf("unexpected IPv6 status: %#v", base.IPv6)
}
}
func TestForwardingBaseInfoReportsMissingBackend(t *testing.T) {
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return nil, errForwardingBackendUnavailable
},
installedProviders: func() []string { return nil },
}
base, err := service.LoadBaseInfo()
if err != nil {
t.Fatal(err)
}
if base.IsExist || base.Name != "-" || base.Backend != "-" {
t.Fatalf("unexpected missing forwarding backend status: %#v", base)
}
}
func TestForwardingBaseInfoReportsUnavailableSelection(t *testing.T) {
wantErr := fmt.Errorf("%w: selected forwarding backend nftables is not installed", errForwardingBackendUnavailable)
service := &ForwardingService{
managerFactory: func() (*forwardClient.Manager, error) {
return nil, wantErr
},
installedProviders: func() []string { return []string{constant.FirewallProviderIptables} },
}
base, err := service.LoadBaseInfo()
if err != nil {
t.Fatal(err)
}
if !base.IsExist || base.Message != wantErr.Error() || base.Name != "-" || base.Backend != "-" {
t.Fatalf("unexpected unavailable forwarding selection status: %#v", base)
}
}
func TestForwardingSearchPreservesAPIShapeAndPagination(t *testing.T) {
adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{
{Num: "1", Family: forwardClient.FamilyIPv6, Protocol: "tcp", Port: "8080", TargetIP: "2001:db8::2", TargetPort: "80", Interface: "eth0"},
{Num: "2", Protocol: "udp", Port: "5353", TargetIP: "127.0.0.1", TargetPort: "53"},
}}
service := forwardingServiceWithAdapter(adapter)
service.rules.(*fakeForwardingRuleRepo).rules[0].ID = 42
total, value, err := service.SearchRules(dto.ForwardRuleSearch{PageInfo: dto.PageInfo{Page: 1, PageSize: 10}, Info: "2001:db8"})
if err != nil {
t.Fatal(err)
}
if total != 1 {
t.Fatalf("got total %d want 1", total)
}
items, ok := value.([]dto.ForwardRule)
if !ok || len(items) != 1 || items[0].ID != 42 || items[0].Port != "8080" ||
items[0].Family != forwardClient.FamilyIPv6 || !items[0].IsDesired || !items[0].IsRuntime ||
items[0].SyncStatus != forwardingSyncConverged {
t.Fatalf("unexpected items: %#v", value)
}
data, err := json.Marshal(items[0])
if err != nil {
t.Fatal(err)
}
var fields map[string]interface{}
if err := json.Unmarshal(data, &fields); err != nil {
t.Fatal(err)
}
wantFields := []string{"id", "chain", "family", "address", "port", "protocol", "strategy", "num", "targetIP", "targetPort", "interface", "usedStatus", "description", "isDesired", "isRuntime", "syncStatus"}
for _, field := range wantFields {
if _, ok := fields[field]; !ok {
t.Fatalf("forward response dropped compatibility field %q: %s", field, data)
}
}
}
func TestForwardingSearchTrimsAndMatchesAllDisplayedFields(t *testing.T) {
adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{
{Family: forwardClient.FamilyIPv4, Protocol: "udp", Port: "5353", TargetIP: "127.0.0.1", TargetPort: "53", Interface: "eth0"},
}}
service := forwardingServiceWithAdapter(adapter)
for _, keyword := range []string{" UDP ", "ipv4", "ETH0", "converged"} {
total, value, err := service.SearchRules(dto.ForwardRuleSearch{
PageInfo: dto.PageInfo{Page: 1, PageSize: 10}, Info: keyword,
})
if err != nil {
t.Fatalf("search %q: %v", keyword, err)
}
if items := value.([]dto.ForwardRule); total != 1 || len(items) != 1 {
t.Fatalf("search %q returned total=%d items=%#v", keyword, total, items)
}
}
}
func TestForwardingSearchMergesDesiredAndRuntimeState(t *testing.T) {
desired := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
runtimeOnly := forwardClient.Rule{
Family: forwardClient.FamilyIPv6, Protocol: "udp", Port: "5353", TargetIP: "2001:db8::2", TargetPort: "53",
}
adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{runtimeOnly}}
service := forwardingServiceWithAdapter(adapter)
service.rules = &fakeForwardingRuleRepo{rules: forwardingRuleModels([]forwardClient.Rule{desired})}
total, value, err := service.SearchRules(dto.ForwardRuleSearch{PageInfo: dto.PageInfo{Page: 1, PageSize: 10}})
if err != nil {
t.Fatal(err)
}
items, ok := value.([]dto.ForwardRule)
if !ok || total != 2 || len(items) != 2 {
t.Fatalf("unexpected forwarding inventory: %#v", value)
}
if !items[0].IsDesired || items[0].IsRuntime || items[0].SyncStatus != forwardingSyncMissing {
t.Fatalf("unexpected missing desired rule: %#v", items[0])
}
if items[1].IsDesired || !items[1].IsRuntime || items[1].SyncStatus != forwardingSyncRuntimeOnly {
t.Fatalf("unexpected runtime-only rule: %#v", items[1])
}
}
func TestForwardingDuplicateIdentityIncludesAddressFamily(t *testing.T) {
adapter := &fakeForwardingAdapter{name: "nftables", rules: []forwardClient.Rule{{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}}}
service := forwardingServiceWithAdapter(adapter)
err := service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{{
Operation: "add", Family: forwardClient.FamilyIPv6, Protocol: "tcp", Port: "8080", TargetIP: "2001:db8::2", TargetPort: "80",
}}})
if err != nil {
t.Fatal(err)
}
if len(adapter.reconciled) != 2 || adapter.reconciled[1].Family != forwardClient.FamilyIPv6 {
t.Fatalf("unexpected reconciled rules: %#v", adapter.reconciled)
}
}
func TestForwardingOperatePreservesDuplicateAndOrderingContracts(t *testing.T) {
existing := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{
{Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80"},
}}
service := forwardingServiceWithAdapter(existing)
err := service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{{
Operation: "add", Protocol: "tcp", Port: "8080", TargetPort: "80",
}}})
if err == nil {
t.Fatal("duplicate forwarding rule must be rejected")
}
adapter := &fakeForwardingAdapter{name: "iptables"}
service = forwardingServiceWithAdapter(adapter)
err = service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{
{Operation: "add", Protocol: "tcp/udp", Port: "9000", TargetIP: "10.0.0.2", TargetPort: "90"},
{Operation: "remove", Num: "1", Protocol: "tcp", Port: "8001", TargetIP: "10.0.0.2", TargetPort: "81"},
{Operation: "remove", Num: "3", Protocol: "tcp", Port: "8003", TargetIP: "10.0.0.2", TargetPort: "83"},
}})
if err != nil {
t.Fatal(err)
}
want := []forwardClient.Rule{
{Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "9000", TargetIP: "10.0.0.2", TargetPort: "90"},
{Family: forwardClient.FamilyIPv4, Protocol: "udp", Port: "9000", TargetIP: "10.0.0.2", TargetPort: "90"},
}
if !reflect.DeepEqual(adapter.reconciled, want) {
t.Fatalf("desired forwarding rules changed\ngot %#v\nwant %#v", adapter.reconciled, want)
}
}
func TestForwardingSearchReturnsAdapterError(t *testing.T) {
wantErr := errors.New("list failed")
service := forwardingServiceWithAdapter(&fakeForwardingAdapter{name: "nftables", listErr: wantErr})
_, _, err := service.SearchRules(dto.ForwardRuleSearch{PageInfo: dto.PageInfo{Page: 1, PageSize: 20}})
if !errors.Is(err, wantErr) {
t.Fatalf("got %v want %v", err, wantErr)
}
}
func TestForwardingOperateReturnsAdapterListError(t *testing.T) {
wantErr := errors.New("list failed")
service := forwardingServiceWithAdapter(&fakeForwardingAdapter{name: "iptables", listErr: wantErr})
err := service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{{
Operation: "remove", Protocol: "tcp", Port: "8080", TargetPort: "80",
}}})
if !errors.Is(err, wantErr) {
t.Fatalf("got %v want %v", err, wantErr)
}
}
func TestForwardingForceDeleteOnlySuppressesRemoveErrors(t *testing.T) {
wantErr := errors.New("operate failed")
removeAdapter := &fakeForwardingAdapter{name: "iptables", operateErr: wantErr}
removeService := forwardingServiceWithAdapter(removeAdapter)
err := removeService.OperateRules(dto.ForwardRuleOperate{ForceDelete: true, Rules: []dto.ForwardRuleOperation{{
Operation: "remove", Protocol: "tcp", Port: "8080", TargetPort: "80",
}}})
if err != nil {
t.Fatalf("forced remove returned %v", err)
}
addAdapter := &fakeForwardingAdapter{name: "iptables", operateErr: wantErr}
addService := forwardingServiceWithAdapter(addAdapter)
err = addService.OperateRules(dto.ForwardRuleOperate{ForceDelete: true, Rules: []dto.ForwardRuleOperation{{
Operation: "add", Protocol: "tcp", Port: "8080", TargetPort: "80",
}}})
if !errors.Is(err, wantErr) {
t.Fatalf("got %v want %v", err, wantErr)
}
}
func TestForwardingForceDeleteExposesRuntimeOnlyRuleAndSyncError(t *testing.T) {
recordForwardingSyncError(nil)
t.Cleanup(func() { recordForwardingSyncError(nil) })
wantErr := errors.New("runtime reconcile failed")
rule := forwardClient.Rule{
Family: forwardClient.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
adapter := &fakeForwardingAdapter{name: "iptables", rules: []forwardClient.Rule{rule}}
service := forwardingServiceWithAdapter(adapter)
adapter.operateErr = wantErr
err := service.OperateRules(dto.ForwardRuleOperate{ForceDelete: true, Rules: []dto.ForwardRuleOperation{{
Operation: "remove", Family: rule.Family, Protocol: rule.Protocol, Port: rule.Port,
TargetIP: rule.TargetIP, TargetPort: rule.TargetPort,
}}})
if err != nil {
t.Fatalf("forced remove returned %v", err)
}
base, err := service.LoadBaseInfo()
if err != nil {
t.Fatal(err)
}
if base.SyncError != wantErr.Error() {
t.Fatalf("sync error = %q, want %q", base.SyncError, wantErr)
}
_, value, err := service.SearchRules(dto.ForwardRuleSearch{PageInfo: dto.PageInfo{Page: 1, PageSize: 10}})
if err != nil {
t.Fatal(err)
}
items := value.([]dto.ForwardRule)
if len(items) != 1 || items[0].IsDesired || !items[0].IsRuntime || items[0].SyncStatus != forwardingSyncRuntimeOnly {
t.Fatalf("unexpected runtime-only inventory after forced delete: %#v", items)
}
}
func TestForwardingOperateKeepsDesiredStateWhenRuntimeReconcileFails(t *testing.T) {
wantErr := errors.New("runtime reconcile failed")
adapter := &fakeForwardingAdapter{name: "iptables", operateErr: wantErr}
rules := &fakeForwardingRuleRepo{}
service := forwardingServiceWithAdapter(adapter)
service.rules = rules
err := service.OperateRules(dto.ForwardRuleOperate{Rules: []dto.ForwardRuleOperation{{
Operation: "add", Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}}})
if !errors.Is(err, wantErr) {
t.Fatalf("got %v want %v", err, wantErr)
}
if len(rules.rules) != 1 || rules.rules[0].Port != "8080" {
t.Fatalf("database desired state was lost after runtime failure: %#v", rules.rules)
}
}
-37
View File
@@ -1,37 +0,0 @@
package migration
import "testing"
func TestAgentDBMigrationsRegisterFirewallUpgradeSteps(t *testing.T) {
migrations := agentDBMigrations()
indexes := make(map[string]int, len(migrations))
for index, migration := range migrations {
if migration == nil || migration.ID == "" {
t.Fatalf("invalid migration at index %d: %#v", index, migration)
}
if previous, exists := indexes[migration.ID]; exists {
t.Fatalf("duplicate migration ID %q at indexes %d and %d", migration.ID, previous, index)
}
indexes[migration.ID] = index
}
tables, hasTables := indexes["20260819-add-firewall-v2-tables"]
status, hasStatus := indexes["20260818-init-docker-port-guard-status"]
selections, hasSelections := indexes["20260826-normalize-firewall-backend-selections"]
policy, hasPolicy := indexes["20260826-simplify-firewall-rule-policy"]
if !hasTables || !hasStatus || !hasSelections || !hasPolicy {
t.Fatalf("firewall upgrade migrations are not registered: %#v", indexes)
}
if _, exists := indexes["20260826-remove-firewall-rule-provider"]; exists {
t.Fatal("firewall provider removal is still registered as a separate migration")
}
if tables > status {
t.Fatalf("firewall table migration index %d runs after status migration index %d", tables, status)
}
if status > selections {
t.Fatalf("firewall status migration index %d runs after selection migration index %d", status, selections)
}
if selections > policy {
t.Fatalf("firewall selection migration index %d runs after policy migration index %d", selections, policy)
}
}
@@ -1,327 +0,0 @@
package migrations
import (
"path/filepath"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func TestAddFirewallRuleTableCreatesUpgradeSchema(t *testing.T) {
db := newFirewallMigrationTestDB(t)
for i := 0; i < 2; i++ {
if err := AddFirewallRuleTable.Migrate(db); err != nil {
t.Fatalf("migrate firewall v2 tables on pass %d: %v", i+1, err)
}
}
for _, table := range []interface{}{&model.FirewallRule{}, &model.DockerPortGuardPolicy{}, &model.ForwardingRule{}} {
if !db.Migrator().HasTable(table) {
t.Fatalf("migration did not create table for %T", table)
}
}
for _, column := range []string{
"provider", "scope_key", "location", "native_kind", "order_index", "order_bucket", "rule_key", "match_key",
} {
if db.Migrator().HasColumn("firewall_rules", column) {
t.Fatalf("new firewall desired-rule schema persisted derived column %s", column)
}
}
for _, column := range []string{"priority", "sequence"} {
if !db.Migrator().HasColumn("firewall_rules", column) {
t.Fatalf("new firewall desired-rule schema omitted placement column %s", column)
}
}
policy := model.DockerPortGuardPolicy{
UUID: "policy-1", Family: "ipv4", HostIP: "0.0.0.0", HostPort: 8080,
Protocol: "tcp", Mode: "allow_all",
}
if err := db.Create(&policy).Error; err != nil {
t.Fatalf("insert first Docker guard policy: %v", err)
}
duplicatePolicy := policy
duplicatePolicy.BaseModel = model.BaseModel{}
duplicatePolicy.UUID = "policy-2"
if err := db.Create(&duplicatePolicy).Error; err == nil {
t.Fatal("Docker guard endpoint uniqueness was not created")
}
forward := model.ForwardingRule{
Family: "ipv4", Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80",
}
if err := db.Create(&forward).Error; err != nil {
t.Fatalf("insert first forwarding rule: %v", err)
}
duplicateForward := forward
duplicateForward.BaseModel = model.BaseModel{}
if err := db.Create(&duplicateForward).Error; err == nil {
t.Fatal("forwarding identity uniqueness was not created")
}
}
func TestInitDockerPortGuardStatusCreatesDefaultOnce(t *testing.T) {
db := newFirewallMigrationTestDB(t)
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatal(err)
}
for i := 0; i < 2; i++ {
if err := InitDockerPortGuardStatus.Migrate(db); err != nil {
t.Fatalf("initialize Docker port guard status on pass %d: %v", i+1, err)
}
}
var settings []model.Setting
if err := db.Where("key = ?", constant.FirewallDockerPortGuardStatusKey).Find(&settings).Error; err != nil {
t.Fatal(err)
}
if len(settings) != 1 || settings[0].Value != constant.StatusDisable {
t.Fatalf("unexpected default Docker port guard settings: %#v", settings)
}
}
func TestInitDockerPortGuardStatusPreservesUpgradeValue(t *testing.T) {
db := newFirewallMigrationTestDB(t)
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatal(err)
}
existing := model.Setting{Key: constant.FirewallDockerPortGuardStatusKey, Value: constant.StatusEnable}
if err := db.Create(&existing).Error; err != nil {
t.Fatal(err)
}
if err := InitDockerPortGuardStatus.Migrate(db); err != nil {
t.Fatal(err)
}
var after model.Setting
if err := db.Where("key = ?", constant.FirewallDockerPortGuardStatusKey).First(&after).Error; err != nil {
t.Fatal(err)
}
if after.ID != existing.ID || after.Value != constant.StatusEnable {
t.Fatalf("migration replaced persisted Docker guard status: before=%#v after=%#v", existing, after)
}
}
func TestNormalizeFirewallBackendSelections(t *testing.T) {
db := newFirewallMigrationTestDB(t)
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatal(err)
}
settings := []model.Setting{
{Key: constant.FirewallDockerBackendKey, Value: constant.FirewallProviderNftables},
{Key: constant.FirewallForwardingBackendKey, Value: constant.FirewallProviderFirewalld},
}
if err := db.Create(&settings).Error; err != nil {
t.Fatal(err)
}
for pass := 1; pass <= 2; pass++ {
if err := NormalizeFirewallBackendSelections.Migrate(db); err != nil {
t.Fatalf("normalize firewall backend selections on pass %d: %v", pass, err)
}
}
for key, want := range map[string]string{
constant.FirewallDockerBackendKey: constant.FirewallProviderNftables,
constant.FirewallForwardingBackendKey: constant.FirewallProviderIptables,
} {
var matches []model.Setting
if err := db.Where("key = ?", key).Find(&matches).Error; err != nil {
t.Fatal(err)
}
if len(matches) != 1 || matches[0].Value != want {
t.Fatalf("setting %s after normalization = %#v, want %q", key, matches, want)
}
}
}
func TestNormalizeFirewallBackendSelectionsPreservesValidForwardingBackend(t *testing.T) {
db := newFirewallMigrationTestDB(t)
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatal(err)
}
existing := model.Setting{
Key: constant.FirewallForwardingBackendKey, Value: constant.FirewallProviderNftables,
}
if err := db.Create(&existing).Error; err != nil {
t.Fatal(err)
}
for pass := 1; pass <= 2; pass++ {
if err := NormalizeFirewallBackendSelections.Migrate(db); err != nil {
t.Fatalf("normalize firewall backend selections on pass %d: %v", pass, err)
}
}
var after model.Setting
if err := db.Where("key = ?", constant.FirewallForwardingBackendKey).First(&after).Error; err != nil {
t.Fatal(err)
}
if after.ID != existing.ID || after.Value != constant.FirewallProviderNftables {
t.Fatalf("migration replaced valid forwarding backend: before=%#v after=%#v", existing, after)
}
}
func TestNormalizeFirewallBackendSelectionsCreatesMissingDefaults(t *testing.T) {
db := newFirewallMigrationTestDB(t)
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatal(err)
}
if err := NormalizeFirewallBackendSelections.Migrate(db); err != nil {
t.Fatal(err)
}
for key, want := range map[string]string{
constant.FirewallDockerBackendKey: "",
constant.FirewallForwardingBackendKey: constant.FirewallProviderIptables,
} {
var matches []model.Setting
if err := db.Where("key = ?", key).Find(&matches).Error; err != nil {
t.Fatal(err)
}
if len(matches) != 1 || matches[0].Value != want {
t.Fatalf("missing setting %s after normalization = %#v, want %q", key, matches, want)
}
}
}
func TestSimplifyFirewallRulePolicyKeepsDesiredRules(t *testing.T) {
db := newFirewallMigrationTestDB(t)
if err := db.Exec(`CREATE TABLE firewall_rules (
uuid text PRIMARY KEY,
provider text NOT NULL,
family text NOT NULL,
protocol text NOT NULL,
scope_key text NOT NULL,
location text NOT NULL,
native_kind text NOT NULL,
priority integer,
order_index integer,
order_bucket text,
match_key text,
rule_key text NOT NULL
)`).Error; err != nil {
t.Fatal(err)
}
if err := db.Exec(
"INSERT INTO firewall_rules (uuid, provider, family, protocol, scope_key, location, native_kind, priority, order_index, order_bucket, match_key, rule_key) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
"policy-1", constant.FirewallProviderIptables, constant.FirewallFamilyIPv4, "tcp",
"iptables:ipv4:filter:1PANEL_BASIC:input", "1PANEL_BASIC", "rule", 10, 3, "legacy", "instance:legacy", "legacy-key",
).Error; err != nil {
t.Fatal(err)
}
if err := db.Exec(
"INSERT INTO firewall_rules (uuid, provider, family, protocol, scope_key, location, native_kind, priority, order_index, order_bucket, match_key, rule_key) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
"native-1", constant.FirewallProviderFirewalld, constant.FirewallFamilyInet, "all",
"firewalld:inet:public:input", "public", "zone_service", nil, nil, "legacy", "", "native-key",
).Error; err != nil {
t.Fatal(err)
}
if err := db.Exec(
"INSERT INTO firewall_rules (uuid, provider, family, protocol, scope_key, location, native_kind, priority, order_index, order_bucket, match_key, rule_key) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
"rich-1", constant.FirewallProviderFirewalld, constant.FirewallFamilyInet, "tcp",
"firewalld:inet:public:input", "public", "rich_rule", -100, nil, "rich_pre", "", "rich-key",
).Error; err != nil {
t.Fatal(err)
}
for _, statement := range []string{
"CREATE UNIQUE INDEX uk_firewall_rules_scope_rule ON firewall_rules(scope_key, rule_key)",
"CREATE UNIQUE INDEX uk_firewall_rules_scope_match ON firewall_rules(scope_key, match_key) WHERE match_key <> ''",
} {
if err := db.Exec(statement).Error; err != nil {
t.Fatal(err)
}
}
for pass := 1; pass <= 2; pass++ {
if err := SimplifyFirewallRulePolicy.Migrate(db); err != nil {
t.Fatalf("simplify firewall rule policy on pass %d: %v", pass, err)
}
}
for _, column := range []string{
"provider", "scope_key", "location", "native_kind", "order_index", "order_bucket", "rule_key", "match_key",
} {
if db.Migrator().HasColumn("firewall_rules", column) {
t.Fatalf("derived column %s was retained in firewall desired rules", column)
}
}
var positional struct {
Priority *int
Sequence *int64
}
if err := db.Table("firewall_rules").Where("uuid = ?", "policy-1").First(&positional).Error; err != nil {
t.Fatal(err)
}
if positional.Priority != nil || positional.Sequence != nil {
t.Fatalf("positional policy migration inferred unsupported placement: %#v", positional)
}
var count int64
if err := db.Table("firewall_rules").Count(&count).Error; err != nil {
t.Fatal(err)
}
if count != 3 {
t.Fatalf("policy migration removed desired rules, count=%d", count)
}
var native struct {
CompatibilityError string
}
if err := db.Table("firewall_rules").Where("uuid = ?", "native-1").First(&native).Error; err != nil {
t.Fatal(err)
}
if native.CompatibilityError == "" {
t.Fatal("legacy provider-native rule was not quarantined from synchronization")
}
var rich struct {
Priority *int
Sequence *int64
}
if err := db.Table("firewall_rules").Where("uuid = ?", "rich-1").First(&rich).Error; err != nil {
t.Fatal(err)
}
if rich.Priority != nil || rich.Sequence != nil {
t.Fatalf("firewalld policy migration retained unsupported placement: %#v", rich)
}
}
func TestSimplifyFirewallRulePolicyContinuesWhenObsoleteColumnCannotBeDropped(t *testing.T) {
db := newFirewallMigrationTestDB(t)
if err := db.Exec(`CREATE TABLE firewall_rules (
uuid text PRIMARY KEY,
provider text NOT NULL,
scope_key text NOT NULL,
location text NOT NULL,
native_kind text NOT NULL,
priority integer
)`).Error; err != nil {
t.Fatal(err)
}
if err := db.Exec("CREATE INDEX unexpected_scope_index ON firewall_rules(scope_key)").Error; err != nil {
t.Fatal(err)
}
if err := SimplifyFirewallRulePolicy.Migrate(db); err != nil {
t.Fatalf("obsolete column cleanup blocked the migration: %v", err)
}
for _, column := range []string{"provider", "location", "native_kind"} {
if db.Migrator().HasColumn("firewall_rules", column) {
t.Fatalf("obsolete column %s was not dropped after an unrelated drop failure", column)
}
}
for _, column := range []string{"compatibility_error", "sequence"} {
if !db.Migrator().HasColumn("firewall_rules", column) {
t.Fatalf("migration omitted required column %s", column)
}
}
}
func newFirewallMigrationTestDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "migration.db")), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatal(err)
}
return db
}
@@ -11,7 +11,6 @@ import (
"github.com/1Panel-dev/1Panel/agent/global"
"github.com/1Panel-dev/1Panel/agent/utils/cmd"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding"
forwardingproviders "github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding/providers"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/iptables_helper"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle"
"gorm.io/gorm"
@@ -109,7 +108,7 @@ func importLegacyForwardingRules(ctx context.Context, db *gorm.DB, rules []forwa
models := make([]model.ForwardingRule, 0, len(rules))
seen := make(map[string]struct{}, len(rules))
for _, rule := range rules {
normalized, err := forwardingproviders.NormalizeRule(rule)
normalized, err := forwarding.NormalizeRule(rule)
if err != nil {
return fmt.Errorf("normalize legacy forwarding rule: %w", err)
}
@@ -160,7 +159,7 @@ func loadLegacyFirewallForwarding() (firewallTransferSource, error) {
return firewallTransferSource{}, err
}
for _, item := range firewalldRules {
if _, err := forwardingproviders.NormalizeRule(item.rule); err != nil {
if _, err := forwarding.NormalizeRule(item.rule); err != nil {
if global.LOG != nil {
global.LOG.Warnf("skip unsupported legacy firewalld forwarding rule %q: %v", item.spec, err)
}
@@ -1,316 +0,0 @@
package utils
import (
"context"
"errors"
"path/filepath"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func TestFirewallTransferImportsCleansAndMarksCompletion(t *testing.T) {
db := newFirewallTransferTestDB(t)
loadCalls, cleanupCalls := 0, 0
transfer := &firewallTransfer{
db: db,
load: func() (firewallTransferSource, error) {
loadCalls++
legacy := legacyFirewalldForward{
rule: forwarding.Rule{Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80"},
spec: "port=8080:proto=tcp:toport=80:toaddr=10.0.0.2",
}
return firewallTransferSource{
rules: []forwarding.Rule{legacy.rule, legacy.rule},
firewalld: []legacyFirewalldForward{legacy},
provider: "iptables",
cleanupOld: func(items []legacyFirewalldForward) error {
cleanupCalls++
if len(items) != 1 {
return errors.New("unexpected cleanup inventory")
}
return nil
},
}, nil
},
}
if err := transfer.run(context.Background()); err != nil {
t.Fatal(err)
}
if loadCalls != 1 || cleanupCalls != 1 {
t.Fatalf("unexpected calls: load=%d cleanup=%d", loadCalls, cleanupCalls)
}
completed, err := firewallTransferCompleted(db)
if err != nil || !completed {
t.Fatalf("firewall transfer completion = %v, err=%v", completed, err)
}
var setting model.Setting
if err := db.Where("key = ?", "IptablesForwardStatus").First(&setting).Error; err != nil {
t.Fatal(err)
}
if setting.Value != constant.StatusEnable {
t.Fatalf("forwarding status = %q", setting.Value)
}
setting = model.Setting{}
if err := db.Where("key = ?", "ForwardingBackend").First(&setting).Error; err != nil {
t.Fatal(err)
}
if setting.Value != "iptables" {
t.Fatalf("forwarding backend = %q", setting.Value)
}
if err := transfer.run(context.Background()); err != nil {
t.Fatal(err)
}
if loadCalls != 1 || cleanupCalls != 1 {
t.Fatalf("completed transfer ran again: load=%d cleanup=%d", loadCalls, cleanupCalls)
}
}
func TestFirewallTransferFailureRemainsRetryable(t *testing.T) {
db := newFirewallTransferTestDB(t)
wantErr := errors.New("cleanup failed")
transfer := &firewallTransfer{
db: db,
load: func() (firewallTransferSource, error) {
legacy := legacyFirewalldForward{rule: forwarding.Rule{
Protocol: "udp", Port: "5353", TargetIP: "127.0.0.1", TargetPort: "53",
}}
return firewallTransferSource{
rules: []forwarding.Rule{legacy.rule},
firewalld: []legacyFirewalldForward{legacy},
provider: "iptables",
cleanupOld: func([]legacyFirewalldForward) error {
return wantErr
},
}, nil
},
}
if err := transfer.run(context.Background()); !errors.Is(err, wantErr) {
t.Fatalf("transfer error = %v, want %v", err, wantErr)
}
completed, err := firewallTransferCompleted(db)
if err != nil {
t.Fatal(err)
}
if completed {
t.Fatal("failed transfer was marked complete")
}
var count int64
if err := db.Model(&model.ForwardingRule{}).Count(&count).Error; err != nil {
t.Fatal(err)
}
if count != 1 {
t.Fatalf("retry inventory count = %d, want 1", count)
}
}
func TestFirewallTransferRetriesCleanupWithoutDuplicatingData(t *testing.T) {
db := newFirewallTransferTestDB(t)
cleanupCalls := 0
transfer := &firewallTransfer{
db: db,
load: func() (firewallTransferSource, error) {
legacy := legacyFirewalldForward{
rule: forwarding.Rule{Protocol: "tcp", Port: "8443", TargetIP: "10.0.0.8", TargetPort: "443"},
spec: "port=8443:proto=tcp:toport=443:toaddr=10.0.0.8",
}
return firewallTransferSource{
rules: []forwarding.Rule{legacy.rule}, firewalld: []legacyFirewalldForward{legacy}, provider: "iptables",
cleanupOld: func([]legacyFirewalldForward) error {
cleanupCalls++
if cleanupCalls == 1 {
return errors.New("temporary cleanup failure")
}
return nil
},
}, nil
},
}
if err := transfer.run(context.Background()); err == nil {
t.Fatal("first cleanup failure was ignored")
}
if err := transfer.run(context.Background()); err != nil {
t.Fatalf("retry firewall transfer: %v", err)
}
var count int64
if err := db.Model(&model.ForwardingRule{}).Count(&count).Error; err != nil {
t.Fatal(err)
}
if count != 1 || cleanupCalls != 2 {
t.Fatalf("retry result count=%d cleanupCalls=%d", count, cleanupCalls)
}
completed, err := firewallTransferCompleted(db)
if err != nil || !completed {
t.Fatalf("retry completion = %v, err=%v", completed, err)
}
}
func TestFirewallTransferEmptyInventoryMarksCompletionWithoutEnabling(t *testing.T) {
db := newFirewallTransferTestDB(t)
transfer := &firewallTransfer{
db: db,
load: func() (firewallTransferSource, error) {
return firewallTransferSource{}, nil
},
}
if err := transfer.run(context.Background()); err != nil {
t.Fatal(err)
}
completed, err := firewallTransferCompleted(db)
if err != nil || !completed {
t.Fatalf("empty transfer completion = %v, err=%v", completed, err)
}
var count int64
if err := db.Model(&model.Setting{}).Where("key IN ?", []string{"IptablesForwardStatus", "ForwardingBackend"}).Count(&count).Error; err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("empty transfer created %d forwarding settings", count)
}
}
func TestFirewallTransferInvalidInventoryIsAtomicAndRetryable(t *testing.T) {
db := newFirewallTransferTestDB(t)
transfer := &firewallTransfer{
db: db,
load: func() (firewallTransferSource, error) {
return firewallTransferSource{rules: []forwarding.Rule{
{Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80"},
{Protocol: "tcp", Port: "invalid", TargetIP: "127.0.0.1", TargetPort: "80"},
}}, nil
},
}
if err := transfer.run(context.Background()); err == nil {
t.Fatal("invalid legacy forwarding inventory was accepted")
}
completed, err := firewallTransferCompleted(db)
if err != nil {
t.Fatal(err)
}
var count int64
if err := db.Model(&model.ForwardingRule{}).Count(&count).Error; err != nil {
t.Fatal(err)
}
if completed || count != 0 {
t.Fatalf("invalid transfer completed=%v imported=%d", completed, count)
}
}
func TestFirewallTransferUpdatesExistingSettings(t *testing.T) {
db := newFirewallTransferTestDB(t)
settings := []model.Setting{
{Key: "IptablesForwardStatus", Value: constant.StatusDisable},
{Key: "ForwardingBackend", Value: "nftables"},
}
if err := db.Create(&settings).Error; err != nil {
t.Fatal(err)
}
transfer := &firewallTransfer{
db: db,
load: func() (firewallTransferSource, error) {
return firewallTransferSource{
rules: []forwarding.Rule{{Protocol: "udp", Port: "5353", TargetIP: "127.0.0.1", TargetPort: "53"}},
provider: "iptables",
}, nil
},
}
if err := transfer.run(context.Background()); err != nil {
t.Fatal(err)
}
for key, want := range map[string]string{
"IptablesForwardStatus": constant.StatusEnable,
"ForwardingBackend": "iptables",
} {
var matches []model.Setting
if err := db.Where("key = ?", key).Find(&matches).Error; err != nil {
t.Fatal(err)
}
if len(matches) != 1 || matches[0].Value != want {
t.Fatalf("setting %s after transfer = %#v, want %q", key, matches, want)
}
}
}
func TestFirewallTransferValidatesDependencies(t *testing.T) {
if err := (&firewallTransfer{}).run(context.Background()); err == nil {
t.Fatal("nil transfer database was accepted")
}
db := newFirewallTransferTestDB(t)
if err := (&firewallTransfer{db: db}).run(context.Background()); err == nil {
t.Fatal("nil legacy loader was accepted")
}
transfer := &firewallTransfer{
db: db,
load: func() (firewallTransferSource, error) {
legacy := legacyFirewalldForward{rule: forwarding.Rule{
Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80",
}}
return firewallTransferSource{rules: []forwarding.Rule{legacy.rule}, firewalld: []legacyFirewalldForward{legacy}}, nil
},
}
if err := transfer.run(context.Background()); err == nil {
t.Fatal("missing firewalld cleanup was accepted")
}
completed, err := firewallTransferCompleted(db)
if err != nil || completed {
t.Fatalf("dependency failure completion=%v err=%v", completed, err)
}
}
func TestParseLegacyFirewalldForwarding(t *testing.T) {
rules := parseLegacyFirewalldForwarding(
"port=8080:proto=tcp:toport=80:toaddr=10.0.0.2\n" +
"port=8443:proto=tcp:toport=443:toaddr=\ninvalid\n",
)
if len(rules) != 2 {
t.Fatalf("parsed %d firewalld rules", len(rules))
}
if rules[0].rule.Family != forwarding.FamilyIPv4 || rules[0].rule.TargetIP != "10.0.0.2" {
t.Fatalf("unexpected remote rule: %#v", rules[0])
}
if rules[1].rule.TargetIP != "127.0.0.1" || rules[1].rule.TargetPort != "443" {
t.Fatalf("unexpected local rule: %#v", rules[1])
}
}
func TestParseLegacyIptablesForwarding(t *testing.T) {
stdout := "1 0 0 DNAT 6 -- eth0 * 0.0.0.0/0 0.0.0.0/0 tcp dpt:8080 to:10.0.0.2:80\n" +
"2 0 0 REDIRECT 17 -- * * 0.0.0.0/0 0.0.0.0/0 udp dpts:5353:5354 redir ports 53\n"
rules := parseLegacyIptablesForwarding(stdout)
if len(rules) != 2 {
t.Fatalf("parsed %d iptables rules", len(rules))
}
if rules[0].Protocol != "tcp" || rules[0].Interface != "eth0" || rules[0].Port != "8080" ||
rules[0].TargetIP != "10.0.0.2" || rules[0].TargetPort != "80" {
t.Fatalf("unexpected DNAT rule: %#v", rules[0])
}
if rules[1].Protocol != "udp" || rules[1].Port != "5353-5354" ||
rules[1].TargetIP != "127.0.0.1" || rules[1].TargetPort != "53" {
t.Fatalf("unexpected REDIRECT rule: %#v", rules[1])
}
}
func newFirewallTransferTestDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "migration.db")), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.Setting{}, &model.ForwardingRule{}); err != nil {
t.Fatal(err)
}
if err := db.Exec("CREATE TABLE migrations (id VARCHAR(255) PRIMARY KEY)").Error; err != nil {
t.Fatal(err)
}
return db
}
@@ -1,365 +0,0 @@
package utils
import (
"context"
"errors"
"path/filepath"
"sort"
"testing"
"github.com/1Panel-dev/1Panel/agent/app/model"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func TestTransferHostFirewallImportsIptablesRecords(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
records := []legacyHostFirewallRecord{
{Type: "port", Protocol: "tcp/udp", SrcIP: "Anywhere", DstPort: "80", Strategy: "accept", Description: "web"},
{Type: "address", SrcIP: "10.0.0.8", Strategy: "drop", Description: "blocked host"},
{Chain: "1PANEL_BASIC_BEFORE", Protocol: "tcp", DstPort: "22", Strategy: "accept", Description: "ssh first"},
{Chain: "1PANEL_INPUT", Protocol: "tcp", DstPort: "9000", Strategy: "accept", Description: "unsupported chain"},
}
if err := db.Table("firewalls").Create(&records).Error; err != nil {
t.Fatal(err)
}
if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil {
t.Fatal(err)
}
rules := loadTransferredHostFirewallRules(t, db)
if len(rules) != 4 {
t.Fatalf("transferred %d rules, want 4", len(rules))
}
assertTransferredRuleDefaults(t, rules)
protocols := make([]string, 0, 2)
foundAddress, foundAdvanced := false, false
for _, rule := range rules {
switch rule.Description {
case "web":
protocols = append(protocols, rule.Protocol)
if rule.Family != string(filter.FamilyIPv4) || rule.DestinationPort != "80" {
t.Fatalf("unexpected port rule: %#v", rule)
}
case "blocked host":
foundAddress = rule.SourceAddress == "10.0.0.8/32" && rule.Action == string(filter.ActionDrop)
case "ssh first":
foundAdvanced = rule.DestinationPort == "22"
case "unsupported chain":
t.Fatal("unsupported legacy chain was imported")
}
}
sort.Strings(protocols)
if len(protocols) != 2 || protocols[0] != "tcp" || protocols[1] != "udp" {
t.Fatalf("port protocols = %#v", protocols)
}
if !foundAddress || !foundAdvanced {
t.Fatalf("address=%v advanced=%v", foundAddress, foundAdvanced)
}
assertHostFirewallTransferCompleted(t, db)
if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil {
t.Fatal(err)
}
if got := len(loadTransferredHostFirewallRules(t, db)); got != 4 {
t.Fatalf("retry transferred %d rules, want 4", got)
}
}
func TestTransferHostFirewallMapsFirewalldRepresentations(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
records := []legacyHostFirewallRecord{
{Type: "port", Protocol: "tcp/udp", DstPort: "80", Strategy: "accept", Description: "zone ports"},
{Type: "port", Protocol: "tcp", DstPort: "81", Strategy: "drop", Description: "dual family deny"},
{Type: "port", Protocol: "tcp", SrcIP: "2001:db8::8", DstPort: "443", Strategy: "accept", Description: "v6 source"},
{Type: "address", SrcIP: "10.0.0.9", Strategy: "drop", Description: "v4 address"},
}
if err := db.Table("firewalls").Create(&records).Error; err != nil {
t.Fatal(err)
}
if err := transferHostFirewall(context.Background(), db, filter.ProviderFirewalld); err != nil {
t.Fatal(err)
}
rules := loadTransferredHostFirewallRules(t, db)
if len(rules) != 6 {
t.Fatalf("transferred %d rules, want 6", len(rules))
}
zonePorts, denyFamilies := 0, make(map[string]bool)
for _, rule := range rules {
switch rule.Description {
case "zone ports":
zonePorts++
if rule.Family != string(filter.FamilyInet) {
t.Fatalf("unexpected zone port: %#v", rule)
}
case "dual family deny":
denyFamilies[rule.Family] = true
case "v6 source":
if rule.Family != string(filter.FamilyIPv6) || rule.SourceAddress != "2001:db8::8/128" {
t.Fatalf("unexpected v6 rule: %#v", rule)
}
}
}
if zonePorts != 2 || !denyFamilies[string(filter.FamilyIPv4)] || !denyFamilies[string(filter.FamilyIPv6)] {
t.Fatalf("zonePorts=%d denyFamilies=%#v", zonePorts, denyFamilies)
}
}
func TestTransferHostFirewallExpandsUFWFamilies(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
records := []legacyHostFirewallRecord{
{Type: "port", Protocol: "tcp", DstPort: "8080", Strategy: "accept", Description: "dual family port"},
{Type: "address", SrcIP: "10.0.0.1-10.0.0.2", Strategy: "drop", Description: "from to"},
}
if err := db.Table("firewalls").Create(&records).Error; err != nil {
t.Fatal(err)
}
if err := transferHostFirewall(context.Background(), db, filter.ProviderUFW); err != nil {
t.Fatal(err)
}
rules := loadTransferredHostFirewallRules(t, db)
if len(rules) != 3 {
t.Fatalf("transferred %d rules, want 3", len(rules))
}
portFamilies := make(map[string]bool)
foundFromTo := false
for _, rule := range rules {
if rule.Description == "dual family port" {
portFamilies[rule.Family] = true
}
if rule.Description == "from to" {
foundFromTo = rule.SourceAddress == "10.0.0.1/32" && rule.DestinationAddress == "10.0.0.2/32"
}
}
if !portFamilies[string(filter.FamilyIPv4)] || !portFamilies[string(filter.FamilyIPv6)] || !foundFromTo {
t.Fatalf("portFamilies=%#v foundFromTo=%v", portFamilies, foundFromTo)
}
}
func TestTransferHostFirewallWithoutLegacyTableOnlyMarksCompletion(t *testing.T) {
db := newHostFirewallTransferTestDB(t, false)
if err := transferHostFirewall(context.Background(), db, filter.ProviderNftables); err != nil {
t.Fatal(err)
}
if got := len(loadTransferredHostFirewallRules(t, db)); got != 0 {
t.Fatalf("transferred %d rules without a legacy table", got)
}
assertHostFirewallTransferCompleted(t, db)
}
func TestTransferHostFirewallRestoresMissingDescription(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
record := legacyHostFirewallRecord{
Type: "port", Protocol: "tcp", DstPort: "443", Strategy: "accept", Description: "legacy tls",
}
if err := db.Table("firewalls").Create(&record).Error; err != nil {
t.Fatal(err)
}
domainRules, err := legacyHostFirewallRules(record, filter.ProviderIptables)
if err != nil {
t.Fatal(err)
}
existing, err := hostFirewallRuleModel(domainRules[0])
if err != nil {
t.Fatal(err)
}
existing.Description = ""
existing.Origin = constant.FirewallRuleOriginCreated
if err := db.Create(&existing).Error; err != nil {
t.Fatal(err)
}
if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil {
t.Fatal(err)
}
rules := loadTransferredHostFirewallRules(t, db)
if len(rules) != 1 || rules[0].Description != "legacy tls" || rules[0].Origin != constant.FirewallRuleOriginCreated {
t.Fatalf("unexpected existing rule after transfer: %#v", rules)
}
}
func TestTransferHostFirewallReadsDeprecatedColumnsOnDirectUpgrade(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
records := []legacyHostFirewallRecord{
{Type: "port", Port: "8080", Address: "Anywhere", Protocol: "tcp", Strategy: "accept", Description: "legacy port"},
{Type: "address", Address: "192.0.2.25", Strategy: "drop", Description: "legacy address"},
{
Type: "port", Port: "9000", Address: "192.0.2.90", DstPort: "9001", SrcIP: "192.0.2.91",
Protocol: "udp", Strategy: "accept", Description: "new columns win",
},
}
if err := db.Table("firewalls").Create(&records).Error; err != nil {
t.Fatal(err)
}
if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil {
t.Fatal(err)
}
rules := loadTransferredHostFirewallRules(t, db)
if len(rules) != 3 {
t.Fatalf("transferred %d direct-upgrade rules, want 3", len(rules))
}
for _, rule := range rules {
switch rule.Description {
case "legacy port":
if rule.DestinationPort != "8080" || rule.SourceAddress != "" {
t.Fatalf("deprecated port columns were not migrated: %#v", rule)
}
case "legacy address":
if rule.SourceAddress != "192.0.2.25/32" || rule.Action != string(filter.ActionDrop) {
t.Fatalf("deprecated address column was not migrated: %#v", rule)
}
case "new columns win":
if rule.DestinationPort != "9001" || rule.SourceAddress != "192.0.2.91/32" {
t.Fatalf("deprecated columns overwrote normalized columns: %#v", rule)
}
}
}
}
func TestTransferHostFirewallCoalescesDuplicatesAndSkipsInvalidRows(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
records := []legacyHostFirewallRecord{
{Type: "port", Protocol: "tcp", DstPort: "443", Strategy: "accept", Description: "first"},
{Type: "port", Protocol: "tcp", DstPort: "443", Strategy: "accept", Description: "latest"},
{Type: "port", Protocol: "tcp", DstPort: "invalid", Strategy: "accept", Description: "invalid"},
}
if err := db.Table("firewalls").Create(&records).Error; err != nil {
t.Fatal(err)
}
if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil {
t.Fatal(err)
}
rules := loadTransferredHostFirewallRules(t, db)
if len(rules) != 1 || rules[0].Description != "latest" || rules[0].DestinationPort != "443" {
t.Fatalf("unexpected duplicate/invalid migration result: %#v", rules)
}
assertHostFirewallTransferCompleted(t, db)
}
func TestTransferHostFirewallFailureDoesNotMarkCompletion(t *testing.T) {
t.Run("unsupported provider", func(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
err := transferHostFirewall(context.Background(), db, filter.Provider("unknown"))
if err == nil {
t.Fatal("unsupported provider was accepted")
}
assertHostFirewallTransferNotCompleted(t, db)
})
t.Run("cancelled context", func(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
record := legacyHostFirewallRecord{Type: "port", Protocol: "tcp", DstPort: "80", Strategy: "accept"}
if err := db.Table("firewalls").Create(&record).Error; err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
err := transferHostFirewall(ctx, db, filter.ProviderIptables)
if !errors.Is(err, context.Canceled) {
t.Fatalf("cancelled transfer error = %v", err)
}
assertHostFirewallTransferNotCompleted(t, db)
})
}
func TestTransferHostFirewallPreservesExistingDescription(t *testing.T) {
db := newHostFirewallTransferTestDB(t, true)
record := legacyHostFirewallRecord{
Type: "port", Protocol: "tcp", DstPort: "443", Strategy: "accept", Description: "legacy description",
}
if err := db.Table("firewalls").Create(&record).Error; err != nil {
t.Fatal(err)
}
domainRules, err := legacyHostFirewallRules(record, filter.ProviderIptables)
if err != nil {
t.Fatal(err)
}
existing, err := hostFirewallRuleModel(domainRules[0])
if err != nil {
t.Fatal(err)
}
existing.Description = "user description"
if err := db.Create(&existing).Error; err != nil {
t.Fatal(err)
}
if err := transferHostFirewall(context.Background(), db, filter.ProviderIptables); err != nil {
t.Fatal(err)
}
rules := loadTransferredHostFirewallRules(t, db)
if len(rules) != 1 || rules[0].Description != "user description" {
t.Fatalf("migration overwrote existing description: %#v", rules)
}
}
func newHostFirewallTransferTestDB(t *testing.T, withLegacyTable bool) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "migration.db")), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.FirewallRule{}); err != nil {
t.Fatal(err)
}
if err := db.Exec("CREATE TABLE migrations (id VARCHAR(255) PRIMARY KEY)").Error; err != nil {
t.Fatal(err)
}
if withLegacyTable {
if err := db.Table("firewalls").AutoMigrate(&legacyHostFirewallRecord{}); err != nil {
t.Fatal(err)
}
}
return db
}
func loadTransferredHostFirewallRules(t *testing.T, db *gorm.DB) []model.FirewallRule {
t.Helper()
var rules []model.FirewallRule
if err := db.Order("family ASC, protocol ASC, destination_port ASC, uuid ASC").Find(&rules).Error; err != nil {
t.Fatal(err)
}
return rules
}
func assertTransferredRuleDefaults(t *testing.T, rules []model.FirewallRule) {
t.Helper()
for _, rule := range rules {
if rule.UUID == "" || rule.Revision != 1 || rule.CompatibilityError != "" ||
rule.Priority != nil || rule.Sequence != nil ||
rule.Origin != constant.FirewallRuleOriginAdopted || rule.Owner != constant.FirewallRuleSourceUser {
t.Fatalf("unexpected transferred defaults: %#v", rule)
}
}
}
func assertHostFirewallTransferCompleted(t *testing.T, db *gorm.DB) {
t.Helper()
completed, err := migrationRecordExists(db, hostFirewallTransferMigrationID)
if err != nil {
t.Fatal(err)
}
if !completed {
t.Fatal("host firewall transfer was not marked complete")
}
}
func assertHostFirewallTransferNotCompleted(t *testing.T, db *gorm.DB) {
t.Helper()
completed, err := migrationRecordExists(db, hostFirewallTransferMigrationID)
if err != nil {
t.Fatal(err)
}
if completed {
t.Fatal("failed host firewall transfer was marked complete")
}
}
+3
View File
@@ -3,6 +3,7 @@ package docker
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"os"
@@ -22,6 +23,8 @@ import (
"github.com/docker/docker/client"
)
var ErrUnavailable = errors.New("Docker is unavailable")
func NewDockerClient() (*client.Client, error) {
var settingItem model.Setting
_ = global.DB.Where("key = ?", "DockerSockPath").First(&settingItem).Error
@@ -1,255 +0,0 @@
package docker_guard
import (
"errors"
"reflect"
"strings"
"testing"
)
type restoreCall struct {
executable string
input string
args []string
}
type recordingRunner struct {
restoreCalls []restoreCall
exists map[string]bool
chains string
dockerRules string
guardRules string
runErr error
}
func (r *recordingRunner) Run(_ string, args ...string) (string, error) {
if r.runErr != nil {
return "", r.runErr
}
if reflect.DeepEqual(args, []string{"-w", "-t", "filter", "-S"}) {
if r.chains != "" {
return r.chains, nil
}
return "-N DOCKER-USER\n-N 1PANEL_DOCKER\n", nil
}
if reflect.DeepEqual(args, []string{"-w", "-t", "filter", "-S", DockerChain}) {
if r.dockerRules != "" {
return r.dockerRules, nil
}
return "-A DOCKER-USER -j 1PANEL_DOCKER\n", nil
}
if reflect.DeepEqual(args, []string{"-w", "-t", "filter", "-S", Chain}) {
return r.guardRules, nil
}
return "", nil
}
func (r *recordingRunner) RunInput(executable, input string, args ...string) (string, error) {
r.restoreCalls = append(r.restoreCalls, restoreCall{executable: executable, input: input, args: args})
return "", nil
}
func (r *recordingRunner) Exists(executable string) bool {
if r.exists != nil {
return r.exists[executable]
}
return executable == "iptables" || executable == "iptables-restore"
}
func TestCompilePolicyUsesOriginalDestination(t *testing.T) {
rules := compilePolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "192.0.2.1", HostPort: 8080, Protocol: "tcp", Mode: ModeSources, Sources: []string{"203.0.113.0/24"}})
got := strings.Join(rules[0], " ")
for _, want := range []string{"--ctorigdst 192.0.2.1", "--ctorigdstport 8080", "-s 203.0.113.0/24", "-j DROP"} {
if !strings.Contains(got, want) {
t.Fatalf("compiled rule %q does not contain %q", got, want)
}
}
}
func TestCompileWildcardDoesNotMatchUnroutableWildcardAddress(t *testing.T) {
rules := compilePolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 53, Protocol: "udp", Mode: ModeAll})
got := strings.Join(rules[0], " ")
if strings.Contains(got, "--ctorigdst ") {
t.Fatalf("wildcard binding must not compile an original destination address: %s", got)
}
if !strings.Contains(got, "--ctorigdstport 53") {
t.Fatalf("original destination port missing: %s", got)
}
}
func TestCompileAllowSourcesReturnsAllowedAndDropsOthers(t *testing.T) {
rules := compilePolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow, Sources: []string{"203.0.113.10/32", "192.0.2.0/24"}})
if len(rules) != 3 {
t.Fatalf("compiled %d rules, want 3", len(rules))
}
for i, source := range []string{"203.0.113.10/32", "192.0.2.0/24"} {
got := strings.Join(rules[i], " ")
if !strings.Contains(got, "-s "+source) || !strings.HasSuffix(got, "-j RETURN") {
t.Fatalf("allow rule = %q", got)
}
}
if got := strings.Join(rules[2], " "); strings.Contains(got, " -s ") || !strings.HasSuffix(got, "-j DROP") {
t.Fatalf("fallback rule = %q", got)
}
}
func TestCompileEmptyAllowSourcesDropsAll(t *testing.T) {
rules := compilePolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow})
if len(rules) != 1 || !strings.HasSuffix(strings.Join(rules[0], " "), "-j DROP") {
t.Fatalf("rules = %#v", rules)
}
}
func TestParseIptablesDockerGuardPolicies(t *testing.T) {
output := strings.Join([]string{
`-A 1PANEL_DOCKER -p tcp -m conntrack --ctorigdstport 8080 -m comment --comment "1panel-docker:deny" -j DROP`,
`-A 1PANEL_DOCKER -p tcp -m conntrack --ctorigdst 192.0.2.10/32 --ctorigdstport 5432 -s 203.0.113.1/32 -m comment --comment "1panel-docker:allow" -j RETURN`,
`-A 1PANEL_DOCKER -p tcp -m conntrack --ctorigdst 192.0.2.10/32 --ctorigdstport 5432 -m comment --comment "1panel-docker:allow" -j DROP`,
}, "\n")
policies, err := parseDockerGuardPolicies(output, FamilyIPv4)
if err != nil {
t.Fatal(err)
}
if len(policies) != 2 || policies[0].UUID != "deny" || policies[0].Mode != ModeAll ||
policies[1].UUID != "allow" || policies[1].HostIP != "192.0.2.10" || policies[1].Mode != ModeAllow ||
!reflect.DeepEqual(policies[1].Sources, []string{"203.0.113.1/32"}) {
t.Fatalf("policies = %#v", policies)
}
}
func TestEffectiveJumpMustBeFirstAndUnique(t *testing.T) {
if !hasFirstUniqueJump("-A DOCKER-USER -j 1PANEL_DOCKER\n-A DOCKER-USER -j RETURN\n") {
t.Fatal("expected first unique jump to be effective")
}
if hasFirstUniqueJump("-A DOCKER-USER -j OTHER\n-A DOCKER-USER -j 1PANEL_DOCKER\n") {
t.Fatal("jump after another rule must not be reported effective")
}
if hasFirstUniqueJump("-A DOCKER-USER -j 1PANEL_DOCKER\n-A DOCKER-USER -j 1PANEL_DOCKER\n") {
t.Fatal("duplicate jumps must not be reported effective")
}
}
func TestReconcileUsesSingleAtomicRestorePerFamily(t *testing.T) {
runner := &recordingRunner{}
manager := NewManagerWithRunner(runner)
policies := []Policy{
{UUID: "first", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, Protocol: "tcp", Mode: ModeAll},
{UUID: "second", Family: FamilyIPv4, HostIP: "192.0.2.10", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow, Sources: []string{"203.0.113.1/32"}},
}
if err := manager.Reconcile(policies); err != nil {
t.Fatal(err)
}
if len(runner.restoreCalls) != 1 {
t.Fatalf("restore calls = %d, want 1", len(runner.restoreCalls))
}
call := runner.restoreCalls[0]
if call.executable != "iptables-restore" || !reflect.DeepEqual(call.args, []string{"--noflush", "--wait"}) {
t.Fatalf("restore call = %#v", call)
}
for _, want := range []string{
"*filter\n",
"-F 1PANEL_DOCKER\n",
"-A 1PANEL_DOCKER -m conntrack --ctstate RELATED,ESTABLISHED -j RETURN\n",
"--ctorigdstport 8080",
"--ctorigdst 192.0.2.10 --ctorigdstport 5432",
"-s 203.0.113.1/32",
"-A 1PANEL_DOCKER -j RETURN\nCOMMIT\n",
} {
if !strings.Contains(call.input, want) {
t.Fatalf("restore input does not contain %q:\n%s", want, call.input)
}
}
if strings.Contains(call.input, "-F DOCKER-USER") {
t.Fatalf("restore input must not flush Docker chains:\n%s", call.input)
}
}
func TestDockerGuardLifecycleRulesBatchCreateAndRebind(t *testing.T) {
output := strings.Join([]string{
"-N " + DockerChain,
"-A " + DockerChain + " -j " + Chain,
"-A " + DockerChain + " -j " + Chain,
}, "\n")
rules := dockerGuardLifecycleRules(output, true, true)
script, err := buildRestoreScript(rules)
if err != nil {
t.Fatal(err)
}
for _, line := range []string{
"-N " + Chain,
"-D " + DockerChain + " -j " + Chain,
"-I " + DockerChain + " 1 -j " + Chain,
} {
if !strings.Contains(script, line+"\n") {
t.Fatalf("lifecycle restore is missing %q:\n%s", line, script)
}
}
if strings.Count(script, "-D "+DockerChain+" -j "+Chain+"\n") != 2 || strings.Count(script, "COMMIT\n") != 1 {
t.Fatalf("duplicate jumps were not removed in one transaction:\n%s", script)
}
}
func TestCleanupRemovesExistingChainWithoutRecreatingIt(t *testing.T) {
runner := &recordingRunner{}
manager := NewManagerWithRunner(runner)
if err := manager.Cleanup(); err != nil {
t.Fatal(err)
}
if len(runner.restoreCalls) != 1 {
t.Fatalf("restore calls = %d, want 1", len(runner.restoreCalls))
}
script := runner.restoreCalls[0].input
if strings.Contains(script, "-N "+Chain+"\n") {
t.Fatalf("cleanup must not recreate the existing chain:\n%s", script)
}
for _, want := range []string{"-F " + Chain + "\n", "-X " + Chain + "\n"} {
if !strings.Contains(script, want) {
t.Fatalf("cleanup restore is missing %q:\n%s", want, script)
}
}
}
func TestReconcileReturnsChainInspectionError(t *testing.T) {
manager := NewManagerWithRunner(&recordingRunner{runErr: errors.New("inspect failed")})
err := manager.Reconcile(nil)
if err == nil {
t.Fatal("expected chain inspection error")
}
var familyErr *FamilyError
if !errors.As(err, &familyErr) || familyErr.Family != FamilyIPv4 {
t.Fatalf("error = %#v, want IPv4 FamilyError", err)
}
}
func TestBuildRestoreScriptRejectsUnsafeTokens(t *testing.T) {
if _, err := buildRestoreScript([][]string{{"-A", Chain, "--comment", "unsafe value"}}); err == nil {
t.Fatal("expected unsafe token to be rejected")
}
}
func TestFamilyStatusExplainsIncompleteStep(t *testing.T) {
tests := []struct {
name string
runner *recordingRunner
wantState string
wantReason string
initialized bool
bound bool
}{
{name: "command missing", runner: &recordingRunner{exists: map[string]bool{}}, wantState: StatusDisabled, wantReason: ReasonCommandMissing},
{name: "Docker chain missing", runner: &recordingRunner{chains: "-N 1PANEL_DOCKER\n"}, wantState: StatusDisabled, wantReason: ReasonDockerChainMissing},
{name: "guard chain missing", runner: &recordingRunner{chains: "-N DOCKER-USER\n"}, wantState: StatusDisabled, wantReason: ReasonGuardChainMissing},
{name: "jump missing", runner: &recordingRunner{dockerRules: "-A DOCKER-USER -j RETURN\n"}, wantState: StatusNotEffective, wantReason: ReasonJumpMissing, initialized: true},
{name: "jump not first", runner: &recordingRunner{dockerRules: "-A DOCKER-USER -j OTHER\n-A DOCKER-USER -j 1PANEL_DOCKER\n"}, wantState: StatusNotEffective, wantReason: ReasonJumpNotFirst, initialized: true},
{name: "jump duplicate", runner: &recordingRunner{dockerRules: "-A DOCKER-USER -j 1PANEL_DOCKER\n-A DOCKER-USER -j 1PANEL_DOCKER\n"}, wantState: StatusNotEffective, wantReason: ReasonJumpDuplicate, initialized: true},
{name: "effective", runner: &recordingRunner{}, wantState: StatusEffective, initialized: true, bound: true},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
status := NewManagerWithRunner(test.runner).Status(FamilyIPv4)
if status.State != test.wantState || status.Reason != test.wantReason || status.Initialized != test.initialized || status.Bound != test.bound {
t.Fatalf("status = %#v", status)
}
})
}
}
@@ -1,210 +0,0 @@
package docker_guard
import (
"errors"
"reflect"
"strings"
"testing"
)
type nftRecordingRunner struct {
objects map[string]bool
baseRules map[string]string
policyRules map[string]string
runCalls [][]string
inputCalls []restoreCall
}
func newNftRecordingRunner() *nftRecordingRunner {
return &nftRecordingRunner{objects: map[string]bool{}, baseRules: map[string]string{}, policyRules: map[string]string{}}
}
func (r *nftRecordingRunner) Run(executable string, args ...string) (string, error) {
r.runCalls = append(r.runCalls, append([]string{executable}, args...))
if executable != "nft" {
return "", errors.New("unexpected executable")
}
plain := args
if len(plain) > 0 && plain[0] == "-a" {
plain = plain[1:]
}
if len(plain) >= 3 && plain[0] == "list" {
key := strings.Join(plain[1:], "|")
if !r.objects[key] {
return "", errors.New("not found")
}
if plain[1] == "chain" && len(plain) == 5 && plain[4] == NftBaseChain {
return r.baseRules[plain[2]], nil
}
if plain[1] == "chain" && len(plain) == 5 && plain[4] == NftChain {
return r.policyRules[plain[2]], nil
}
return "", nil
}
if len(plain) >= 4 && plain[0] == "add" && (plain[1] == "table" || plain[1] == "chain") {
nameLen := 3
if plain[1] == "chain" {
nameLen = 4
}
r.objects[strings.Join(plain[1:1+nameLen], "|")] = true
return "", nil
}
if len(plain) >= 7 && plain[0] == "insert" && plain[1] == "rule" {
r.baseRules[plain[2]] = "jump " + NftChain + " # handle 1\n"
return "", nil
}
return "", nil
}
func (r *nftRecordingRunner) RunInput(executable, input string, args ...string) (string, error) {
r.inputCalls = append(r.inputCalls, restoreCall{executable: executable, input: input, args: args})
for _, line := range strings.Split(input, "\n") {
fields := strings.Fields(line)
if len(fields) >= 4 && fields[0] == "add" && fields[1] == "table" {
r.objects["table|"+fields[2]+"|"+fields[3]] = true
}
if len(fields) >= 5 && fields[0] == "add" && fields[1] == "chain" {
r.objects["chain|"+fields[2]+"|"+fields[3]+"|"+fields[4]] = true
}
if len(fields) >= 7 && fields[0] == "insert" && fields[1] == "rule" {
r.baseRules[fields[2]] = "jump " + NftChain + " # handle 1\n"
}
}
return "", nil
}
func (r *nftRecordingRunner) Exists(executable string) bool { return executable == "nft" }
func (r *nftRecordingRunner) addTable(family, table string) {
r.objects["table|"+family+"|"+table] = true
}
func (r *nftRecordingRunner) addChain(family, table, chain string) {
r.objects["chain|"+family+"|"+table+"|"+chain] = true
}
func TestCompileNftPolicyUsesOriginalDestination(t *testing.T) {
rules := compileNftPolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "192.0.2.1", HostPort: 8080, Protocol: "tcp", Mode: ModeSources, Sources: []string{"203.0.113.0/24"}})
got := strings.Join(rules[0], " ")
for _, want := range []string{"meta l4proto tcp", "ct original ip daddr 192.0.2.1", "ct original proto-dst 8080", "ip saddr 203.0.113.0/24", "drop"} {
if !strings.Contains(got, want) {
t.Fatalf("compiled rule %q does not contain %q", got, want)
}
}
}
func TestCompileNftWildcardDoesNotMatchWildcardAddress(t *testing.T) {
rules := compileNftPolicy(Policy{UUID: "id", Family: FamilyIPv6, HostIP: "::", HostPort: 53, Protocol: "udp", Mode: ModeAll})
got := strings.Join(rules[0], " ")
if strings.Contains(got, "ct original ip6 daddr") {
t.Fatalf("wildcard binding must not compile an original destination address: %s", got)
}
if !strings.Contains(got, "ct original proto-dst 53") {
t.Fatalf("original destination port missing: %s", got)
}
}
func TestCompileNftEmptyAllowSourcesDropsAll(t *testing.T) {
rules := compileNftPolicy(Policy{UUID: "id", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow})
if len(rules) != 1 || !strings.HasSuffix(strings.Join(rules[0], " "), "drop") {
t.Fatalf("rules = %#v", rules)
}
}
func TestParseNftablesDockerGuardPolicies(t *testing.T) {
output := strings.Join([]string{
`meta l4proto udp ct original proto-dst 53 comment "1panel-docker:deny" drop # handle 3`,
`meta l4proto tcp ct original ip daddr 192.0.2.10 ct original proto-dst 5432 ip saddr 203.0.113.1/32 comment "1panel-docker:allow" return # handle 4`,
`meta l4proto tcp ct original ip daddr 192.0.2.10 ct original proto-dst 5432 comment "1panel-docker:allow" drop # handle 5`,
}, "\n")
policies, err := parseDockerGuardPolicies(output, FamilyIPv4)
if err != nil {
t.Fatal(err)
}
if len(policies) != 2 || policies[0].UUID != "deny" || policies[0].Mode != ModeAll ||
policies[1].UUID != "allow" || policies[1].HostIP != "192.0.2.10" || policies[1].Mode != ModeAllow ||
!reflect.DeepEqual(policies[1].Sources, []string{"203.0.113.1/32"}) {
t.Fatalf("policies = %#v", policies)
}
}
func TestNftInitializeCreatesOwnedChainsBeforeDockerForwardRules(t *testing.T) {
runner := newNftRecordingRunner()
runner.addTable("ip", dockerNftTable)
manager := NewNftablesManagerWithRunner(runner)
if err := manager.Initialize(nil); err != nil {
t.Fatal(err)
}
if len(runner.inputCalls) != 2 {
t.Fatalf("batch calls = %d, want lifecycle plus rule restore", len(runner.inputCalls))
}
all := runner.inputCalls[0].input
for _, want := range []string{
"add table ip " + NftTable,
"add chain ip " + NftTable + " " + NftBaseChain + " { type filter hook forward priority filter - 1 ; policy accept ; }",
"add chain ip " + NftTable + " " + NftChain,
"insert rule ip " + NftTable + " " + NftBaseChain + " jump " + NftChain,
} {
if !strings.Contains(all, want) {
t.Fatalf("nft lifecycle batch does not contain %q:\n%s", want, all)
}
}
if !strings.Contains(runner.inputCalls[1].input, "flush chain ip "+NftTable+" "+NftChain) {
t.Fatalf("policy restore batch was not executed:\n%s", runner.inputCalls[1].input)
}
}
func TestNftReconcileUsesSingleAtomicScriptPerFamily(t *testing.T) {
runner := newNftRecordingRunner()
runner.addChain("ip", NftTable, NftChain)
manager := NewNftablesManagerWithRunner(runner)
policies := []Policy{
{UUID: "first", Family: FamilyIPv4, HostIP: "0.0.0.0", HostPort: 8080, Protocol: "tcp", Mode: ModeAll},
{UUID: "second", Family: FamilyIPv4, HostIP: "192.0.2.10", HostPort: 5432, Protocol: "tcp", Mode: ModeAllow, Sources: []string{"203.0.113.1/32"}},
}
if err := manager.Reconcile(policies); err != nil {
t.Fatal(err)
}
if len(runner.inputCalls) != 1 {
t.Fatalf("restore calls = %d, want 1", len(runner.inputCalls))
}
call := runner.inputCalls[0]
if call.executable != "nft" || !reflect.DeepEqual(call.args, []string{"-f", "-"}) {
t.Fatalf("restore call = %#v", call)
}
for _, want := range []string{
"flush chain ip " + NftTable + " " + NftChain,
"ct state { established,related } return",
"ct original proto-dst 8080 comment \"1panel-docker:first\" drop",
"ct original ip daddr 192.0.2.10 ct original proto-dst 5432 ip saddr 203.0.113.1/32",
"add rule ip " + NftTable + " " + NftChain + " return",
} {
if !strings.Contains(call.input, want) {
t.Fatalf("restore input does not contain %q:\n%s", want, call.input)
}
}
}
func TestNftFamilyStatusExplainsIncompleteStep(t *testing.T) {
runner := newNftRecordingRunner()
runner.addTable("ip", dockerNftTable)
runner.addTable("ip", NftTable)
runner.addChain("ip", NftTable, NftBaseChain)
runner.addChain("ip", NftTable, NftChain)
runner.baseRules["ip"] = "counter packets 0 bytes 0 # handle 1\njump " + NftChain + " # handle 2\n"
status := NewNftablesManagerWithRunner(runner).Status(FamilyIPv4)
if status.State != StatusNotEffective || status.Reason != ReasonJumpNotFirst || !status.Initialized {
t.Fatalf("status = %#v", status)
}
runner.baseRules["ip"] = "jump " + NftChain + " # handle 2\n"
status = NewNftablesManagerWithRunner(runner).Status(FamilyIPv4)
if status.State != StatusEffective || !status.Initialized || !status.Bound || !status.Effective {
t.Fatalf("status = %#v", status)
}
}
func TestBuildNftScriptRejectsUnsafeTokens(t *testing.T) {
if _, err := buildNftScript([][]string{{"add", "rule", "unsafe value"}}); err == nil {
t.Fatal("expected unsafe token to be rejected")
}
}
+131
View File
@@ -0,0 +1,131 @@
package docker_guard
import (
"encoding/json"
"errors"
"fmt"
"net/netip"
"sort"
"strconv"
"strings"
)
var ErrInvalidPolicy = errors.New("invalid Docker port guard request")
func NormalizePolicy(policy Policy) (Policy, error) {
policy.Family = strings.ToLower(strings.TrimSpace(policy.Family))
policy.HostIP = strings.TrimSpace(policy.HostIP)
policy.Protocol = strings.ToLower(strings.TrimSpace(policy.Protocol))
policy.Mode = strings.ToLower(strings.TrimSpace(policy.Mode))
if policy.HostPort == 0 ||
(policy.Protocol != "tcp" && policy.Protocol != "udp") ||
(policy.Family != FamilyIPv4 && policy.Family != FamilyIPv6) ||
(policy.Mode != ModeAll && policy.Mode != ModeSources && policy.Mode != ModeAllow) {
return Policy{}, fmt.Errorf("%w: invalid policy fields", ErrInvalidPolicy)
}
address, err := netip.ParseAddr(policy.HostIP)
if err != nil || (policy.Family == FamilyIPv4) != address.Is4() {
return Policy{}, fmt.Errorf("%w: host IP does not match address family", ErrInvalidPolicy)
}
normalizedSources := make([]string, 0, len(policy.Sources))
seen := make(map[string]struct{}, len(policy.Sources))
for _, source := range policy.Sources {
source = strings.TrimSpace(source)
if source == "" {
continue
}
prefix, err := netip.ParsePrefix(source)
if err != nil {
if sourceAddress, addressErr := netip.ParseAddr(source); addressErr == nil {
bits := 128
if sourceAddress.Is4() {
bits = 32
}
prefix = netip.PrefixFrom(sourceAddress, bits)
} else {
return Policy{}, fmt.Errorf("%w: invalid source address %q", ErrInvalidPolicy, source)
}
}
if (policy.Family == FamilyIPv4) != prefix.Addr().Is4() {
return Policy{}, fmt.Errorf("%w: source %q does not match address family", ErrInvalidPolicy, source)
}
canonical := prefix.Masked().String()
if _, exists := seen[canonical]; !exists {
seen[canonical] = struct{}{}
normalizedSources = append(normalizedSources, canonical)
}
}
if policy.Mode == ModeSources && len(normalizedSources) == 0 {
return Policy{}, fmt.Errorf("%w: deny_sources requires at least one source", ErrInvalidPolicy)
}
if policy.Mode == ModeAll {
normalizedSources = []string{}
}
sort.Strings(normalizedSources)
policy.Sources = normalizedSources
return policy, nil
}
func NormalizePolicyUUIDs(values []string) ([]string, error) {
uuids := make([]string, 0, len(values))
seen := make(map[string]struct{}, len(values))
for _, policyUUID := range values {
policyUUID = strings.TrimSpace(policyUUID)
if policyUUID == "" {
return nil, fmt.Errorf("%w: policy UUID cannot be empty", ErrInvalidPolicy)
}
if _, exists := seen[policyUUID]; exists {
continue
}
seen[policyUUID] = struct{}{}
uuids = append(uuids, policyUUID)
}
if len(uuids) == 0 {
return nil, fmt.Errorf("%w: policy UUIDs cannot be empty", ErrInvalidPolicy)
}
return uuids, nil
}
func PolicySyncKey(policy Policy) string {
mode := policy.Mode
if mode == ModeAllow && len(policy.Sources) == 0 {
mode = ModeAll
}
sources := append([]string(nil), policy.Sources...)
sort.Strings(sources)
return strings.Join([]string{
policy.UUID, policy.Family, CanonicalHost(policy.HostIP), strconv.Itoa(int(policy.HostPort)),
policy.Protocol, mode, strings.Join(sources, ","),
}, "\x00")
}
func PolicyStatesEqual(left, right []Policy) bool {
if len(left) != len(right) {
return false
}
counts := make(map[string]int, len(left))
for _, policy := range left {
counts[PolicySyncKey(policy)]++
}
for _, policy := range right {
key := PolicySyncKey(policy)
if counts[key] == 0 {
return false
}
counts[key]--
}
return true
}
func CanonicalHost(value string) string {
if address, err := netip.ParseAddr(value); err == nil {
return address.String()
}
return value
}
func DecodeSources(value string) []string {
result := []string{}
_ = json.Unmarshal([]byte(value), &result)
return result
}
@@ -0,0 +1,83 @@
package docker_guard
import (
"fmt"
"github.com/1Panel-dev/1Panel/agent/constant"
)
type Runtime interface {
Initialize([]Policy) error
Bind() error
Reconcile([]Policy) error
Unbind() error
Cleanup() error
Initialized(string) (bool, error)
Status(string) FamilyStatus
ListPolicies() ([]Policy, error)
}
func NewRuntime(provider string) Runtime {
if provider == constant.FirewallProviderNftables {
return NewNftablesManager()
}
return NewManager()
}
func Verify(runtime Runtime, desired []Policy) error {
actual, err := runtime.ListPolicies()
if err != nil {
return fmt.Errorf("verify synchronized Docker firewall policies: %w", err)
}
if !PolicyStatesEqual(actual, desired) {
return fmt.Errorf("verify synchronized Docker firewall policies: target policies do not match the database")
}
return nil
}
func ReconcileTarget(backend string, policies []Policy, runtime Runtime) error {
families := make(map[string]struct{}, len(policies))
needsInitialize, needsBind := false, false
for _, policy := range policies {
families[policy.Family] = struct{}{}
}
if len(families) == 0 {
initialized := false
for _, family := range []string{FamilyIPv4, FamilyIPv6} {
status := runtime.Status(family)
if status.Reason == ReasonInspectFailed {
return fmt.Errorf("inspect Docker firewall target %s for %s failed", backend, family)
}
initialized = initialized || status.Initialized
}
if initialized {
return runtime.Reconcile(nil)
}
return nil
}
for family := range families {
status := runtime.Status(family)
needsInitialize = needsInitialize || !status.Initialized
needsBind = needsBind || !status.Bound || !status.Effective
}
var err error
if needsInitialize {
err = runtime.Initialize(policies)
} else {
if needsBind {
err = runtime.Bind()
}
if err == nil {
err = runtime.Reconcile(policies)
}
}
if err != nil {
return err
}
for family := range families {
if !runtime.Status(family).Effective {
return fmt.Errorf("Docker firewall target %s is not effective for %s", backend, family)
}
}
return nil
}
+170
View File
@@ -0,0 +1,170 @@
package filter
import (
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"fmt"
"strings"
)
type CheckFlagCodec struct {
secret []byte
version int
}
type checkFlagClaims struct {
Version int `json:"version"`
Provider Provider `json:"provider"`
ScopeKey string `json:"scopeKey"`
RuleDigest string `json:"ruleDigest"`
SnapshotRevision string `json:"snapshotRevision"`
ManagedRevision string `json:"managedRevision"`
Decision CheckDecision `json:"decision"`
Classification CheckClassification `json:"classification"`
AllowedActions []CheckAction `json:"allowedActions"`
AdoptionCandidates []checkFlagAdoptionCandidate `json:"adoptionCandidates,omitempty"`
}
type checkFlagAdoptionCandidate struct {
InstanceKey string `json:"instanceKey"`
Locator Locator `json:"locator"`
}
type CreateAuthorization struct {
Operation ChangeOperation
Locator *Locator
}
func NewCheckFlagCodec(secret []byte, version int) *CheckFlagCodec {
return &CheckFlagCodec{secret: append([]byte(nil), secret...), version: version}
}
func (c *CheckFlagCodec) Sign(result RuleCheckResult, snapshot Snapshot, managedRevision string) (string, error) {
ruleDigest, err := ruleDigest(result.RequestedRule)
if err != nil {
return "", err
}
claims := checkFlagClaims{
Version: c.version,
Provider: result.RequestedRule.Scope.Provider,
ScopeKey: result.RequestedRule.Scope.Key(),
RuleDigest: ruleDigest,
SnapshotRevision: snapshot.Revision,
ManagedRevision: managedRevision,
Decision: result.Decision,
Classification: result.Classification,
AllowedActions: result.AllowedActions,
}
if result.Classification == CheckClassificationExactExternal {
claims.AdoptionCandidates = make([]checkFlagAdoptionCandidate, 0, len(result.Candidates))
for _, candidate := range result.Candidates {
claims.AdoptionCandidates = append(claims.AdoptionCandidates, checkFlagAdoptionCandidate{
InstanceKey: candidate.InstanceKey,
Locator: candidate.Locator,
})
}
}
payload, err := json.Marshal(claims)
if err != nil {
return "", err
}
signature := c.signature(payload)
return base64.RawURLEncoding.EncodeToString(payload) + "." + base64.RawURLEncoding.EncodeToString(signature), nil
}
func (c *CheckFlagCodec) Authorize(
checkFlag string,
action CheckAction,
adoptInstanceKey string,
rule FirewallRule,
snapshot Snapshot,
managedRevision string,
) (CreateAuthorization, error) {
claims, err := c.parse(checkFlag)
if err != nil {
return CreateAuthorization{}, err
}
ruleDigest, err := ruleDigest(rule)
if err != nil {
return CreateAuthorization{}, err
}
if claims.Version != c.version ||
claims.Provider != rule.Scope.Provider ||
claims.ScopeKey != rule.Scope.Key() ||
claims.RuleDigest != ruleDigest ||
claims.SnapshotRevision != snapshot.Revision ||
claims.ManagedRevision != managedRevision {
return CreateAuthorization{}, fmt.Errorf("%w: firewall or managed rules changed", ErrRuleCheckRequired)
}
if claims.Decision != CheckDecisionReady && claims.Decision != CheckDecisionConfirmationRequired {
return CreateAuthorization{}, ErrRuleOperation
}
if !containsCheckAction(claims.AllowedActions, action) {
return CreateAuthorization{}, ErrRuleOperation
}
switch action {
case CheckActionCreate, CheckActionCreateAnyway:
if strings.TrimSpace(adoptInstanceKey) != "" {
return CreateAuthorization{}, ErrRuleOperation
}
return CreateAuthorization{Operation: ChangeCreate}, nil
case CheckActionAdopt, CheckActionSelectAdopt:
for _, candidate := range claims.AdoptionCandidates {
if candidate.InstanceKey == adoptInstanceKey && adoptInstanceKey != "" {
locator := candidate.Locator
return CreateAuthorization{Operation: ChangeAdopt, Locator: &locator}, nil
}
}
return CreateAuthorization{}, ErrRuleOperation
default:
return CreateAuthorization{}, ErrRuleOperation
}
}
func (c *CheckFlagCodec) parse(checkFlag string) (checkFlagClaims, error) {
parts := strings.Split(strings.TrimSpace(checkFlag), ".")
if len(parts) != 2 || parts[0] == "" || parts[1] == "" {
return checkFlagClaims{}, ErrRuleCheckRequired
}
payload, err := base64.RawURLEncoding.DecodeString(parts[0])
if err != nil {
return checkFlagClaims{}, ErrRuleCheckRequired
}
signature, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil || !hmac.Equal(signature, c.signature(payload)) {
return checkFlagClaims{}, ErrRuleCheckRequired
}
var claims checkFlagClaims
if err := json.Unmarshal(payload, &claims); err != nil {
return checkFlagClaims{}, ErrRuleCheckRequired
}
return claims, nil
}
func (c *CheckFlagCodec) signature(payload []byte) []byte {
mac := hmac.New(sha256.New, c.secret)
_, _ = mac.Write(payload)
return mac.Sum(nil)
}
func ruleDigest(rule FirewallRule) (string, error) {
payload, err := json.Marshal(rule)
if err != nil {
return "", err
}
sum := sha256.Sum256(payload)
return hex.EncodeToString(sum[:]), nil
}
func containsCheckAction(actions []CheckAction, expected CheckAction) bool {
for _, action := range actions {
if action == expected {
return true
}
}
return false
}
-366
View File
@@ -1,366 +0,0 @@
package filter
import "testing"
func TestCheckCreateRequestsAdoptionForEquivalentExternalRule(t *testing.T) {
rule := checkAddressRule("172.16.10.111", ActionDrop)
snapshot := checkSnapshot(t, rule)
plan, err := CheckCreate(snapshot, rule, nil, "")
if err != nil {
t.Fatalf("plan create: %v", err)
}
if plan.Decision != CheckDecisionConfirmationRequired || plan.Classification != CheckClassificationExactExternal || plan.Reason != "equivalent_external_rule" {
t.Fatalf("unexpected adoption check: %#v", plan)
}
if len(plan.Candidates) != 1 || len(plan.AllowedActions) != 2 || plan.AllowedActions[0] != CheckActionAdopt || plan.Candidates[0].InstanceKey == "" {
t.Fatalf("adoption details missing: %#v", plan)
}
}
func TestCheckCreateTreatsManagedDuplicateAsIdempotent(t *testing.T) {
rule := checkPortRule("22", ActionAccept)
marker := "onepanel:created:ssh"
observed := checkObservedRule(rule, marker, 1)
snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{observed})
if err != nil {
t.Fatalf("snapshot: %v", err)
}
ruleKey, _ := RuleKey(rule)
desired := DesiredRule{
UUID: "managed", Rule: rule, RuleKey: ruleKey, Origin: RuleOriginCreated, Marker: marker,
}
plan, err := CheckCreate(snapshot, rule, []DesiredRule{desired}, "")
if err != nil {
t.Fatalf("plan managed duplicate: %v", err)
}
if plan.Decision != CheckDecisionNoChange || plan.Classification != CheckClassificationExactManaged || plan.ExistingRuleUUID != desired.UUID {
t.Fatalf("managed duplicate was not idempotent: %#v", plan)
}
}
func TestCheckCreateBlocksMissingManagedRuleInsteadOfCreatingDuplicate(t *testing.T) {
rule := checkPortRule("22", ActionAccept)
snapshot, err := NewSnapshot(rule.Scope, nil)
if err != nil {
t.Fatalf("snapshot: %v", err)
}
ruleKey, _ := RuleKey(rule)
desired := DesiredRule{UUID: "managed", Rule: rule, RuleKey: ruleKey, Origin: RuleOriginCreated}
plan, err := CheckCreate(snapshot, rule, []DesiredRule{desired}, "")
if err != nil {
t.Fatalf("plan missing managed rule: %v", err)
}
if plan.Decision != CheckDecisionBlocked || plan.Reason != "managed_rule_drifted" {
t.Fatalf("missing managed rule was recreated: %#v", plan)
}
}
func TestCheckCreateRequiresCandidateSelectionForDuplicates(t *testing.T) {
rule := checkPortRule("3306", ActionAccept)
first := checkObservedRule(rule, "", 1)
second := checkObservedRule(rule, "", 2)
snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{first, second})
if err != nil {
t.Fatalf("snapshot: %v", err)
}
plan, err := CheckCreate(snapshot, rule, nil, "")
if err != nil {
t.Fatalf("plan duplicate candidates: %v", err)
}
if plan.Decision != CheckDecisionConfirmationRequired || plan.Reason != "multiple_equivalent_external_rules" || len(plan.Candidates) != 2 || plan.AllowedActions[0] != CheckActionSelectAdopt {
t.Fatalf("check guessed an equivalent candidate: %#v", plan)
}
if plan.Candidates[0].InstanceKey == "" || plan.Candidates[1].InstanceKey == "" || plan.Candidates[0].InstanceKey == plan.Candidates[1].InstanceKey {
t.Fatalf("candidate instance keys must be present and unique: %+v", plan.Candidates)
}
if _, err := FindCandidate(plan.Candidates, ""); err == nil {
t.Fatalf("expected candidate selection to be required, got %v", err)
}
secondKey, _ := InstanceKey(second)
candidate, err := FindCandidate(plan.Candidates, secondKey)
if err != nil || candidate.Locator.Position == nil || *candidate.Locator.Position != 2 {
t.Fatalf("selected candidate was not found: candidate=%#v err=%v", candidate, err)
}
}
func TestCheckCreateBlocksAdoptionOfProtectedRule(t *testing.T) {
rule := checkPortRule("22", ActionAccept)
observed := checkObservedRule(rule, "", 1)
observed.Protected = true
snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{observed})
if err != nil {
t.Fatalf("snapshot: %v", err)
}
plan, err := CheckCreate(snapshot, rule, nil, "")
if err != nil {
t.Fatalf("plan protected rule: %v", err)
}
if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationProtected || plan.Reason != "protected_rule" || len(plan.AllowedActions) != 0 {
t.Fatalf("protected rule could be adopted: %#v", plan)
}
}
func TestCheckCreateBlocksRuntimePermanentMismatch(t *testing.T) {
rule := checkPortRule("443", ActionAccept)
observed := checkObservedRule(rule, "", 1)
observed.Persistence = PersistenceStatusRuntimeOnly
snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{observed})
if err != nil {
t.Fatalf("snapshot: %v", err)
}
plan, err := CheckCreate(snapshot, rule, nil, "")
if err != nil {
t.Fatalf("plan persistence drift: %v", err)
}
if plan.Decision != CheckDecisionBlocked || plan.Reason != "runtime_permanent_mismatch" {
t.Fatalf("runtime-only rule could be adopted: %#v", plan)
}
}
func TestCheckCreateDoesNotLetOpaqueFirewalldServiceBlockUnrelatedRule(t *testing.T) {
rule := FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput},
NativeKind: NativeKindZonePort, Protocol: "tcp", DestinationPort: "8080", Action: ActionAccept,
}
service := ObservedRule{
Rule: FirewallRule{Scope: rule.Scope, NativeKind: NativeKindZoneService},
Locator: Locator{Provider: ProviderFirewalld, ScopeKey: rule.Scope.Key(), Canonical: "service:ssh"},
ParseStatus: ParseStatusOpaque, Raw: "ssh", Persistence: PersistenceStatusConverged,
}
snapshot, err := NewSnapshot(rule.Scope, []ObservedRule{service})
if err != nil {
t.Fatalf("snapshot: %v", err)
}
plan, err := CheckCreate(snapshot, rule, nil, "")
if err != nil {
t.Fatalf("plan with opaque service: %v", err)
}
if plan.Decision != CheckDecisionReady || plan.Classification != CheckClassificationNone {
t.Fatalf("opaque service blocked unrelated native port: %#v", plan)
}
opaqueRich := service
opaqueRich.Rule.NativeKind = NativeKindRichRule
opaqueRich.Locator.Canonical = `rich:rule log prefix="audit" accept`
opaqueRich.Raw = `rule log prefix="audit" accept`
snapshot, err = NewSnapshot(rule.Scope, []ObservedRule{opaqueRich})
if err != nil {
t.Fatalf("opaque rich snapshot: %v", err)
}
plan, err = CheckCreate(snapshot, rule, nil, "")
if err != nil {
t.Fatalf("plan with opaque rich rule: %v", err)
}
if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationUnsupported {
t.Fatalf("opaque rich rule did not block an unsafe plan: %#v", plan)
}
}
func TestCheckCreateClassifiesCoverageAndConflict(t *testing.T) {
requested := checkAddressRule("172.16.10.111", ActionDrop)
covering := checkAddressRule("172.16.10.0/24", ActionDrop)
coveredSnapshot := checkSnapshot(t, covering)
plan, err := CheckCreate(coveredSnapshot, requested, nil, "")
if err != nil {
t.Fatalf("plan covered rule: %v", err)
}
if plan.Classification != CheckClassificationCovered || plan.Decision != CheckDecisionConfirmationRequired || plan.AllowedActions[0] != CheckActionCreateAnyway {
t.Fatalf("unexpected covered plan: %#v", plan)
}
conflicting := checkAddressRule("172.16.10.0/24", ActionAccept)
conflictSnapshot := checkSnapshot(t, conflicting)
plan, err = CheckCreate(conflictSnapshot, requested, nil, "")
if err != nil {
t.Fatalf("plan conflicting rule: %v", err)
}
if plan.Classification != CheckClassificationConflict || plan.Decision != CheckDecisionBlocked {
t.Fatalf("unexpected conflict plan: %#v", plan)
}
}
func TestCheckCreateAllowsPartialOverlapWithOppositeActionAfterConfirmation(t *testing.T) {
scope := Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: FirewalldInputZone, Direction: DirectionInput}
existingRule := FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: FirewalldInputZone, Direction: DirectionInput},
NativeKind: NativeKindRichRule, Protocol: "all", SourceAddress: "1.1.1.1", Action: ActionDrop,
}
requested := FirewallRule{
Scope: scope, NativeKind: NativeKindZonePort, Protocol: "tcp", DestinationPort: "8080", Action: ActionAccept,
}
snapshot := checkSnapshot(t, existingRule)
plan, err := CheckCreate(snapshot, requested, nil, "")
if err != nil {
t.Fatalf("plan partially overlapping rule: %v", err)
}
if plan.Decision != CheckDecisionConfirmationRequired || plan.Classification != CheckClassificationConflict ||
plan.Reason != "partially_overlapping_rule_with_different_action" ||
len(plan.AllowedActions) == 0 || plan.AllowedActions[0] != CheckActionCreateAnyway {
t.Fatalf("partial overlap was not confirmable: %#v", plan)
}
}
func TestRuleCoverageAndOverlapRespectAddressFamilies(t *testing.T) {
base := Scope{Provider: ProviderFirewalld, Zone: FirewalldInputZone, Direction: DirectionInput}
ipv4 := FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: base.Zone, Direction: base.Direction},
NativeKind: NativeKindRichRule, Protocol: "all", SourceAddress: "1.1.1.1", Action: ActionDrop,
}
ipv6 := FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv6, Zone: base.Zone, Direction: base.Direction},
NativeKind: NativeKindRichRule, Protocol: "tcp", DestinationPort: "8080", Action: ActionAccept,
}
if RulesOverlap(ipv4, ipv6) || RuleCovers(ipv4, ipv6) {
t.Fatal("disjoint IPv4 and IPv6 rules must not overlap")
}
inet := ipv6
inet.Scope.Family = FamilyInet
inet.SourceAddress = ""
if !RulesOverlap(ipv4, inet) {
t.Fatal("inet rule should overlap an IPv4 rule")
}
if RuleCovers(ipv4, inet) {
t.Fatal("IPv4 rule must not cover a dual-stack inet rule")
}
}
func TestCheckCreateAllowsOrderedFirewalldDenyBeforeNativePort(t *testing.T) {
scope := Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: "public", Direction: DirectionInput}
existingRule := FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput},
NativeKind: NativeKindZonePort, Protocol: "tcp", DestinationPort: "3306", Action: ActionAccept,
}
priority := -100
requested := FirewallRule{
Scope: scope, NativeKind: NativeKindRichRule, Protocol: "tcp", SourceAddress: "172.16.10.111",
DestinationPort: "3306", Action: ActionDrop, Priority: &priority,
}
existing := checkObservedRule(existingRule, "", 1)
existing.Persistence = PersistenceStatusConverged
snapshot, err := NewSnapshot(scope, []ObservedRule{existing})
if err != nil {
t.Fatalf("snapshot: %v", err)
}
plan, err := CheckCreate(snapshot, requested, nil, "")
if err != nil {
t.Fatalf("plan ordered deny: %v", err)
}
if plan.Decision != CheckDecisionReady || plan.Classification != CheckClassificationNone {
t.Fatalf("negative-priority deny was blocked: %#v", plan)
}
existing.Protected = true
protectedSnapshot, _ := NewSnapshot(scope, []ObservedRule{existing})
plan, err = CheckCreate(protectedSnapshot, requested, nil, "")
if err != nil {
t.Fatalf("plan protected overlap: %v", err)
}
if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationConflict {
t.Fatalf("protected port overlap was allowed: %#v", plan)
}
}
func TestRuleCoverageAndOverlapUseNormalizedRanges(t *testing.T) {
existing := checkPortRule("8000-9000", ActionAccept)
requested := checkPortRule("8080", ActionAccept)
if !RuleCovers(existing, requested) || !RulesOverlap(existing, requested) {
t.Fatal("expected port range to cover and overlap contained port")
}
differentProtocol := requested
differentProtocol.Protocol = "udp"
if RulesOverlap(existing, differentProtocol) {
t.Fatal("different transport protocols should not overlap")
}
}
func TestCheckCreateBlocksRuleTargetingCurrentManagementClient(t *testing.T) {
rule := checkAddressRule("203.0.113.9", ActionDrop)
snapshot, err := NewSnapshot(rule.Scope, nil)
if err != nil {
t.Fatalf("snapshot: %v", err)
}
plan, err := CheckCreate(snapshot, rule, nil, "203.0.113.9")
if err != nil {
t.Fatalf("plan create: %v", err)
}
if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationProtected || plan.Reason != "current_management_connection" {
t.Fatalf("unexpected current connection decision: %#v", plan)
}
}
func TestCheckCreateBlocksProtectedManagementPortWithoutObservedAllowRule(t *testing.T) {
rule := checkPortRule("22", ActionDrop)
snapshot, err := NewSnapshot(rule.Scope, nil)
if err != nil {
t.Fatalf("snapshot: %v", err)
}
plan, err := CheckCreate(snapshot, rule, nil, "203.0.113.9", PortWhitelist{
Family: "ipv4", Port: "22", Protocol: "tcp",
})
if err != nil {
t.Fatalf("plan create: %v", err)
}
if plan.Decision != CheckDecisionBlocked || plan.Classification != CheckClassificationProtected || plan.Reason != "current_management_connection" {
t.Fatalf("protected management port was not blocked: %#v", plan)
}
}
func TestCheckCreateAllowsProtectedPortDenyForUnrelatedSource(t *testing.T) {
rule := checkPortRule("22", ActionDrop)
rule.SourceAddress = "198.51.100.8/32"
snapshot, err := NewSnapshot(rule.Scope, nil)
if err != nil {
t.Fatalf("snapshot: %v", err)
}
plan, err := CheckCreate(snapshot, rule, nil, "203.0.113.9", PortWhitelist{
Family: "ipv4", Port: "22", Protocol: "tcp",
})
if err != nil {
t.Fatalf("plan create: %v", err)
}
if plan.Decision != CheckDecisionReady {
t.Fatalf("unrelated source was incorrectly blocked: %#v", plan)
}
}
func checkSnapshot(t *testing.T, rules ...FirewallRule) Snapshot {
t.Helper()
observed := make([]ObservedRule, 0, len(rules))
for index, rule := range rules {
observed = append(observed, checkObservedRule(rule, "", index+1))
}
snapshot, err := NewSnapshot(rules[0].Scope, observed)
if err != nil {
t.Fatalf("snapshot: %v", err)
}
return snapshot
}
func checkObservedRule(rule FirewallRule, marker string, position int) ObservedRule {
return ObservedRule{
Rule: rule,
Locator: Locator{Provider: rule.Scope.Provider, ScopeKey: rule.Scope.Key(), Position: &position},
Marker: marker, ParseStatus: ParseStatusSupported,
}
}
func checkAddressRule(address string, action Action) FirewallRule {
return FirewallRule{
Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput},
NativeKind: NativeKindRule,
Protocol: "all",
SourceAddress: address,
Action: action,
}
}
func checkPortRule(port string, action Action) FirewallRule {
return FirewallRule{
Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput},
NativeKind: NativeKindRule,
Protocol: "tcp",
DestinationPort: port,
Action: action,
}
}
@@ -1,195 +0,0 @@
package filter
import (
"errors"
"testing"
)
func TestRuleKeyUsesSemanticFields(t *testing.T) {
priority := 10
base := FirewallRule{
UUID: "first",
Scope: Scope{
Provider: ProviderFirewalld,
Family: FamilyIPv4,
Zone: "public",
Direction: DirectionInput,
},
NativeKind: NativeKindRichRule,
Protocol: "TCP",
SourceAddress: "172.16.10.111",
DestinationPort: "3306:3306",
Action: "ALLOW",
Priority: &priority,
Description: "first description",
}
other := base
other.UUID = "second"
other.Description = "description is not identity"
other.SourceAddress = "172.16.10.111/32"
other.DestinationPort = "3306"
firstKey, err := RuleKey(base)
if err != nil {
t.Fatalf("first rule key: %v", err)
}
secondKey, err := RuleKey(other)
if err != nil {
t.Fatalf("second rule key: %v", err)
}
if firstKey != secondKey {
t.Fatalf("equivalent rules produced different keys:\n%s\n%s", firstKey, secondKey)
}
changedPriority := 11
other.Priority = &changedPriority
changedKey, err := RuleKey(other)
if err != nil {
t.Fatalf("changed rule key: %v", err)
}
if changedKey == firstKey {
t.Fatal("priority change did not change rule key")
}
}
func TestRuleKeyKeepsFamilyInsideSharedFirewalldScope(t *testing.T) {
ipv4 := FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: "public", Direction: DirectionInput},
NativeKind: NativeKindRichRule, Protocol: "tcp", Action: ActionAccept,
}
ipv6 := ipv4
ipv6.Scope.Family = FamilyIPv6
first, err := RuleKey(ipv4)
if err != nil {
t.Fatalf("IPv4 key: %v", err)
}
second, err := RuleKey(ipv6)
if err != nil {
t.Fatalf("IPv6 key: %v", err)
}
if first == second {
t.Fatal("firewalld family-specific rich rules collapsed to one identity")
}
}
func TestInstanceKeyIncludesLocator(t *testing.T) {
positionOne := 1
positionTwo := 2
rule := ObservedRule{
Rule: FirewallRule{
Scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput},
Protocol: "tcp",
DestinationPort: "22",
Action: ActionAccept,
},
Marker: "1panel-rule:test",
ParseStatus: ParseStatusSupported,
Locator: Locator{Position: &positionOne},
}
first, err := InstanceKey(rule)
if err != nil {
t.Fatalf("first instance key: %v", err)
}
rule.Locator.Position = &positionTwo
second, err := InstanceKey(rule)
if err != nil {
t.Fatalf("second instance key: %v", err)
}
if first == second {
t.Fatal("position change did not change instance key")
}
rule.Locator.Position = &positionOne
rule.Persistence = PersistenceStatusRuntimeOnly
runtimeOnly, err := InstanceKey(rule)
if err != nil {
t.Fatalf("runtime-only instance key: %v", err)
}
if runtimeOnly == first {
t.Fatal("runtime/permanent presence did not change instance key")
}
}
func TestSnapshotRevisionIsInputOrderIndependent(t *testing.T) {
scope := Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput}
positionOne := 1
positionTwo := 2
rules := []ObservedRule{
{
Rule: FirewallRule{Scope: scope, Protocol: "tcp", DestinationPort: "22", Action: ActionAccept},
Locator: Locator{Position: &positionOne},
ParseStatus: ParseStatusSupported,
},
{
Rule: FirewallRule{Scope: scope, Protocol: "tcp", DestinationPort: "80", Action: ActionAccept},
Locator: Locator{Position: &positionTwo},
ParseStatus: ParseStatusSupported,
},
}
first, err := SnapshotRevision(scope, rules)
if err != nil {
t.Fatalf("first snapshot revision: %v", err)
}
reversed := []ObservedRule{rules[1], rules[0]}
second, err := SnapshotRevision(scope, reversed)
if err != nil {
t.Fatalf("second snapshot revision: %v", err)
}
if first != second {
t.Fatalf("slice order changed snapshot revision:\n%s\n%s", first, second)
}
rules[0].Locator.Position = &positionTwo
changed, err := SnapshotRevision(scope, rules)
if err != nil {
t.Fatalf("changed snapshot revision: %v", err)
}
if changed == first {
t.Fatal("locator position change did not change snapshot revision")
}
}
func TestSnapshotRevisionSupportsOpaqueRules(t *testing.T) {
scope := Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput}
position := 1
rules := []ObservedRule{
{
Rule: FirewallRule{Scope: scope},
Locator: Locator{Position: &position, Canonical: "service:ssh"},
ParseStatus: ParseStatusOpaque,
Raw: "services: ssh",
},
}
revision, err := SnapshotRevision(scope, rules)
if err != nil {
t.Fatalf("opaque snapshot revision: %v", err)
}
if revision == "" {
t.Fatal("expected non-empty opaque snapshot revision")
}
instanceKey, err := InstanceKey(rules[0])
if err != nil || instanceKey == "" {
t.Fatalf("opaque instance key: key=%q err=%v", instanceKey, err)
}
rules[0].Persistence = PersistenceStatusPermanentOnly
changed, err := SnapshotRevision(scope, rules)
if err != nil {
t.Fatalf("opaque persistence revision: %v", err)
}
if changed == revision {
t.Fatal("opaque runtime/permanent presence did not change snapshot revision")
}
}
func TestSnapshotRevisionRequiresPositionForOrderedProvider(t *testing.T) {
scope := Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput}
_, err := SnapshotRevision(scope, []ObservedRule{
{
Rule: FirewallRule{Scope: scope, Protocol: "tcp", DestinationPort: "22", Action: ActionAccept},
ParseStatus: ParseStatusSupported,
},
})
if !errors.Is(err, ErrInvalidRule) {
t.Fatalf("expected missing position error, got %v", err)
}
}
@@ -1,256 +0,0 @@
package filter
import "testing"
func TestMergeInventoryPreservesObservedOrderAndOwnership(t *testing.T) {
ssh := inventoryTestRule("22")
http := inventoryTestRule("80")
https := inventoryTestRule("443")
sshObserved := inventoryObservedRule(t, ssh, "onepanel:created:ssh", 1)
httpObserved := inventoryObservedRule(t, http, "", 2)
protectedKey, _ := RuleKey(http)
items, err := MergeInventory(InventoryMergeInput{
Observed: []ObservedRule{httpObserved, sshObserved},
Desired: []DesiredRule{
{UUID: "ssh", Rule: ssh, Origin: RuleOriginCreated, Marker: "onepanel:created:ssh"},
{UUID: "https", Rule: https, Origin: RuleOriginAdopted},
},
ProtectedObservedKeys: map[string]struct{}{protectedKey: {}},
})
if err != nil {
t.Fatalf("merge inventory: %v", err)
}
if len(items) != 3 {
t.Fatalf("unexpected item count: %#v", items)
}
if items[0].Rule.DestinationPort != "80" || items[0].State != InventoryStateProtected || items[0].Desired != nil {
t.Fatalf("observed order/protected classification changed: %#v", items[0])
}
if items[1].Rule.DestinationPort != "22" || items[1].State != InventoryStateManaged || items[1].Match != InventoryMatchExact {
t.Fatalf("managed rule was not matched: %#v", items[1])
}
if items[2].Rule.DestinationPort != "443" || items[2].State != InventoryStateDrifted || items[2].Match != InventoryMatchMissing || items[2].Observed != nil {
t.Fatalf("missing desired rule was not appended as drifted: %#v", items[2])
}
}
func TestMergeInventoryUsesDesiredDescriptionForManagedRule(t *testing.T) {
rule := inventoryTestRule("8080")
desired := rule
desired.Description = "user description"
observed := inventoryObservedRule(t, rule, "1panel-rule:web", 1)
items, err := MergeInventory(InventoryMergeInput{
Observed: []ObservedRule{observed},
Desired: []DesiredRule{{
UUID: "web", Rule: desired, Origin: RuleOriginCreated, Marker: observed.Marker,
}},
})
if err != nil {
t.Fatalf("merge inventory: %v", err)
}
if len(items) != 1 || items[0].Rule.Description != desired.Description {
t.Fatalf("managed description was not taken from desired state: %#v", items)
}
if items[0].Observed == nil || items[0].Observed.Rule.Description != "" {
t.Fatalf("observed runtime rule was modified: %#v", items[0].Observed)
}
}
func TestMergeInventoryUsesInstanceBeforeSemanticIdentity(t *testing.T) {
rule := inventoryTestRule("22")
first := inventoryObservedRule(t, rule, "", 1)
second := inventoryObservedRule(t, rule, "", 2)
secondKey, err := InstanceKey(second)
if err != nil {
t.Fatalf("instance key: %v", err)
}
items, err := MergeInventory(InventoryMergeInput{
Observed: []ObservedRule{first, second},
Desired: []DesiredRule{{
UUID: "adopted", Rule: rule, Origin: RuleOriginAdopted, ObservedInstanceKey: secondKey,
}},
})
if err != nil {
t.Fatalf("merge inventory: %v", err)
}
if items[0].State != InventoryStateExternal || items[1].State != InventoryStateAdopted || items[1].Observed.Locator.Position == nil || *items[1].Observed.Locator.Position != 2 {
t.Fatalf("instance locator did not select the intended candidate: %#v", items)
}
}
func TestMergeInventoryDoesNotGuessBetweenEquivalentRules(t *testing.T) {
rule := inventoryTestRule("3306")
items, err := MergeInventory(InventoryMergeInput{
Observed: []ObservedRule{
inventoryObservedRule(t, rule, "", 1),
inventoryObservedRule(t, rule, "", 2),
},
Desired: []DesiredRule{{UUID: "managed", Rule: rule, Origin: RuleOriginCreated}},
})
if err != nil {
t.Fatalf("merge inventory: %v", err)
}
if len(items) != 3 || items[0].State != InventoryStateExternal || items[1].State != InventoryStateExternal || items[2].State != InventoryStateDrifted || items[2].Match != InventoryMatchAmbiguous {
t.Fatalf("ambiguous candidates were guessed: %#v", items)
}
}
func TestMergeInventoryUsesMarkerToReportChangedManagedRule(t *testing.T) {
desiredRule := inventoryTestRule("80")
changedRule := inventoryTestRule("8080")
marker := "onepanel:created:web"
items, err := MergeInventory(InventoryMergeInput{
Observed: []ObservedRule{inventoryObservedRule(t, changedRule, marker, 1)},
Desired: []DesiredRule{{
UUID: "web", Rule: desiredRule, Origin: RuleOriginCreated, Marker: marker,
}},
})
if err != nil {
t.Fatalf("merge inventory: %v", err)
}
if len(items) != 1 || items[0].State != InventoryStateDrifted || items[0].Match != InventoryMatchChanged || items[0].Observed == nil || items[0].Desired == nil {
t.Fatalf("marker-owned semantic drift was not retained: %#v", items)
}
}
func TestMergeInventoryMarkerSurvivesTransientLocatorChange(t *testing.T) {
rule := inventoryTestRule("443")
marker := "1panel-rule:https"
previous := inventoryObservedRule(t, rule, marker, 2)
previousKey, err := InstanceKey(previous)
if err != nil {
t.Fatalf("previous instance key: %v", err)
}
current := inventoryObservedRule(t, rule, marker, 7)
items, err := MergeInventory(InventoryMergeInput{
Observed: []ObservedRule{current},
Desired: []DesiredRule{{
UUID: "https", Rule: rule, Origin: RuleOriginCreated,
Marker: marker, ObservedInstanceKey: previousKey,
}},
})
if err != nil {
t.Fatalf("merge locator drift: %v", err)
}
if len(items) != 1 || items[0].State != InventoryStateManaged || items[0].Match != InventoryMatchExact ||
items[0].Observed == nil || items[0].Observed.Locator.Position == nil || *items[0].Observed.Locator.Position != 7 {
t.Fatalf("stable marker did not survive transient locator change: %#v", items)
}
}
func TestMergeInventoryKeepsOpaqueRulesExternal(t *testing.T) {
opaque := ObservedRule{
Rule: FirewallRule{Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput}},
Locator: Locator{Provider: ProviderFirewalld, ScopeKey: "firewalld:public:input", NativeID: "service:ssh"},
ParseStatus: ParseStatusOpaque,
Raw: "ssh service",
}
items, err := MergeInventory(InventoryMergeInput{Observed: []ObservedRule{opaque}})
if err != nil {
t.Fatalf("merge opaque inventory: %v", err)
}
if len(items) != 1 || items[0].State != InventoryStateExternal || items[0].Match != InventoryMatchOpaque || items[0].Observed.Raw != opaque.Raw {
t.Fatalf("opaque rule was not preserved: %#v", items)
}
}
func TestMergeInventoryUsesAdapterProtectedClassification(t *testing.T) {
rule := inventoryTestRule("22")
observed := inventoryObservedRule(t, rule, "", 1)
observed.Protected = true
items, err := MergeInventory(InventoryMergeInput{Observed: []ObservedRule{observed}})
if err != nil {
t.Fatalf("merge protected inventory: %v", err)
}
if len(items) != 1 || items[0].State != InventoryStateProtected || items[0].Observed == nil || !items[0].Observed.Protected {
t.Fatalf("adapter protected classification was lost: %#v", items)
}
}
func TestMergeInventoryReportsRuntimePermanentDrift(t *testing.T) {
rule := inventoryTestRule("443")
observed := inventoryObservedRule(t, rule, "onepanel:created:https", 1)
observed.Persistence = PersistenceStatusPermanentOnly
items, err := MergeInventory(InventoryMergeInput{
Observed: []ObservedRule{observed},
Desired: []DesiredRule{{
UUID: "https", Rule: rule, Origin: RuleOriginCreated, Marker: observed.Marker,
}},
})
if err != nil {
t.Fatalf("merge persistence drift: %v", err)
}
if len(items) != 1 || items[0].State != InventoryStateDrifted || items[0].Match != InventoryMatchExact {
t.Fatalf("runtime/permanent drift was hidden: %#v", items)
}
}
func TestAttachRuntimeUsageDoesNotChangeOwnership(t *testing.T) {
rule := inventoryTestRule("8080")
items := []InventoryItem{{Rule: rule, State: InventoryStateAdopted, Match: InventoryMatchExact}}
usage := map[string]RuntimeUsage{
RuntimeUsageKey(rule): {UsedBy: []string{"nginx", "", "nginx", "1panel"}, Reason: "listener"},
}
result := AttachRuntimeUsage(items, usage)
if result[0].State != InventoryStateAdopted || result[0].Usage == nil || !result[0].Usage.Used {
t.Fatalf("usage changed ownership or was not attached: %#v", result[0])
}
if len(result[0].Usage.UsedBy) != 2 || result[0].Usage.UsedBy[0] != "1panel" || result[0].Usage.UsedBy[1] != "nginx" {
t.Fatalf("usage owners were not normalized: %#v", result[0].Usage)
}
if items[0].Usage != nil {
t.Fatal("AttachRuntimeUsage mutated its input")
}
}
func TestAttachRuntimeUsageAggregatesPortRanges(t *testing.T) {
rule := inventoryTestRule("8000-8010")
items := []InventoryItem{{Rule: rule}}
usage := map[string]RuntimeUsage{
"tcp\x008001": {Used: true, UsedBy: []string{"app"}, Reason: "application"},
"tcp\x008009": {Used: true, UsedBy: []string{"worker"}, Reason: "listener"},
"udp\x008005": {Used: true, UsedBy: []string{"ignored"}},
}
result := AttachRuntimeUsage(items, usage)
if result[0].Usage == nil || len(result[0].Usage.UsedBy) != 2 || result[0].Usage.Reason != "multiple" {
t.Fatalf("range usage was not aggregated: %#v", result[0].Usage)
}
}
func TestMergeInventoryRejectsStoredRuleKeyMismatch(t *testing.T) {
_, err := MergeInventory(InventoryMergeInput{Desired: []DesiredRule{{
UUID: "broken", Rule: inventoryTestRule("53"), RuleKey: "sha256:stale", Origin: RuleOriginCreated,
}}})
if err == nil {
t.Fatal("expected stored rule key mismatch to fail")
}
}
func inventoryTestRule(port string) FirewallRule {
return FirewallRule{
Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput},
NativeKind: NativeKindRule,
Protocol: "tcp",
DestinationPort: port,
Action: ActionAccept,
}
}
func inventoryObservedRule(t *testing.T, rule FirewallRule, marker string, position int) ObservedRule {
t.Helper()
return ObservedRule{
Rule: rule,
Locator: Locator{
Provider: rule.Scope.Provider,
ScopeKey: rule.Scope.Key(),
Position: &position,
},
Marker: marker,
ParseStatus: ParseStatusSupported,
}
}
+29
View File
@@ -125,6 +125,35 @@ type Scope struct {
Direction Direction `json:"direction"`
}
func ManagedInputScopes(provider Provider) []Scope {
base := Scope{Provider: provider, Direction: DirectionInput}
switch provider {
case ProviderIptables, ProviderNftables:
result := make([]Scope, 0, 6)
for _, family := range []Family{FamilyIPv4, FamilyIPv6} {
for _, chain := range []string{BasicBeforeChain, IptablesInputChain, BasicAfterChain} {
scope := base
scope.Family, scope.Table, scope.Chain = family, "filter", chain
result = append(result, scope)
}
}
return result
case ProviderFirewalld:
base.Family, base.Zone = FamilyInet, FirewalldInputZone
return []Scope{base}
case ProviderUFW:
result := make([]Scope, 0, 2)
for _, family := range []Family{FamilyIPv4, FamilyIPv6} {
scope := base
scope.Family, scope.Chain = family, UFWInputChain
result = append(result, scope)
}
return result
default:
return nil
}
}
func (s Scope) Normalize() Scope {
s.Provider = Provider(strings.ToLower(strings.TrimSpace(string(s.Provider))))
s.Family = Family(strings.ToLower(strings.TrimSpace(string(s.Family))))
-85
View File
@@ -1,85 +0,0 @@
package filter
import (
"errors"
"testing"
)
func TestScopeValidateMVP(t *testing.T) {
tests := []struct {
name string
scope Scope
key string
wantErr error
}{
{
name: "iptables basic chain",
scope: Scope{Provider: "IPTABLES", Family: FamilyIPv4, Table: "FILTER", Chain: "1panel_basic", Direction: DirectionInput},
key: "iptables:ipv4:filter:1PANEL_BASIC:input",
},
{
name: "firewalld public",
scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "PUBLIC", Direction: DirectionInput},
key: "firewalld:public:input",
},
{
name: "ufw incoming default chain",
scope: Scope{Provider: ProviderUFW, Family: FamilyIPv6, Direction: DirectionInput},
key: "ufw:incoming:ipv6",
},
{
name: "firewalld private unsupported",
scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "private", Direction: DirectionInput},
wantErr: ErrUnsupportedScope,
},
{
name: "iptables external chain unsupported",
scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "DOCKER", Direction: DirectionInput},
wantErr: ErrUnsupportedScope,
},
{
name: "ufw output unsupported",
scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: Direction("output")},
wantErr: ErrInvalidScope,
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
err := test.scope.ValidateMVP()
if test.wantErr != nil {
if !errors.Is(err, test.wantErr) {
t.Fatalf("expected %v, got %v", test.wantErr, err)
}
return
}
if err != nil {
t.Fatalf("validate scope: %v", err)
}
if got := test.scope.Key(); got != test.key {
t.Fatalf("expected key %q, got %q", test.key, got)
}
})
}
}
func TestCapabilitiesSupportsMVPScope(t *testing.T) {
capabilities := Capabilities{Scopes: MVPScopePatterns()}
if !capabilities.SupportsScope(Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput}) {
t.Fatal("expected UFW incoming IPv4 scope to be supported")
}
if capabilities.SupportsScope(Scope{Provider: ProviderUFW, Family: FamilyIPv6, Chain: "outgoing", Direction: Direction("output")}) {
t.Fatal("did not expect UFW outgoing scope to be supported")
}
if capabilities.SupportsScope(Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "private", Direction: DirectionInput}) {
t.Fatal("did not expect firewalld private zone to be supported")
}
}
func TestFirewalldFamiliesSharePublicExecutionScope(t *testing.T) {
ipv4 := Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: "public", Direction: DirectionInput}
ipv6 := Scope{Provider: ProviderFirewalld, Family: FamilyIPv6, Zone: "public", Direction: DirectionInput}
if ipv4.Key() != ipv6.Key() || ipv4.Key() != "firewalld:public:input" {
t.Fatalf("firewalld public pipeline was split by family: ipv4=%q ipv6=%q", ipv4.Key(), ipv6.Key())
}
}
@@ -1,270 +0,0 @@
package filter
import (
"errors"
"fmt"
"strings"
"testing"
)
func TestNormalizeRule(t *testing.T) {
rule, err := NormalizeRule(FirewallRule{
Scope: Scope{
Provider: ProviderUFW,
Family: FamilyIPv4,
Direction: DirectionInput,
},
Protocol: " TCP ",
SourceAddress: "172.16.10.111",
DestinationAddress: "0.0.0.0/0",
DestinationPort: "080:080",
Interface: " * ",
Action: "ALLOW",
ConnectionStates: []string{"NEW", "established", "new"},
})
if err != nil {
t.Fatalf("normalize rule: %v", err)
}
if rule.Scope.Key() != "ufw:incoming:ipv4" {
t.Fatalf("unexpected scope key %q", rule.Scope.Key())
}
if rule.NativeKind != NativeKindUFWRule {
t.Fatalf("unexpected native kind %q", rule.NativeKind)
}
if rule.Protocol != "tcp" || rule.Action != ActionAccept {
t.Fatalf("unexpected protocol/action: %s/%s", rule.Protocol, rule.Action)
}
if rule.SourceAddress != "172.16.10.111/32" || rule.DestinationAddress != "" {
t.Fatalf("unexpected normalized addresses: %q -> %q", rule.SourceAddress, rule.DestinationAddress)
}
if rule.DestinationPort != "80" {
t.Fatalf("unexpected destination port %q", rule.DestinationPort)
}
if rule.Interface != "" {
t.Fatalf("wildcard interface was not normalized: %q", rule.Interface)
}
if len(rule.ConnectionStates) != 2 || rule.ConnectionStates[0] != "established" || rule.ConnectionStates[1] != "new" {
t.Fatalf("unexpected states: %#v", rule.ConnectionStates)
}
}
func TestNormalizeRuleRejectsCompositeAndFamilyMismatch(t *testing.T) {
base := FirewallRule{
Scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput},
Protocol: "tcp/udp",
DestinationPort: "80",
Action: ActionAccept,
}
if _, err := NormalizeRule(base); !errors.Is(err, ErrCompositeRule) {
t.Fatalf("expected composite rule error, got %v", err)
}
base.Protocol = "tcp"
base.SourceAddress = "2001:db8::1"
if _, err := NormalizeRule(base); !errors.Is(err, ErrInvalidRule) {
t.Fatalf("expected family validation error, got %v", err)
}
}
func TestNormalizeRuleRejectsInvalidConnectionState(t *testing.T) {
rule := FirewallRule{
Scope: Scope{
Provider: ProviderNftables, Family: FamilyIPv4, Table: "filter",
Chain: IptablesInputChain, Direction: DirectionInput,
},
Protocol: "tcp", Action: ActionAccept,
ConnectionStates: []string{"established } accept\nflush ruleset"},
}
if _, err := NormalizeRule(rule); !errors.Is(err, ErrInvalidRule) {
t.Fatalf("expected invalid connection state error, got %v", err)
}
}
func TestNormalizeNativeDestinationPortSets(t *testing.T) {
for _, provider := range []Provider{ProviderIptables, ProviderUFW} {
scope := Scope{Provider: provider, Family: FamilyIPv4, Direction: DirectionInput}
if provider == ProviderIptables {
scope.Table = "filter"
scope.Chain = IptablesInputChain
}
rule, err := NormalizeRule(FirewallRule{
Scope: scope, Protocol: "tcp", DestinationPort: "080,443,8080:8090,443", Action: ActionAccept,
})
if err != nil {
t.Fatalf("normalize %s port set: %v", provider, err)
}
if rule.DestinationPort != "80,443,8080-8090" {
t.Fatalf("unexpected %s port set: %q", provider, rule.DestinationPort)
}
}
_, err := NormalizeRule(FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: FirewalldInputZone, Direction: DirectionInput},
Protocol: "tcp", DestinationPort: "80,443", Action: ActionAccept,
})
if !errors.Is(err, ErrCompositeRule) {
t.Fatalf("expected firewalld port set expansion, got %v", err)
}
tooMany := "1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16"
_, err = NormalizeRule(FirewallRule{
Scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput},
Protocol: "tcp", DestinationPort: tooMany, Action: ActionAccept,
})
if !errors.Is(err, ErrInvalidRule) {
t.Fatalf("expected native port-set limit error, got %v", err)
}
}
func TestNormalizeUFWAllowsAllProtocolsForDestinationPort(t *testing.T) {
rule, err := NormalizeRule(FirewallRule{
Scope: Scope{Provider: ProviderUFW, Family: FamilyIPv4, Direction: DirectionInput},
Protocol: "all", DestinationPort: "53", Action: ActionAccept,
})
if err != nil {
t.Fatalf("normalize UFW all-protocol port rule: %v", err)
}
if rule.Protocol != "all" || rule.DestinationPort != "53" {
t.Fatalf("unexpected normalized rule: %#v", rule)
}
for _, provider := range []Provider{ProviderIptables, ProviderFirewalld} {
testRule := rule
testRule.Scope.Provider = provider
switch provider {
case ProviderIptables:
testRule.Scope.Table = "filter"
testRule.Scope.Chain = IptablesInputChain
case ProviderFirewalld:
testRule.Scope.Family = FamilyInet
testRule.Scope.Zone = FirewalldInputZone
testRule.Scope.Chain = ""
}
if _, err = NormalizeRule(testRule); !errors.Is(err, ErrInvalidRule) {
t.Fatalf("expected %s all-protocol port rejection, got %v", provider, err)
}
}
}
func TestExpandAtomicRules(t *testing.T) {
rules, err := ExpandAtomicRules(FirewallRule{
Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: "1PANEL_BASIC", Direction: DirectionInput},
Protocol: "tcp/udp",
SourceAddress: "172.16.10.111, 172.16.10.112",
DestinationPort: "80,443,80",
Action: ActionAccept,
})
if err != nil {
t.Fatalf("expand atomic rules: %v", err)
}
if len(rules) != 4 {
t.Fatalf("expected 4 rules with native iptables port sets, got %d", len(rules))
}
for _, rule := range rules {
if rule.Protocol == "tcp/udp" || rule.SourceAddress == "" || rule.DestinationPort != "80,443" {
t.Fatalf("rule was not expanded correctly: %#v", rule)
}
}
}
func TestExpandAtomicRulesKeepsNativePortSetsAndSplitsFirewalld(t *testing.T) {
iptablesRules, err := ExpandAtomicRules(FirewallRule{
Scope: Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput},
Protocol: "tcp", DestinationPort: "80,443", Action: ActionAccept,
})
if err != nil || len(iptablesRules) != 1 || iptablesRules[0].DestinationPort != "80,443" {
t.Fatalf("iptables port set was expanded: rules=%#v err=%v", iptablesRules, err)
}
firewalldRules, err := ExpandAtomicRules(FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: FirewalldInputZone, Direction: DirectionInput},
Protocol: "tcp", DestinationPort: "80,443", Action: ActionDrop,
})
if err != nil || len(firewalldRules) != 2 || firewalldRules[0].DestinationPort != "80" || firewalldRules[1].DestinationPort != "443" {
t.Fatalf("firewalld port set was not expanded: rules=%#v err=%v", firewalldRules, err)
}
}
func TestExpandUFWInetByFamily(t *testing.T) {
rules, err := ExpandAtomicRules(FirewallRule{
Scope: Scope{Provider: ProviderUFW, Family: FamilyInet, Direction: DirectionInput},
Protocol: "tcp",
DestinationPort: "22",
Action: ActionAccept,
})
if err != nil {
t.Fatalf("expand UFW families: %v", err)
}
if len(rules) != 2 || rules[0].Scope.Family != FamilyIPv4 || rules[1].Scope.Family != FamilyIPv6 {
t.Fatalf("unexpected UFW family expansion: %#v", rules)
}
rules, err = ExpandAtomicRules(FirewallRule{
Scope: Scope{Provider: ProviderUFW, Family: FamilyInet, Direction: DirectionInput},
Protocol: "all",
SourceAddress: "172.16.10.111",
Action: ActionDrop,
})
if err != nil {
t.Fatalf("expand family-specific UFW rule: %v", err)
}
if len(rules) != 1 || rules[0].Scope.Family != FamilyIPv4 {
t.Fatalf("expected one IPv4 rule, got %#v", rules)
}
rules, err = ExpandAtomicRules(FirewallRule{
Scope: Scope{Provider: ProviderUFW, Family: FamilyInet, Direction: DirectionInput},
Protocol: "all",
SourceAddress: "::/0",
Action: ActionDrop,
})
if err != nil {
t.Fatalf("expand IPv6-any UFW rule: %v", err)
}
if len(rules) != 1 || rules[0].Scope.Family != FamilyIPv6 {
t.Fatalf("expected one IPv6 rule, got %#v", rules)
}
}
func TestNormalizeFirewalldNativeKindSetsExecutionBucket(t *testing.T) {
priority := -100
rich, err := NormalizeRule(FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyIPv4, Zone: "public", Direction: DirectionInput},
NativeKind: NativeKindRichRule, Protocol: "tcp", DestinationPort: "3306", Action: ActionDrop, Priority: &priority,
})
if err != nil {
t.Fatalf("normalize firewalld rich rule: %v", err)
}
if rich.OrderBucket != OrderBucketRichPre || rich.Priority == nil || *rich.Priority != -100 {
t.Fatalf("unexpected rich rule placement: %#v", rich)
}
zonePort, err := NormalizeRule(FirewallRule{
Scope: Scope{Provider: ProviderFirewalld, Family: FamilyInet, Zone: "public", Direction: DirectionInput},
NativeKind: NativeKindZonePort, Protocol: "tcp", DestinationPort: "3306", Action: ActionAccept, Priority: &priority,
})
if err != nil {
t.Fatalf("normalize firewalld zone port: %v", err)
}
if zonePort.OrderBucket != OrderBucketZonePrimitiveAllow || zonePort.Priority != nil {
t.Fatalf("native port exposed fake priority: %#v", zonePort)
}
}
func TestExpandAtomicRulesLimit(t *testing.T) {
var addresses strings.Builder
for i := 1; i <= 65; i++ {
if i > 1 {
addresses.WriteString(",")
}
addresses.WriteString(fmt.Sprintf("10.0.0.%d", i))
}
_, err := ExpandAtomicRules(FirewallRule{
Scope: Scope{Provider: ProviderUFW, Family: FamilyInet, Direction: DirectionInput},
Protocol: "tcp/udp",
SourceAddress: addresses.String(),
Action: ActionAccept,
})
if !errors.Is(err, ErrExpansionLimit) {
t.Fatalf("expected expansion limit error, got %v", err)
}
}
@@ -1,579 +0,0 @@
package firewalld
import (
"context"
"errors"
"reflect"
"sort"
"strings"
"sync"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
func TestPrepareRuleChoosesNativePortOrRichRule(t *testing.T) {
adapter := NewAdapterWithReader(newFakeCommandReader())
zonePort, err := adapter.PrepareRule(filter.FirewallRule{
Scope: testScope(filter.FamilyInet), Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept,
})
if err != nil {
t.Fatalf("prepare native port: %v", err)
}
if zonePort.NativeKind != filter.NativeKindZonePort || zonePort.OrderBucket != filter.OrderBucketZonePrimitiveAllow || zonePort.Priority != nil {
t.Fatalf("simple allow did not remain a native port: %#v", zonePort)
}
priority := -100
rich, err := adapter.PrepareRule(filter.FirewallRule{
Scope: testScope(filter.FamilyIPv4), Protocol: "tcp", SourceAddress: "172.16.10.111", DestinationPort: "443",
Action: filter.ActionDrop, Priority: &priority,
})
if err != nil {
t.Fatalf("prepare rich rule: %v", err)
}
if rich.NativeKind != filter.NativeKindRichRule || rich.OrderBucket != filter.OrderBucketRichPre || rich.SourceAddress != "172.16.10.111/32" {
t.Fatalf("address deny did not become a rich rule: %#v", rich)
}
}
func TestLegacyFirewalldRejectsOnlyExplicitRichRulePriority(t *testing.T) {
reader := newFakeCommandReader()
reader.outputs["--version"] = "0.6.3\n"
adapter := NewAdapterWithReader(reader)
priority := -100
rule := filter.FirewallRule{
Scope: testScope(filter.FamilyIPv4), NativeKind: filter.NativeKindRichRule,
Protocol: "tcp", DestinationPort: "443", Action: filter.ActionDrop, Priority: &priority,
}
if err := adapter.CheckRule(context.Background(), rule); !errors.Is(err, filter.ErrUnsupportedScope) {
t.Fatalf("legacy firewalld accepted explicit priority: %v", err)
}
priority = 0
if err := adapter.CheckRule(context.Background(), rule); err != nil {
t.Fatalf("legacy firewalld rejected priority-zero compatibility rule: %v", err)
}
rule.UUID = "legacy-rich-rule"
snapshot, err := filter.NewSnapshot(testScope(filter.FamilyIPv4), nil)
if err != nil {
t.Fatalf("build legacy snapshot: %v", err)
}
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile legacy compatibility rule: %v", err)
}
if option := plan.Rules[0].Commands[0].Args[1]; strings.Contains(option, "priority=") {
t.Fatalf("legacy compatibility command included priority: %s", option)
}
capabilities, err := adapter.Capabilities(context.Background())
if err != nil || capabilities.ExplicitPriority {
t.Fatalf("legacy firewalld exposed explicit priority: %#v err=%v", capabilities, err)
}
}
func TestModernFirewalldSupportsExplicitRichRulePriority(t *testing.T) {
reader := newFakeCommandReader()
reader.outputs["--version"] = "1.3.4\n"
adapter := NewAdapterWithReader(reader)
priority := 100
rule := filter.FirewallRule{
Scope: testScope(filter.FamilyIPv4), NativeKind: filter.NativeKindRichRule,
Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, Priority: &priority,
}
if err := adapter.CheckRule(context.Background(), rule); err != nil {
t.Fatalf("modern firewalld rejected explicit priority: %v", err)
}
capabilities, err := adapter.Capabilities(context.Background())
if err != nil || !capabilities.ExplicitPriority {
t.Fatalf("modern firewalld hid explicit priority: %#v err=%v", capabilities, err)
}
}
func TestParseFirewalldVersion(t *testing.T) {
for _, test := range []struct {
version string
major, minor int
}{
{version: "0.6.3", major: 0, minor: 6},
{version: "0.7.0-1.el7", major: 0, minor: 7},
{version: "2.1.0", major: 2, minor: 1},
} {
major, minor, err := parseFirewalldVersion(test.version)
if err != nil || major != test.major || minor != test.minor {
t.Fatalf("parse %q: got %d.%d err=%v", test.version, major, minor, err)
}
}
}
func TestCompileCreateUsesExplicitPublicRuntimeAndPermanentCommands(t *testing.T) {
adapter := NewAdapterWithReader(newFakeCommandReader())
snapshot, _ := filter.NewSnapshot(testScope(filter.FamilyInet), nil)
rule := filter.FirewallRule{
UUID: "https", Scope: testScope(filter.FamilyInet), Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept,
}
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile native port: %v", err)
}
want := [][]string{
{"--zone=public", "--add-port=443/tcp"},
{"--permanent", "--zone=public", "--add-port=443/tcp"},
}
if len(plan.Rules) != 1 || len(plan.Rules[0].Commands) != 2 ||
!reflect.DeepEqual(plan.Rules[0].Commands[0].Args, want[0]) || !reflect.DeepEqual(plan.Rules[0].Commands[1].Args, want[1]) {
t.Fatalf("unexpected native port plan: %#v", plan)
}
if plan.Rules[0].Expected.Rule.NativeKind != filter.NativeKindZonePort || plan.Rules[0].Expected.Marker != "" {
t.Fatalf("firewalld plan invented marker or representation: %#v", plan.Rules[0].Expected)
}
richSnapshot, _ := filter.NewSnapshot(testScope(filter.FamilyIPv4), nil)
priority := -100
rich := filter.FirewallRule{
UUID: "blocked-ip", Scope: testScope(filter.FamilyIPv4), Protocol: "tcp", SourceAddress: "172.16.10.111",
DestinationPort: "3306", Action: filter.ActionDrop, Priority: &priority,
}
plan, err = adapter.Compile(richSnapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rich}})
if err != nil {
t.Fatalf("compile rich rule: %v", err)
}
option := `--add-rich-rule=rule family="ipv4" priority="-100" source address="172.16.10.111/32" port port="3306" protocol="tcp" drop`
if plan.Rules[0].Commands[0].Args[1] != option || plan.Rules[0].Expected.Rule.NativeKind != filter.NativeKindRichRule {
t.Fatalf("unexpected rich rule plan: %#v", plan.Rules[0])
}
}
func TestCompileAdoptsCanonicalExternalRuleWithoutSystemMutation(t *testing.T) {
reader := newFakeCommandReader()
reader.set(false, "--list-ports", "8080/tcp\n")
reader.set(true, "--list-ports", "8080/tcp\n")
adapter := NewAdapterWithReader(reader)
snapshot, err := adapter.Observe(context.Background(), testScope(filter.FamilyInet))
if err != nil {
t.Fatalf("observe external port: %v", err)
}
rule := snapshot.Rules[0].Rule
rule.UUID = "adopted"
locator := snapshot.Rules[0].Locator
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeAdopt, After: &rule, Locator: &locator}})
if err != nil {
t.Fatalf("compile adoption: %v", err)
}
if len(plan.Rules[0].Commands) != 0 || plan.Rules[0].Previous == nil || plan.Rules[0].Expected.Locator.Canonical != "port:8080/tcp" {
t.Fatalf("canonical adoption mutated firewalld: %#v", plan.Rules[0])
}
if _, err := adapter.Apply(context.Background(), plan); err != nil {
t.Fatalf("no-op canonical adoption required a writer: %v", err)
}
}
func TestCompileUpdateAndDeleteValidateManagedCanonicalTarget(t *testing.T) {
reader := newFakeCommandReader()
reader.set(false, "--list-ports", "8080/tcp\n")
reader.set(true, "--list-ports", "8080/tcp\n")
adapter := NewAdapterWithReader(reader)
snapshot, _ := adapter.Observe(context.Background(), testScope(filter.FamilyInet))
before := snapshot.Rules[0].Rule
before.UUID = "owned"
after := before
after.DestinationPort = "9090"
locator := snapshot.Rules[0].Locator
update, err := adapter.Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeUpdate, Before: &before, After: &after, Locator: &locator,
}})
if err != nil {
t.Fatalf("compile update: %v", err)
}
if len(update.Rules[0].Commands) != 4 || len(update.Rules[0].RollbackCommands) != 4 ||
update.Rules[0].Commands[0].Args[1] != "--remove-port=8080/tcp" || update.Rules[0].Commands[2].Args[1] != "--add-port=9090/tcp" {
t.Fatalf("unexpected update plan: %#v", update.Rules[0])
}
deletePlan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeDelete, Before: &before, Locator: &locator}})
if err != nil {
t.Fatalf("compile delete: %v", err)
}
if len(deletePlan.Rules[0].Commands) != 2 || deletePlan.Rules[0].Previous == nil || deletePlan.Rules[0].Commands[1].Args[2] != "--remove-port=8080/tcp" {
t.Fatalf("unexpected delete plan: %#v", deletePlan.Rules[0])
}
verified, err := adapter.Verify(context.Background(), deletePlan)
if err != nil || verified.Matched {
t.Fatalf("delete verified while canonical target remained: result=%#v err=%v", verified, err)
}
reader.set(false, "--list-ports", "")
reader.set(true, "--list-ports", "")
verified, err = adapter.Verify(context.Background(), deletePlan)
if err != nil || !verified.Matched {
t.Fatalf("deleted canonical target did not verify: result=%#v err=%v", verified, err)
}
protected := snapshot
protected.Rules = append([]filter.ObservedRule(nil), snapshot.Rules...)
protected.Rules[0].Protected = true
if _, err := adapter.Compile(protected, []filter.DesiredChange{{Operation: filter.ChangeDelete, Before: &before, Locator: &locator}}); !errors.Is(err, filter.ErrProtectedRule) {
t.Fatalf("expected protected delete rejection, got %v", err)
}
}
func TestApplyCompensatesOnlySuccessfulFirewalldSteps(t *testing.T) {
reader := newFakeCommandReader()
writer := &fakeCommandWriter{failAt: 2, err: errors.New("permanent write failed")}
adapter := NewAdapterWithBackend(reader, writer)
snapshot, _ := adapter.Observe(context.Background(), testScope(filter.FamilyInet))
rule := filter.FirewallRule{UUID: "web", Scope: testScope(filter.FamilyInet), Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
plan, _ := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if _, err := adapter.Apply(context.Background(), plan); err == nil || !strings.Contains(err.Error(), "permanent write failed") {
t.Fatalf("expected apply failure, got %v", err)
}
if len(writer.commands) != 3 || writer.commands[2].Args[1] != "--remove-port=80/tcp" {
t.Fatalf("runtime write was not compensated precisely: %#v", writer.commands)
}
}
func TestRollbackReversesFullyAppliedFirewalldPlan(t *testing.T) {
reader := newFakeCommandReader()
writer := &fakeCommandWriter{}
adapter := NewAdapterWithBackend(reader, writer)
snapshot, _ := adapter.Observe(context.Background(), testScope(filter.FamilyInet))
rule := filter.FirewallRule{UUID: "rollback", Scope: testScope(filter.FamilyInet), Protocol: "tcp", DestinationPort: "8080", Action: filter.ActionAccept}
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile rollback plan: %v", err)
}
if err := adapter.Rollback(context.Background(), plan); err != nil {
t.Fatalf("rollback applied plan: %v", err)
}
if len(writer.commands) != 2 || writer.commands[0].Args[0] != "--permanent" ||
writer.commands[0].Args[2] != "--remove-port=8080/tcp" || writer.commands[1].Args[1] != "--remove-port=8080/tcp" {
t.Fatalf("unexpected rollback writes: %#v", writer.commands)
}
}
func TestVerifyRequiresConvergedCanonicalRule(t *testing.T) {
reader := newFakeCommandReader()
adapter := NewAdapterWithReader(reader)
snapshot, _ := adapter.Observe(context.Background(), testScope(filter.FamilyInet))
rule := filter.FirewallRule{UUID: "dns", Scope: testScope(filter.FamilyInet), Protocol: "udp", DestinationPort: "53", Action: filter.ActionAccept}
plan, _ := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
reader.set(false, "--list-ports", "53/udp\n")
verified, err := adapter.Verify(context.Background(), plan)
if err != nil || verified.Matched {
t.Fatalf("runtime-only rule verified as converged: result=%#v err=%v", verified, err)
}
reader.set(true, "--list-ports", "53/udp\n")
verified, err = adapter.Verify(context.Background(), plan)
if err != nil || !verified.Matched {
t.Fatalf("converged rule did not verify: result=%#v err=%v", verified, err)
}
}
func TestCompileRejectsBroadDenyAndPersistenceDrift(t *testing.T) {
adapter := NewAdapterWithReader(newFakeCommandReader())
scope := testScope(filter.FamilyInet)
snapshot, _ := filter.NewSnapshot(scope, nil)
rule := filter.FirewallRule{UUID: "deny-all", Scope: scope, Protocol: "all", Action: filter.ActionDrop}
if _, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}}); !errors.Is(err, filter.ErrLockoutRisk) {
t.Fatalf("expected broad deny rejection, got %v", err)
}
port := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindZonePort, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
observed := observedForRule(port)
observed.Persistence = filter.PersistenceStatusRuntimeOnly
drifted, _ := filter.NewSnapshot(scope, []filter.ObservedRule{observed})
create := filter.FirewallRule{UUID: "https", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept}
if _, err := adapter.Compile(drifted, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &create}}); !errors.Is(err, filter.ErrRuleStale) {
t.Fatalf("expected runtime/permanent drift rejection, got %v", err)
}
}
func TestObservePublicInetMergesNativeObjectsAndReportsScopeNotices(t *testing.T) {
reader := newFakeCommandReader()
reader.set(false, "--list-ports", "22/tcp 53/udp 123/sctp\n")
reader.set(true, "--list-ports", "22/tcp 80/tcp 123/sctp\n")
reader.set(false, "--list-rich-rules", `rule port port="8080" protocol="tcp" accept`+"\n")
reader.set(true, "--list-rich-rules", `rule port port="8080" protocol="tcp" accept`+"\n")
reader.set(false, "--list-services", "ssh dhcpv6-client\n")
reader.set(true, "--list-services", "dhcpv6-client ssh\n")
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), testScope(filter.FamilyInet))
if err != nil {
t.Fatalf("observe public inet: %v", err)
}
if len(snapshot.Rules) != 7 {
t.Fatalf("unexpected firewalld object count: %#v", snapshot.Rules)
}
assertPresence(t, snapshot.Rules, "port:22/tcp", filter.PersistenceStatusConverged)
assertPresence(t, snapshot.Rules, "port:53/udp", filter.PersistenceStatusRuntimeOnly)
assertPresence(t, snapshot.Rules, "port:80/tcp", filter.PersistenceStatusPermanentOnly)
sctp := findObserved(snapshot.Rules, "port:123/sctp")
if sctp == nil || sctp.ParseStatus != filter.ParseStatusOpaque {
t.Fatalf("unsupported native port was guessed: %#v", sctp)
}
rich := findObserved(snapshot.Rules, `rich:rule port port="8080" protocol="tcp" accept`)
if rich == nil || rich.ParseStatus != filter.ParseStatusSupported || rich.Rule.NativeKind != filter.NativeKindRichRule || rich.Rule.OrderBucket != filter.OrderBucketRichZeroAllow {
t.Fatalf("family-neutral rich rule was not normalized: %#v", rich)
}
service := findObserved(snapshot.Rules, "service:ssh")
if service == nil || service.ParseStatus != filter.ParseStatusOpaque || service.Rule.NativeKind != filter.NativeKindZoneService ||
service.Rule.Protocol != "" || service.Rule.DestinationPort != "" || service.Rule.Description != "ssh" || service.Raw != "ssh" {
t.Fatalf("service object was not preserved as opaque: %#v", service)
}
dhcpv6Client := findObserved(snapshot.Rules, "service:dhcpv6-client")
if dhcpv6Client == nil || dhcpv6Client.Rule.Protocol != "" || dhcpv6Client.Rule.DestinationPort != "" ||
dhcpv6Client.Rule.Description != "dhcpv6-client" || dhcpv6Client.Raw != "dhcpv6-client" {
t.Fatalf("dhcpv6 service was exposed as an allow-all rule: %#v", dhcpv6Client)
}
if !hasNotice(snapshot.Notices, filter.ScopeNoticeRuntimePermanentMismatch, "ports") {
t.Fatalf("scope notices missing: %#v", snapshot.Notices)
}
}
func TestObserveReadsRuntimeAndPermanentWithListAll(t *testing.T) {
reader := newFakeCommandReader()
reader.set(false, "--list-ports", "22/tcp\n")
reader.set(true, "--list-ports", "22/tcp 80/tcp\n")
if _, err := NewAdapterWithReader(reader).Observe(context.Background(), testScope(filter.FamilyInet)); err != nil {
t.Fatalf("observe public zone: %v", err)
}
want := []string{
"--zone=public\x00--list-all",
"--permanent\x00--zone=public\x00--list-all",
}
got := reader.readCalls()
sort.Strings(got)
sort.Strings(want)
if !reflect.DeepEqual(got, want) {
t.Fatalf("unexpected firewall-cmd reads: got %#v, want %#v", got, want)
}
}
func TestParseZoneOutput(t *testing.T) {
output := `public (active)
target: default
interfaces: eth0
services: cockpit ssh
ports: 22/tcp 53/udp
protocols:
forward: yes
masquerade: no
forward-ports:
source-ports:
icmp-blocks:
rich rules:
rule family="ipv4" source address="10.0.0.1" accept
rule port port="8080" protocol="tcp" accept`
got := parseZoneOutput(output)
want := zoneOutput{
ports: "22/tcp 53/udp",
services: "cockpit ssh",
active: true,
rich: strings.Join([]string{
`rule family="ipv4" source address="10.0.0.1" accept`,
`rule port port="8080" protocol="tcp" accept`,
}, "\n"),
}
if !reflect.DeepEqual(got, want) {
t.Fatalf("unexpected parsed zone output: got %#v, want %#v", got, want)
}
}
func TestObservePublicPipelineKeepsRuleFamiliesInOneZoneScope(t *testing.T) {
reader := newFakeCommandReader()
rich := strings.Join([]string{
`rule family="ipv4" priority="-100" source address="172.16.10.111" port port="3306" protocol="tcp" drop`,
`rule family="ipv4" destination address="10.0.0.1" reject`,
`rule family="ipv4" source address="10.0.0.0/8" protocol value="tcp" accept`,
`rule family="ipv4" log prefix="audit" accept`,
`rule family="ipv6" source address="2001:db8::1" accept`,
}, "\n") + "\n"
reader.set(false, "--list-rich-rules", rich)
reader.set(true, "--list-rich-rules", rich)
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), testScope(filter.FamilyIPv4))
if err != nil {
t.Fatalf("observe public IPv4: %v", err)
}
if snapshot.Scope.Family != filter.FamilyInet || len(snapshot.Rules) != 5 {
t.Fatalf("unexpected public pipeline: %#v", snapshot)
}
first := snapshot.Rules[0]
if first.ParseStatus != filter.ParseStatusSupported || first.Rule.SourceAddress != "172.16.10.111/32" || first.Rule.DestinationPort != "3306" ||
first.Rule.Priority == nil || *first.Rule.Priority != -100 || first.Rule.OrderBucket != filter.OrderBucketRichPre {
t.Fatalf("priority rich rule was not normalized: %#v", first)
}
opaque := findObserved(snapshot.Rules, `rich:rule family="ipv4" log prefix="audit" accept`)
if opaque == nil || opaque.ParseStatus != filter.ParseStatusOpaque || opaque.Raw == "" {
t.Fatalf("unsupported rich rule was not kept opaque: %#v", opaque)
}
protocolRule := findObserved(snapshot.Rules, `rich:rule family="ipv4" source address="10.0.0.0/8" protocol value="tcp" accept`)
if protocolRule == nil || protocolRule.ParseStatus != filter.ParseStatusSupported || protocolRule.Rule.Protocol != "tcp" || protocolRule.Rule.DestinationPort != "" {
t.Fatalf("protocol-only rich rule was not parsed: %#v", protocolRule)
}
ipv6 := findObserved(snapshot.Rules, `rich:rule family="ipv6" source address="2001:db8::1/128" accept`)
if ipv6 == nil || ipv6.Rule.Scope.Family != filter.FamilyIPv6 {
t.Fatalf("IPv6 rich rule was not retained in the public execution scope: %#v", ipv6)
}
}
func TestNativeDetailRunsDedicatedFirewalldCommand(t *testing.T) {
reader := newFakeCommandReader()
info := "ssh\n ports: 22/tcp\n protocols:\n source-ports:\n helpers:\n destination:"
reader.setServiceInfo(false, "ssh", info)
reader.setServiceInfo(true, "ssh", info)
adapter := NewAdapterWithReader(reader)
runtimeInfo, err := adapter.NativeDetail(context.Background(), "ssh", false)
if err != nil {
t.Fatalf("read runtime service info: %v", err)
}
permanentInfo, err := adapter.NativeDetail(context.Background(), "ssh", true)
if err != nil {
t.Fatalf("read permanent service info: %v", err)
}
if runtimeInfo != info || permanentInfo != info {
t.Fatalf("service info output was not preserved: runtime=%q permanent=%q", runtimeInfo, permanentInfo)
}
}
func TestNativeDetailRejectsInvalidServiceName(t *testing.T) {
adapter := NewAdapterWithReader(newFakeCommandReader())
if _, err := adapter.NativeDetail(context.Background(), "../ssh", false); !errors.Is(err, filter.ErrInvalidRule) {
t.Fatalf("expected invalid service rejection, got %v", err)
}
}
func TestObserveRejectsZoneOutsidePublic(t *testing.T) {
scope := testScope(filter.FamilyInet)
scope.Zone = "work"
_, err := NewAdapterWithReader(newFakeCommandReader()).Observe(context.Background(), scope)
if !errors.Is(err, filter.ErrUnsupportedScope) {
t.Fatalf("expected unsupported scope, got %v", err)
}
}
func TestZoneNoticesReportInactivePublic(t *testing.T) {
notices := publicZoneNotices(zoneOutput{}, zoneOutput{})
if !hasNotice(notices, filter.ScopeNoticeManagedScopeInactive, "") {
t.Fatalf("inactive public notices missing: %#v", notices)
}
}
type fakeCommandReader struct {
mu sync.Mutex
outputs map[string]string
errors map[string]error
calls []string
zones map[bool]zoneOutput
}
type fakeCommandWriter struct {
commands []filter.NativeCommand
failAt int
err error
}
func (f *fakeCommandWriter) Run(_ context.Context, command filter.NativeCommand) error {
f.commands = append(f.commands, command)
if f.failAt > 0 && len(f.commands) == f.failAt {
return f.err
}
return nil
}
func newFakeCommandReader() *fakeCommandReader {
return &fakeCommandReader{
outputs: make(map[string]string), errors: make(map[string]error), zones: map[bool]zoneOutput{false: {active: true}},
}
}
func (f *fakeCommandReader) set(permanent bool, option, output string) {
zone := f.zones[permanent]
switch option {
case "--list-ports":
zone.ports = strings.TrimSpace(output)
case "--list-rich-rules":
zone.rich = strings.TrimSpace(output)
case "--list-services":
zone.services = strings.TrimSpace(output)
}
f.zones[permanent] = zone
}
func (f *fakeCommandReader) setServiceInfo(permanent bool, service, output string) {
args := make([]string, 0, 2)
if permanent {
args = append(args, "--permanent")
}
args = append(args, "--info-service="+service)
f.outputs[strings.Join(args, "\x00")] = output
}
func (f *fakeCommandReader) Read(_ context.Context, args ...string) (string, error) {
key := strings.Join(args, "\x00")
f.mu.Lock()
f.calls = append(f.calls, key)
f.mu.Unlock()
if err := f.errors[key]; err != nil {
return "", err
}
if len(args) >= 2 && args[len(args)-2] == "--zone=public" && args[len(args)-1] == "--list-all" {
permanent := args[0] == "--permanent"
zone := f.zones[permanent]
header := "public"
if zone.active {
header += " (active)"
}
output := header + "\n services: " + zone.services + "\n ports: " + zone.ports + "\n rich rules:"
if zone.rich != "" {
output += "\n " + strings.ReplaceAll(zone.rich, "\n", "\n ")
}
return output + "\n", nil
}
output, exists := f.outputs[key]
if !exists {
return "", errors.New("unexpected firewall-cmd call: " + strings.Join(args, " "))
}
return output, nil
}
func (f *fakeCommandReader) readCalls() []string {
f.mu.Lock()
defer f.mu.Unlock()
return append([]string(nil), f.calls...)
}
func testScope(family filter.Family) filter.Scope {
return filter.Scope{Provider: filter.ProviderFirewalld, Family: family, Zone: "public", Direction: filter.DirectionInput}
}
func findObserved(rules []filter.ObservedRule, canonical string) *filter.ObservedRule {
for index := range rules {
if rules[index].Locator.Canonical == canonical {
return &rules[index]
}
}
return nil
}
func assertPresence(t *testing.T, rules []filter.ObservedRule, canonical string, expected filter.PersistenceStatus) {
t.Helper()
rule := findObserved(rules, canonical)
if rule == nil || rule.Persistence != expected {
t.Fatalf("unexpected presence for %s: %#v", canonical, rule)
}
}
func hasNotice(notices []filter.ScopeNotice, code filter.ScopeNoticeCode, value string) bool {
for _, notice := range notices {
if notice.Code != code {
continue
}
if value == "" {
return true
}
for _, candidate := range notice.Values {
if candidate == value {
return true
}
}
}
return false
}
@@ -1,753 +0,0 @@
package iptables
import (
"context"
"errors"
"fmt"
"reflect"
"strings"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
func TestCompileCreateAndAdoptCommands(t *testing.T) {
scope := testScope("1PANEL_BASIC")
snapshot, err := filter.NewSnapshot(scope, nil)
if err != nil {
t.Fatalf("snapshot: %v", err)
}
rule := filter.FirewallRule{
UUID: "rule-1", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp",
SourceAddress: "10.0.0.0/8", DestinationPort: "443", Action: filter.ActionAccept,
}
adapter := NewAdapterWithReader(&fakeRuleReader{})
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile create: %v", err)
}
command := plan.Rules[0].Commands[0]
if len(plan.Rules) != 1 || command.Executable != "iptables-restore" ||
!strings.Contains(command.Stdin, "-A 1PANEL_BASIC -p tcp -s 10.0.0.0/8 --dport 443 -m comment --comment 1panel-rule:rule-1 -j ACCEPT") ||
plan.Rules[0].Expected.Marker != "1panel-rule:rule-1" {
t.Fatalf("unexpected create plan: %#v", plan)
}
position := 1
adoptSnapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{executorObserved(rule, position)})
plan, err = adapter.Compile(adoptSnapshot, []filter.DesiredChange{{Operation: filter.ChangeAdopt, After: &rule, Locator: &filter.Locator{Position: &position}}})
if err != nil {
t.Fatalf("compile adopt: %v", err)
}
if plan.Rules[0].Commands[0].Executable != "iptables-restore" ||
!strings.Contains(plan.Rules[0].Commands[0].Stdin, "1panel-rule:rule-1") {
t.Fatalf("adoption did not replace the selected position: %#v", plan.Rules[0])
}
}
func TestIPv6ObserveCompileAndCapabilities(t *testing.T) {
scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6)
reader := &fakeRuleReader{output: "-A 1PANEL_BASIC -p ipv6-icmp -s 2001:db8::/64 -m comment --comment \"1panel-rule:ping6\" -j ACCEPT"}
adapter := NewAdapterWithReader(reader)
snapshot, err := adapter.Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe IPv6 rules: %v", err)
}
if len(snapshot.Rules) != 1 || snapshot.Rules[0].Rule.Protocol != "icmpv6" ||
snapshot.Rules[0].Rule.SourceAddress != "2001:db8::/64" {
t.Fatalf("unexpected IPv6 snapshot: %#v", snapshot)
}
rule := filter.FirewallRule{
UUID: "ping6", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "icmpv6",
SourceAddress: "2001:db8::/64", Action: filter.ActionAccept,
}
empty, _ := filter.NewSnapshot(scope, nil)
plan, err := adapter.Compile(empty, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile IPv6 rule: %v", err)
}
command := plan.Rules[0].Commands[0]
if command.Executable != "ip6tables-restore" || !strings.Contains(command.Stdin, "-p ipv6-icmp -s 2001:db8::/64") {
t.Fatalf("unexpected IPv6 command: %#v", command)
}
capabilities, err := adapter.Capabilities(context.Background())
if err != nil || !capabilities.SupportsScope(scope) || !capabilities.SupportsScope(testScope("1PANEL_BASIC")) {
t.Fatalf("IPv4/IPv6 capabilities are incomplete: %#v err=%v", capabilities, err)
}
}
func TestContainsChainDeclarationRejectsMissingManagedChain(t *testing.T) {
if !containsChainDeclaration("-N 1PANEL_BASIC\n-A 1PANEL_BASIC -j ACCEPT\n", "1PANEL_BASIC") {
t.Fatal("managed chain declaration was not detected")
}
if containsChainDeclaration("", "1PANEL_BASIC") ||
containsChainDeclaration("-N 1PANEL_BASIC_OTHER\n", "1PANEL_BASIC") {
t.Fatal("missing managed chain was treated as initialized")
}
}
func TestObserveReportsMissingManagedChainWithoutHidingOtherFamilies(t *testing.T) {
scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6)
adapter := NewAdapterWithReader(&fakeRuleReader{err: fmt.Errorf("%w: IPv6 chain is missing", filter.ErrProviderUnavailable)})
snapshot, err := adapter.Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe missing managed chain: %v", err)
}
if len(snapshot.Rules) != 0 || len(snapshot.Notices) != 1 ||
snapshot.Notices[0].Code != filter.ScopeNoticeManagedScopeMissing {
t.Fatalf("unexpected missing-chain snapshot: %#v", snapshot)
}
}
func TestCompileRejectsICMPFamilyMismatch(t *testing.T) {
for _, test := range []struct {
family filter.Family
protocol string
}{
{family: filter.FamilyIPv4, protocol: "icmpv6"},
{family: filter.FamilyIPv6, protocol: "icmp"},
} {
scope := testScopeFamily("1PANEL_BASIC", test.family)
snapshot, _ := filter.NewSnapshot(scope, nil)
rule := filter.FirewallRule{
UUID: "icmp", Scope: scope, NativeKind: filter.NativeKindRule,
Protocol: test.protocol, Action: filter.ActionAccept,
}
_, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if !errors.Is(err, filter.ErrInvalidRule) {
t.Fatalf("expected %s/%s to be rejected, got %v", test.family, test.protocol, err)
}
}
}
func TestApplyRejectsCrossFamilyExecutable(t *testing.T) {
scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6)
snapshot, _ := filter.NewSnapshot(scope, nil)
rule := filter.FirewallRule{
UUID: "web6", Scope: scope, NativeKind: filter.NativeKindRule,
Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept,
}
writer := &fakeRuleWriter{}
adapter := NewAdapterWithBackend(&fakeRuleReader{}, writer)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile IPv6 rule: %v", err)
}
plan.Rules[0].Commands[0].Executable = "iptables"
if _, err := adapter.Apply(context.Background(), plan); !errors.Is(err, filter.ErrInvalidRule) {
t.Fatalf("expected cross-family command rejection, got %v", err)
}
if len(writer.commands) != 0 {
t.Fatalf("cross-family command reached the writer: %#v", writer.commands)
}
}
func TestCompileRejectsBroadIPv4AndIPv6Deny(t *testing.T) {
for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} {
scope := testScopeFamily("1PANEL_BASIC", family)
snapshot, _ := filter.NewSnapshot(scope, nil)
rule := filter.FirewallRule{
UUID: "deny-all", Scope: scope, NativeKind: filter.NativeKindRule,
Protocol: "all", Action: filter.ActionDrop,
}
_, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if !errors.Is(err, filter.ErrLockoutRisk) {
t.Fatalf("expected broad %s deny to be rejected, got %v", family, err)
}
}
}
func TestCompileInsertsAllowBeforeTerminalDrop(t *testing.T) {
scope := testScope("1PANEL_BASIC_AFTER")
dropTCP := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", Action: filter.ActionDrop}
dropUDP := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "udp", Action: filter.ActionDrop}
snapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{executorObserved(dropTCP, 1), executorObserved(dropUDP, 2)})
allow := filter.FirewallRule{UUID: "dns", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "udp", DestinationPort: "53", Action: filter.ActionAccept}
plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &allow}})
if err != nil {
t.Fatalf("compile terminal insert: %v", err)
}
script := plan.Rules[0].Commands[0].Stdin
if strings.Index(script, "1panel-rule:dns") > strings.Index(script, "-p tcp -j DROP") {
t.Fatalf("allow rule was inserted after terminal drop: %#v", plan.Rules[0].Commands[0])
}
}
func TestCompileCreateUsesRequestedPosition(t *testing.T) {
scope := testScope("1PANEL_BASIC")
first := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
second := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "81", Action: filter.ActionAccept}
snapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{executorObserved(first, 1), executorObserved(second, 2)})
position := int64(2)
rule := filter.FirewallRule{
UUID: "inserted", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp",
DestinationPort: "443", Action: filter.ActionAccept, OrderIndex: &position,
}
plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile positioned create: %v", err)
}
script := plan.Rules[0].Commands[0].Stdin
if strings.Index(script, "1panel-rule:inserted") < strings.Index(script, "--dport 80") ||
strings.Index(script, "1panel-rule:inserted") > strings.Index(script, "--dport 81") {
t.Fatalf("create ignored requested position: %#v", plan.Rules[0].Commands[0])
}
}
func TestMultiportCheckCompileAndObserve(t *testing.T) {
scope := testScope("1PANEL_BASIC")
reader := &fakeRuleReader{}
adapter := NewAdapterWithReader(reader)
rule := filter.FirewallRule{
UUID: "web", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp",
DestinationPort: "80,443,8080-8090", Action: filter.ActionAccept,
}
if err := adapter.CheckRule(context.Background(), rule); err != nil {
t.Fatalf("check multiport: %v", err)
}
if !reflect.DeepEqual(reader.checks, []filter.Family{filter.FamilyIPv4}) {
t.Fatalf("unexpected multiport checks: %#v", reader.checks)
}
if err := adapter.CheckRule(context.Background(), rule); err != nil || len(reader.checks) != 1 {
t.Fatalf("successful multiport capability was not cached: checks=%#v err=%v", reader.checks, err)
}
snapshot, _ := filter.NewSnapshot(scope, nil)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile multiport: %v", err)
}
script := plan.Rules[0].Commands[0].Stdin
if !strings.Contains(script, "-m multiport --dports 80,443,8080:8090") {
t.Fatalf("unexpected multiport command: %s", script)
}
reader.output = `-A 1PANEL_BASIC -p tcp -m multiport --dports 80,443,8080:8090 -j ACCEPT -m comment --comment "web"`
observed, err := adapter.Observe(context.Background(), scope)
if err != nil || len(observed.Rules) != 1 || observed.Rules[0].ParseStatus != filter.ParseStatusSupported ||
observed.Rules[0].Rule.DestinationPort != "80,443,8080-8090" {
t.Fatalf("multiport was not observed: snapshot=%#v err=%v", observed, err)
}
failingReader := &fakeRuleReader{multiportErr: errors.New("extension missing")}
if err := NewAdapterWithReader(failingReader).CheckRule(context.Background(), rule); !errors.Is(err, filter.ErrUnsupportedScope) {
t.Fatalf("expected unavailable multiport error, got %v", err)
}
}
func TestCompileRejectsPositionFromOtherFamilySnapshot(t *testing.T) {
snapshotScope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv4)
snapshot, _ := filter.NewSnapshot(snapshotScope, nil)
position := int64(1)
rule := filter.FirewallRule{
UUID: "ipv6-rule",
Scope: testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6), NativeKind: filter.NativeKindRule,
Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept, OrderIndex: &position,
}
_, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeCreate,
After: &rule,
}})
if !errors.Is(err, filter.ErrUnsupportedScope) {
t.Fatalf("expected cross-family position to be rejected, got %v", err)
}
}
func TestCompileReordersManagedRuleWithinChain(t *testing.T) {
scope := testScope("1PANEL_BASIC")
first := filter.FirewallRule{UUID: "first", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
second := filter.FirewallRule{UUID: "second", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "81", Action: filter.ActionAccept}
third := filter.FirewallRule{UUID: "third", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "82", Action: filter.ActionAccept}
rules := []filter.ObservedRule{executorObserved(first, 1), executorObserved(second, 2), executorObserved(third, 3)}
for index := range rules {
rules[index].Marker = "1panel-rule:" + rules[index].Rule.UUID
}
snapshot, _ := filter.NewSnapshot(scope, rules)
target := int64(3)
after := first
after.OrderIndex = &target
plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeReorder, Before: &first, After: &after, Locator: &rules[0].Locator,
}})
if err != nil {
t.Fatalf("compile reorder: %v", err)
}
rulePlan := plan.Rules[0]
if len(rulePlan.Commands) != 1 || rulePlan.Commands[0].Executable != "iptables-restore" ||
strings.Index(rulePlan.Commands[0].Stdin, "1panel-rule:first") < strings.Index(rulePlan.Commands[0].Stdin, "1panel-rule:third") ||
rulePlan.Expected.Locator.Position == nil || *rulePlan.Expected.Locator.Position != 3 {
t.Fatalf("unexpected reorder plan: %#v", rulePlan)
}
}
func TestCompileUpdateMovesAndChangesManagedRule(t *testing.T) {
scope := testScope("1PANEL_BASIC")
first := filter.FirewallRule{UUID: "first", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
second := filter.FirewallRule{UUID: "second", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "81", Action: filter.ActionAccept}
third := filter.FirewallRule{UUID: "third", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "82", Action: filter.ActionAccept}
rules := []filter.ObservedRule{executorObserved(first, 1), executorObserved(second, 2), executorObserved(third, 3)}
for index := range rules {
rules[index].Marker = "1panel-rule:" + rules[index].Rule.UUID
}
snapshot, _ := filter.NewSnapshot(scope, rules)
target := int64(3)
after := first
after.DestinationPort = "443"
after.OrderIndex = &target
plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeUpdate, Before: &first, After: &after, Locator: &rules[0].Locator,
}})
if err != nil {
t.Fatalf("compile positioned update: %v", err)
}
rulePlan := plan.Rules[0]
if len(rulePlan.Commands) != 1 || rulePlan.Commands[0].Executable != "iptables-restore" ||
!strings.Contains(rulePlan.Commands[0].Stdin, "--dport 443") || rulePlan.Expected.Locator.Position == nil ||
*rulePlan.Expected.Locator.Position != 3 {
t.Fatalf("unexpected positioned update plan: %#v", rulePlan)
}
}
func TestCompileBlocksReorderAcrossExternalOrProtectedRule(t *testing.T) {
scope := testScope("1PANEL_BASIC")
first := filter.FirewallRule{UUID: "first", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
middle := filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "81", Action: filter.ActionAccept}
last := filter.FirewallRule{UUID: "last", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "82", Action: filter.ActionAccept}
rules := []filter.ObservedRule{executorObserved(first, 1), executorObserved(middle, 2), executorObserved(last, 3)}
rules[0].Marker = "1panel-rule:first"
rules[2].Marker = "1panel-rule:last"
snapshot, _ := filter.NewSnapshot(scope, rules)
target := int64(3)
after := first
after.OrderIndex = &target
change := []filter.DesiredChange{{Operation: filter.ChangeReorder, Before: &first, After: &after, Locator: &rules[0].Locator}}
if _, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, change); !errors.Is(err, filter.ErrUnsupportedScope) {
t.Fatalf("expected external boundary rejection, got %v", err)
}
rules[1].Marker = "1panel-rule:middle"
rules[1].Protected = true
snapshot, _ = filter.NewSnapshot(scope, rules)
change[0].Locator = &snapshot.Rules[0].Locator
if _, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, change); !errors.Is(err, filter.ErrProtectedRule) {
t.Fatalf("expected protected boundary rejection, got %v", err)
}
}
func TestApplyUsesCompiledSnapshotWithoutSecondRead(t *testing.T) {
scope := testScope("1PANEL_BASIC")
initial, _ := filter.NewSnapshot(scope, nil)
rule := filter.FirewallRule{UUID: "web", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
reader := &fakeRuleReader{}
writer := &fakeRuleWriter{}
adapter := NewAdapterWithBackend(reader, writer)
plan, err := adapter.Compile(initial, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile create: %v", err)
}
reader.output = "-A 1PANEL_BASIC -p tcp --dport 22 -j ACCEPT"
if _, err := adapter.Apply(context.Background(), plan); err != nil {
t.Fatalf("apply compiled plan: %v", err)
}
if len(writer.commands) != 1 || !writer.saved {
t.Fatalf("compiled plan was not applied: %#v", writer)
}
}
func TestBatchCreateUsesSingleRestoreForOwnedChain(t *testing.T) {
scope := testScope("1PANEL_BASIC")
reader := &fakeRuleReader{output: `-A 1PANEL_BASIC -p tcp --dport 22 -m comment --comment external-ssh -j ACCEPT`}
adapter := NewAdapterWithReader(reader)
snapshot, err := adapter.Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe initial chain: %v", err)
}
first := filter.FirewallRule{
UUID: "web", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp",
DestinationPort: "80", Action: filter.ActionAccept,
}
second := filter.FirewallRule{
UUID: "tls", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp",
DestinationPort: "443", Action: filter.ActionAccept,
}
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{
{Operation: filter.ChangeCreate, After: &first},
{Operation: filter.ChangeCreate, After: &second},
})
if err != nil {
t.Fatalf("compile batch: %v", err)
}
if len(plan.Rules) != 2 || len(plan.Rules[0].Commands) != 1 || len(plan.Rules[1].Commands) != 0 {
t.Fatalf("batch did not compile to one restore command: %#v", plan)
}
command := plan.Rules[0].Commands[0]
if command.Executable != "iptables-restore" || !reflect.DeepEqual(command.Args, []string{"--noflush", "--wait"}) {
t.Fatalf("unexpected restore command: %#v", command)
}
for _, want := range []string{
"*filter\n-F 1PANEL_BASIC\n",
"-A 1PANEL_BASIC -p tcp --dport 22 -m comment --comment external-ssh -j ACCEPT\n",
"--dport 80 -m comment --comment 1panel-rule:web -j ACCEPT\n",
"--dport 443 -m comment --comment 1panel-rule:tls -j ACCEPT\n",
"COMMIT\n",
} {
if !strings.Contains(command.Stdin, want) {
t.Fatalf("restore input does not contain %q:\n%s", want, command.Stdin)
}
}
if strings.Contains(command.Stdin, "-F INPUT") {
t.Fatalf("restore input flushes an unmanaged chain:\n%s", command.Stdin)
}
writer := &fakeRuleWriter{}
adapter = NewAdapterWithBackend(reader, writer)
if _, err = adapter.Apply(context.Background(), plan); err != nil {
t.Fatalf("apply batch: %v", err)
}
if len(writer.commands) != 1 || writer.saveCalls != 1 {
t.Fatalf("batch was not restored and persisted once: %#v", writer)
}
if err = adapter.Rollback(context.Background(), plan); err != nil {
t.Fatalf("rollback batch: %v", err)
}
if len(writer.commands) != 2 || !strings.Contains(writer.commands[1].Stdin, "external-ssh") ||
strings.Contains(writer.commands[1].Stdin, "1panel-rule:web") || writer.saveCalls != 2 {
t.Fatalf("rollback did not restore the original chain: %#v", writer)
}
}
func TestIPv6BatchCreateUsesIPv6Restore(t *testing.T) {
scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6)
snapshot, _ := filter.NewSnapshot(scope, nil)
first := filter.FirewallRule{UUID: "one", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
second := filter.FirewallRule{UUID: "two", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept}
plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{
{Operation: filter.ChangeCreate, After: &first},
{Operation: filter.ChangeCreate, After: &second},
})
if err != nil {
t.Fatalf("compile IPv6 batch: %v", err)
}
if got := plan.Rules[0].Commands[0].Executable; got != "ip6tables-restore" {
t.Fatalf("IPv6 batch executable = %q", got)
}
}
func TestBatchDeleteUsesSingleRestoreAndPreservesExternalRules(t *testing.T) {
scope := testScope("1PANEL_BASIC")
reader := &fakeRuleReader{output: strings.Join([]string{
`-A 1PANEL_BASIC -p tcp --dport 80 -m comment --comment 1panel-rule:web -j ACCEPT`,
`-A 1PANEL_BASIC -p tcp --dport 22 -m comment --comment external-ssh -j ACCEPT`,
`-A 1PANEL_BASIC -p tcp --dport 443 -m comment --comment 1panel-rule:tls -j ACCEPT`,
}, "\n")}
adapter := NewAdapterWithReader(reader)
snapshot, err := adapter.Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe initial chain: %v", err)
}
web := snapshot.Rules[0].Rule
web.UUID = "web"
tls := snapshot.Rules[2].Rule
tls.UUID = "tls"
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{
{Operation: filter.ChangeDelete, Before: &tls, Locator: &snapshot.Rules[2].Locator},
{Operation: filter.ChangeDelete, Before: &web, Locator: &snapshot.Rules[0].Locator},
})
if err != nil {
t.Fatalf("compile batch delete: %v", err)
}
command := plan.Rules[0].Commands[0]
if command.Executable != "iptables-restore" || strings.Contains(command.Stdin, "1panel-rule:web") ||
strings.Contains(command.Stdin, "1panel-rule:tls") || !strings.Contains(command.Stdin, "external-ssh") {
t.Fatalf("unexpected batch delete restore input:\n%s", command.Stdin)
}
rollback := plan.Rules[0].RollbackCommands[0].Stdin
if !strings.Contains(rollback, "1panel-rule:web") || !strings.Contains(rollback, "1panel-rule:tls") ||
!strings.Contains(rollback, "external-ssh") {
t.Fatalf("batch delete rollback does not contain the original chain:\n%s", rollback)
}
}
func TestApplyAndVerifyMarker(t *testing.T) {
scope := testScope("1PANEL_BASIC")
reader := &fakeRuleReader{}
writer := &fakeRuleWriter{}
adapter := NewAdapterWithBackend(reader, writer)
snapshot, _ := filter.NewSnapshot(scope, nil)
rule := filter.FirewallRule{UUID: "ssh", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "22", Action: filter.ActionAccept}
plan, _ := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if _, err := adapter.Apply(context.Background(), plan); err != nil {
t.Fatalf("apply plan: %v", err)
}
if len(writer.commands) != 1 || !writer.saved {
t.Fatalf("command was not executed and persisted: %#v", writer)
}
reader.output = `-A 1PANEL_BASIC -p tcp --dport 22 -m comment --comment "1panel-rule:ssh" -j ACCEPT`
verified, err := adapter.Verify(context.Background(), plan)
if err != nil || !verified.Matched {
t.Fatalf("verify marker: result=%#v err=%v", verified, err)
}
}
func TestApplyCompensatesWhenPersistenceFails(t *testing.T) {
scope := testScope("1PANEL_BASIC")
reader := &fakeRuleReader{}
writer := &fakeRuleWriter{saveErrors: []error{errors.New("disk full"), nil}}
adapter := NewAdapterWithBackend(reader, writer)
snapshot, _ := filter.NewSnapshot(scope, nil)
rule := filter.FirewallRule{UUID: "ssh", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "22", Action: filter.ActionAccept}
plan, _ := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if _, err := adapter.Apply(context.Background(), plan); err == nil || err.Error() != "disk full" {
t.Fatalf("expected original persistence error, got %v", err)
}
if len(writer.commands) != 2 || writer.commands[0].Executable != "iptables-restore" ||
writer.commands[1].Executable != "iptables-restore" ||
!strings.Contains(writer.commands[0].Stdin, "1panel-rule:ssh") ||
strings.Contains(writer.commands[1].Stdin, "1panel-rule:ssh") || writer.saveCalls != 2 {
t.Fatalf("failed write was not compensated and persisted: %#v", writer)
}
}
func TestRollbackReversesFullyAppliedIptablesPlan(t *testing.T) {
scope := testScope("1PANEL_BASIC")
writer := &fakeRuleWriter{}
adapter := NewAdapterWithBackend(&fakeRuleReader{}, writer)
snapshot, _ := filter.NewSnapshot(scope, nil)
rule := filter.FirewallRule{UUID: "rollback", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "8080", Action: filter.ActionAccept}
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile rollback plan: %v", err)
}
if err := adapter.Rollback(context.Background(), plan); err != nil {
t.Fatalf("rollback applied plan: %v", err)
}
if len(writer.commands) != 1 || writer.commands[0].Executable != "iptables-restore" ||
strings.Contains(writer.commands[0].Stdin, "1panel-rule:rollback") || writer.saveCalls != 1 {
t.Fatalf("unexpected rollback writes: %#v", writer)
}
}
func TestApplyCompensatesFailedBatchReorder(t *testing.T) {
scope := testScope("1PANEL_BASIC")
reader := &fakeRuleReader{output: "-A 1PANEL_BASIC -p tcp --dport 80 -m comment --comment 1panel-rule:first -j ACCEPT\n" +
"-A 1PANEL_BASIC -p tcp --dport 81 -m comment --comment 1panel-rule:second -j ACCEPT"}
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe reorder snapshot: %v", err)
}
first := snapshot.Rules[0].Rule
first.UUID = "first"
target := int64(2)
after := first
after.OrderIndex = &target
plan, err := NewAdapterWithReader(reader).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeReorder, Before: &first, After: &after, Locator: &snapshot.Rules[0].Locator,
}})
if err != nil {
t.Fatalf("compile reorder: %v", err)
}
writer := &fakeRuleWriter{runErrors: []error{errors.New("restore failed"), nil}}
adapter := NewAdapterWithBackend(reader, writer)
if _, err := adapter.Apply(context.Background(), plan); err == nil || err.Error() != "restore failed" {
t.Fatalf("expected reorder restore failure, got %v", err)
}
if len(writer.commands) != 1 || writer.commands[0].Executable != "iptables-restore" || writer.saveCalls != 1 {
t.Fatalf("failed batch reorder issued partial commands: %#v", writer)
}
}
func TestVerifyDeleteRequiresMarkerToDisappear(t *testing.T) {
scope := testScope("1PANEL_BASIC")
rule := filter.FirewallRule{UUID: "ssh", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "tcp", DestinationPort: "22", Action: filter.ActionAccept}
position := 1
observed := executorObserved(rule, position)
observed.Marker = "1panel-rule:ssh"
snapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{observed})
plan, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{
{Operation: filter.ChangeDelete, Before: &rule, Locator: &observed.Locator},
})
if err != nil {
t.Fatalf("compile delete: %v", err)
}
reader := &fakeRuleReader{output: `-A 1PANEL_BASIC -p tcp --dport 2200 -m comment --comment "1panel-rule:ssh" -j ACCEPT`}
verified, err := NewAdapterWithReader(reader).Verify(context.Background(), plan)
if err != nil || verified.Matched {
t.Fatalf("delete verification ignored a surviving marker: result=%#v err=%v", verified, err)
}
}
func TestCompileRejectsProtectedMutation(t *testing.T) {
scope := testScope("1PANEL_BASIC_BEFORE")
rule := filter.FirewallRule{UUID: "loopback", Scope: scope, NativeKind: filter.NativeKindRule, Protocol: "all", Interface: "lo", Action: filter.ActionAccept}
position := 1
observed := executorObserved(rule, position)
observed.Protected = true
snapshot, _ := filter.NewSnapshot(scope, []filter.ObservedRule{observed})
_, err := NewAdapterWithReader(&fakeRuleReader{}).Compile(snapshot, []filter.DesiredChange{
{Operation: filter.ChangeDelete, Before: &rule, Locator: &observed.Locator},
})
if !errors.Is(err, filter.ErrProtectedRule) {
t.Fatalf("expected protected mutation rejection, got %v", err)
}
}
func TestObserveMergesIPPortAndCombinedRulesInNativeOrder(t *testing.T) {
scope := testScope("1PANEL_BASIC")
reader := &fakeRuleReader{output: `-N 1PANEL_BASIC
-A 1PANEL_BASIC -s 172.16.10.111/32 -j DROP
-A 1PANEL_BASIC -p tcp -m tcp --dport 22 -j ACCEPT -m comment --comment "ssh"
-A 1PANEL_BASIC -p tcp -m tcp -s 10.0.0.0/8 --dport 443 -j ACCEPT -m comment --comment "1panel-rule:managed"
`}
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe rules: %v", err)
}
if len(snapshot.Rules) != 3 || snapshot.Revision == "" {
t.Fatalf("unexpected snapshot: %#v", snapshot)
}
if snapshot.Rules[0].Rule.SourceAddress != "172.16.10.111/32" || snapshot.Rules[0].Rule.DestinationPort != "" {
t.Fatalf("IP rule was not preserved: %#v", snapshot.Rules[0])
}
if snapshot.Rules[1].Rule.DestinationPort != "22" || snapshot.Rules[1].Rule.Description != "ssh" {
t.Fatalf("port rule was not preserved: %#v", snapshot.Rules[1])
}
if snapshot.Rules[2].Rule.SourceAddress != "10.0.0.0/8" || snapshot.Rules[2].Rule.DestinationPort != "443" || snapshot.Rules[2].Marker != "1panel-rule:managed" {
t.Fatalf("combined managed rule was not preserved: %#v", snapshot.Rules[2])
}
for index, observed := range snapshot.Rules {
if observed.Locator.Position == nil || *observed.Locator.Position != index+1 {
t.Fatalf("native position was not preserved: %#v", observed.Locator)
}
}
}
func TestObserveKeepsUnsupportedRulesOpaque(t *testing.T) {
scope := testScope("1PANEL_BASIC_BEFORE")
reader := &fakeRuleReader{output: `-A 1PANEL_BASIC_BEFORE -m limit --limit 5/min -j ACCEPT
-A 1PANEL_BASIC_BEFORE -p tcp -m multiport --dports 80,443 -j ACCEPT
`}
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe opaque rules: %v", err)
}
if len(snapshot.Rules) != 2 || snapshot.Rules[0].ParseStatus != filter.ParseStatusOpaque || snapshot.Rules[1].ParseStatus != filter.ParseStatusSupported ||
snapshot.Rules[1].Rule.DestinationPort != "80,443" {
t.Fatalf("unsupported rules were guessed: %#v", snapshot.Rules)
}
if snapshot.Rules[0].Raw == "" || snapshot.Rules[0].Locator.Canonical == "" {
t.Fatalf("opaque diagnostics were not retained: %#v", snapshot.Rules[0])
}
}
func TestObserveAcceptsDefaultIPv6RejectRepresentation(t *testing.T) {
scope := testScopeFamily("1PANEL_BASIC", filter.FamilyIPv6)
reader := &fakeRuleReader{output: "-A 1PANEL_BASIC -s 2001:db8::/64 -j REJECT --reject-with icmp6-port-unreachable"}
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe IPv6 reject: %v", err)
}
if len(snapshot.Rules) != 1 || snapshot.Rules[0].ParseStatus != filter.ParseStatusSupported ||
snapshot.Rules[0].Rule.Action != filter.ActionReject {
t.Fatalf("default IPv6 reject became opaque: %#v", snapshot.Rules)
}
reader.output = "-A 1PANEL_BASIC -s 2001:db8::/64 -j REJECT --reject-with icmp6-adm-prohibited"
snapshot, err = NewAdapterWithReader(reader).Observe(context.Background(), scope)
if err != nil || len(snapshot.Rules) != 1 || snapshot.Rules[0].ParseStatus != filter.ParseStatusOpaque {
t.Fatalf("non-default IPv6 reject semantics were guessed: snapshot=%#v err=%v", snapshot, err)
}
}
func TestObserveNormalizesConnectionStateSafetyRule(t *testing.T) {
scope := testScope("1PANEL_BASIC_BEFORE")
reader := &fakeRuleReader{output: `-A 1PANEL_BASIC_BEFORE -m conntrack --ctstate RELATED,ESTABLISHED -j ACCEPT -m comment --comment "ESTABLISHED Whitelist"`}
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe state rule: %v", err)
}
states := snapshot.Rules[0].Rule.ConnectionStates
if len(states) != 2 || states[0] != "established" || states[1] != "related" {
t.Fatalf("connection states were not normalized: %#v", states)
}
if !snapshot.Rules[0].Protected {
t.Fatalf("established whitelist was not protected: %#v", snapshot.Rules[0])
}
}
func TestObserveProtectsSystemPresetChains(t *testing.T) {
before := testScope("1PANEL_BASIC_BEFORE")
beforeSnapshot, err := NewAdapterWithReader(&fakeRuleReader{output: `-A 1PANEL_BASIC_BEFORE -p tcp --dport 8080 -j ACCEPT`}).Observe(context.Background(), before)
if err != nil || len(beforeSnapshot.Rules) != 1 || !beforeSnapshot.Rules[0].Protected {
t.Fatalf("BEFORE preset rule was not protected: snapshot=%#v err=%v", beforeSnapshot, err)
}
after := testScope("1PANEL_BASIC_AFTER")
afterSnapshot, err := NewAdapterWithReader(&fakeRuleReader{output: `-A 1PANEL_BASIC_AFTER -p tcp --dport 8080 -j ACCEPT`}).Observe(context.Background(), after)
if err != nil || len(afterSnapshot.Rules) != 1 || !afterSnapshot.Rules[0].Protected {
t.Fatalf("AFTER preset rule was not protected: snapshot=%#v err=%v", afterSnapshot, err)
}
}
func TestObserveRejectsScopeOutsideOwnedChains(t *testing.T) {
scope := testScope("INPUT")
_, err := NewAdapterWithReader(&fakeRuleReader{}).Observe(context.Background(), scope)
if err == nil {
t.Fatal("expected INPUT scope to be rejected")
}
}
type fakeRuleReader struct {
output string
err error
multiportErr error
checks []filter.Family
}
type fakeRuleWriter struct {
commands []filter.NativeCommand
saved bool
saveCalls int
saveErrors []error
runErrors []error
}
func (f *fakeRuleWriter) Run(_ context.Context, command filter.NativeCommand) error {
f.commands = append(f.commands, command)
if len(f.runErrors) >= len(f.commands) {
return f.runErrors[len(f.commands)-1]
}
return nil
}
func (f *fakeRuleWriter) Save(context.Context, filter.Scope) error {
f.saveCalls++
f.saved = true
if len(f.saveErrors) >= f.saveCalls {
return f.saveErrors[f.saveCalls-1]
}
return nil
}
func executorObserved(rule filter.FirewallRule, position int) filter.ObservedRule {
return filter.ObservedRule{
Rule: rule, ParseStatus: filter.ParseStatusSupported,
Locator: filter.Locator{Provider: filter.ProviderIptables, ScopeKey: rule.Scope.Key(), Position: &position},
}
}
func (f *fakeRuleReader) ListChain(context.Context, filter.Scope) (string, error) {
return f.output, f.err
}
func (f *fakeRuleReader) CheckMultiport(_ context.Context, family filter.Family) error {
f.checks = append(f.checks, family)
return f.multiportErr
}
func testScope(chain string) filter.Scope {
return testScopeFamily(chain, filter.FamilyIPv4)
}
func testScopeFamily(chain string, family filter.Family) filter.Scope {
return filter.Scope{
Provider: filter.ProviderIptables, Family: family, Table: "filter", Chain: chain, Direction: filter.DirectionInput,
}
}
@@ -1,243 +0,0 @@
package nftables
import (
"context"
"errors"
"strings"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
type fakeBackend struct {
output string
commands []filter.NativeCommand
saves int
failAt int
}
func (f *fakeBackend) ListChain(context.Context, filter.Scope) (string, error) { return f.output, nil }
func (f *fakeBackend) Run(_ context.Context, command filter.NativeCommand) error {
f.commands = append(f.commands, command)
if f.failAt != 0 && len(f.commands) == f.failAt {
return errors.New("run failed")
}
return nil
}
func (f *fakeBackend) Save(context.Context) error { f.saves++; return nil }
func TestObserveNativeNftablesRules(t *testing.T) {
scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC")
backend := &fakeBackend{output: `table ip nft_1panel_filter {
chain NFT_1PANEL_BASIC {
meta l4proto tcp ip saddr 10.0.0.0/8 tcp dport 443 accept comment "1panel-rule:web" # handle 12
meta l4proto udp udp dport 53 drop # handle 13
meta l4proto tcp ct state established,related tcp dport 8443 accept comment "1panel-rule:stateful" # handle 14
}
}`}
snapshot, err := NewAdapterWithBackend(backend).Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe: %v", err)
}
if len(snapshot.Rules) != 3 || snapshot.Rules[0].Marker != "1panel-rule:web" || snapshot.Rules[0].Locator.NativeID != "12" {
t.Fatalf("unexpected snapshot: %#v", snapshot)
}
if snapshot.Rules[0].Rule.SourceAddress != "10.0.0.0/8" || snapshot.Rules[0].Rule.DestinationPort != "443" || snapshot.Rules[1].Rule.Action != filter.ActionDrop {
t.Fatalf("unexpected parsed rules: %#v", snapshot.Rules)
}
if len(snapshot.Rules[2].Rule.ConnectionStates) != 2 || snapshot.Rules[2].Rule.ConnectionStates[0] != "established" {
t.Fatalf("connection states were not parsed: %#v", snapshot.Rules[2])
}
}
func TestObserveIgnoresChainDeclarationHandle(t *testing.T) {
scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC")
backend := &fakeBackend{output: `table ip nft_1panel_filter {
chain NFT_1PANEL_BASIC { # handle 3
}
}`}
snapshot, err := NewAdapterWithBackend(backend).Observe(context.Background(), scope)
if err != nil {
t.Fatalf("observe empty chain: %v", err)
}
if len(snapshot.Rules) != 0 {
t.Fatalf("chain declaration was parsed as a rule: %#v", snapshot.Rules)
}
}
func TestApplyCompensatesFailedRulesetTransaction(t *testing.T) {
scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC")
snapshot, err := filter.NewSnapshot(scope, nil)
if err != nil {
t.Fatal(err)
}
rule := filter.FirewallRule{UUID: "web", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept}
backend := &fakeBackend{failAt: 1}
adapter := NewAdapterWithBackend(backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile: %v", err)
}
if _, err := adapter.Apply(context.Background(), plan); err == nil {
t.Fatal("expected ruleset transaction failure")
}
if len(backend.commands) != 2 {
t.Fatalf("expected failed transaction and rollback transaction; got %#v", backend.commands)
}
if backend.saves != 1 {
t.Fatalf("rollback was not persisted: saves=%d", backend.saves)
}
}
func TestCompileApplyAndRollbackRebuildOwnedChain(t *testing.T) {
scope := testScope(filter.FamilyIPv6, "1PANEL_BASIC")
snapshot, err := filter.NewSnapshot(scope, nil)
if err != nil {
t.Fatal(err)
}
rule := filter.FirewallRule{UUID: "dns6", Scope: scope, Protocol: "udp", DestinationPort: "53", Action: filter.ActionAccept}
backend := &fakeBackend{}
adapter := NewAdapterWithBackend(backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile: %v", err)
}
commands := plan.Rules[0].Commands
if len(commands) != 1 || commands[0].Executable != "nft" || strings.Join(commands[0].Args, " ") != "-f -" {
t.Fatalf("unexpected commands: %#v", commands)
}
joined := commands[0].Stdin
for _, want := range []string{"add rule ip6 nft_1panel_filter NFT_1PANEL_BASIC", "udp dport 53", `comment "1panel-rule:dns6"`} {
if !strings.Contains(joined, want) {
t.Fatalf("command %q does not contain %q", joined, want)
}
}
if _, err := adapter.Apply(context.Background(), plan); err != nil {
t.Fatalf("apply: %v", err)
}
if len(backend.commands) != 1 || backend.saves != 1 {
t.Fatalf("unexpected apply calls: commands=%#v saves=%d", backend.commands, backend.saves)
}
if err := adapter.Rollback(context.Background(), plan); err != nil {
t.Fatalf("rollback: %v", err)
}
if len(backend.commands) != 2 || backend.commands[1].Stdin != "flush chain ip6 nft_1panel_filter NFT_1PANEL_BASIC\n" || backend.saves != 2 {
t.Fatalf("unexpected rollback calls: commands=%#v saves=%d", backend.commands, backend.saves)
}
}
func TestCompileAndApplyBatchCreateUsesOneRulesetTransaction(t *testing.T) {
scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC")
snapshot, err := filter.NewSnapshot(scope, nil)
if err != nil {
t.Fatal(err)
}
web := filter.FirewallRule{UUID: "web", Scope: scope, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
tls := filter.FirewallRule{UUID: "tls", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept}
backend := &fakeBackend{}
adapter := NewAdapterWithBackend(backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{
{Operation: filter.ChangeCreate, After: &web},
{Operation: filter.ChangeCreate, After: &tls},
})
if err != nil {
t.Fatalf("compile batch create: %v", err)
}
if len(plan.Rules) != 2 || len(plan.Rules[0].Commands) != 1 || len(plan.Rules[1].Commands) != 0 {
t.Fatalf("batch was not compiled into one transaction: %#v", plan.Rules)
}
script := plan.Rules[0].Commands[0].Stdin
if strings.Count(script, "flush chain ") != 1 || strings.Count(script, "add rule ") != 2 ||
!strings.Contains(script, "1panel-rule:web") || !strings.Contains(script, "1panel-rule:tls") {
t.Fatalf("unexpected batch ruleset:\n%s", script)
}
result, err := adapter.Apply(context.Background(), plan)
if err != nil {
t.Fatalf("apply batch create: %v", err)
}
if len(result.Applied) != 2 || len(backend.commands) != 1 || backend.saves != 1 {
t.Fatalf("batch create was not applied once: result=%#v commands=%#v saves=%d", result, backend.commands, backend.saves)
}
}
func TestCompileBatchDeletePreservesExternalRules(t *testing.T) {
scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC")
positionOne, positionTwo, positionThree := 1, 2, 3
web := filter.FirewallRule{UUID: "web", Scope: scope, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
tls := filter.FirewallRule{UUID: "tls", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept}
external := filter.ObservedRule{
Rule: filter.FirewallRule{Scope: scope, NativeKind: filter.NativeKindOpaque}, Raw: "fib saddr . iif oif missing drop",
ParseStatus: filter.ParseStatusOpaque, Locator: filter.Locator{Provider: filter.ProviderNftables, ScopeKey: scope.Key(), Position: &positionOne},
}
webObserved := observedRule(web, "1panel-rule:web", positionTwo, "")
tlsObserved := observedRule(tls, "1panel-rule:tls", positionThree, "")
snapshot, err := filter.NewSnapshot(scope, []filter.ObservedRule{external, webObserved, tlsObserved})
if err != nil {
t.Fatal(err)
}
adapter := NewAdapterWithBackend(&fakeBackend{})
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{
{Operation: filter.ChangeDelete, Before: &tls, Locator: &tlsObserved.Locator},
{Operation: filter.ChangeDelete, Before: &web, Locator: &webObserved.Locator},
})
if err != nil {
t.Fatalf("compile batch delete: %v", err)
}
applyScript := plan.Rules[0].Commands[0].Stdin
if strings.Count(applyScript, "flush chain ") != 1 || strings.Count(applyScript, "add rule ") != 1 ||
!strings.Contains(applyScript, external.Raw) || strings.Contains(applyScript, "1panel-rule:web") || strings.Contains(applyScript, "1panel-rule:tls") {
t.Fatalf("unexpected batch delete ruleset:\n%s", applyScript)
}
rollbackScript := plan.Rules[0].RollbackCommands[0].Stdin
if !strings.Contains(rollbackScript, "1panel-rule:web") || !strings.Contains(rollbackScript, "1panel-rule:tls") ||
!strings.Contains(rollbackScript, external.Raw) {
t.Fatalf("rollback does not restore the original ruleset:\n%s", rollbackScript)
}
}
func TestVerifyBatchCreateChecksEveryRule(t *testing.T) {
scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC")
snapshot, err := filter.NewSnapshot(scope, nil)
if err != nil {
t.Fatal(err)
}
web := filter.FirewallRule{UUID: "web", Scope: scope, Protocol: "tcp", DestinationPort: "80", Action: filter.ActionAccept}
tls := filter.FirewallRule{UUID: "tls", Scope: scope, Protocol: "tcp", DestinationPort: "443", Action: filter.ActionAccept}
backend := &fakeBackend{output: strings.Join([]string{
`meta l4proto tcp tcp dport 80 accept comment "1panel-rule:web" # handle 10`,
`meta l4proto tcp tcp dport 443 accept comment "1panel-rule:tls" # handle 11`,
}, "\n")}
adapter := NewAdapterWithBackend(backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{
{Operation: filter.ChangeCreate, After: &web},
{Operation: filter.ChangeCreate, After: &tls},
})
if err != nil {
t.Fatalf("compile batch create: %v", err)
}
verified, err := adapter.Verify(context.Background(), plan)
if err != nil || !verified.Matched {
t.Fatalf("verify complete batch: result=%#v err=%v", verified, err)
}
backend.output = `meta l4proto tcp tcp dport 80 accept comment "1panel-rule:web" # handle 10`
verified, err = adapter.Verify(context.Background(), plan)
if err != nil || verified.Matched {
t.Fatalf("missing batch member passed verification: result=%#v err=%v", verified, err)
}
}
func TestParseOpaqueRulePreservesRawExpression(t *testing.T) {
scope := testScope(filter.FamilyIPv4, "1PANEL_BASIC")
backend := &fakeBackend{output: "fib saddr . iif oif missing drop # handle 9\n"}
snapshot, err := NewAdapterWithBackend(backend).Observe(context.Background(), scope)
if err != nil || len(snapshot.Rules) != 1 {
t.Fatalf("observe opaque: snapshot=%#v err=%v", snapshot, err)
}
if snapshot.Rules[0].ParseStatus != filter.ParseStatusOpaque || snapshot.Rules[0].Raw == "" {
t.Fatalf("opaque rule was not preserved: %#v", snapshot.Rules[0])
}
}
func testScope(family filter.Family, chain string) filter.Scope {
return filter.Scope{Provider: filter.ProviderNftables, Family: family, Table: "filter", Chain: chain, Direction: filter.DirectionInput}
}
@@ -1,782 +0,0 @@
package ufw
import (
"context"
"errors"
"fmt"
"reflect"
"strings"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
const numberedFixture = `Status: active
To Action From
-- ------ ----
[ 1] 22/tcp ALLOW IN Anywhere # 1panel-rule:rule-v4
[ 2] Anywhere DENY IN 172.16.10.111
[ 3] OpenSSH ALLOW IN Anywhere
[ 4] 443/tcp (v6) ALLOW IN Anywhere (v6) # web v6
[ 5] Anywhere (v6) on eth0 REJECT IN 2001:db8::/64 (v6)
[ 6] Anywhere ALLOW FWD Anywhere on lxdbr0
[ 7] 53 ALLOW IN Anywhere
[ 8] 6000:6010/udp on eth1 ALLOW IN 10.0.0.0/8
[ 9] 22,80,443/tcp ALLOW IN Anywhere
[10] 2222/tcp LIMIT IN Anywhere # SSH limit
[11] 25/tcp DENY IN 2001:db8::/32
[12] a-very-long-application-profile-name ALLOW IN Anywhere
[13] Anywhere on eth0 ALLOW IN Anywhere (log)
[14] 443/tcp ALLOW OUT Anywhere on eth0 (out)
`
type fakeReader struct {
outputs map[string]string
errors map[string]error
calls [][]string
}
type scriptedBackend struct {
numbered []string
writes []filter.NativeCommand
failAt int
}
func (b *scriptedBackend) Read(_ context.Context, args ...string) (string, error) {
if stringsKey(args) == "status verbose" {
return "Status: active\nDefault: deny (incoming), allow (outgoing), deny (routed)\n", nil
}
if stringsKey(args) != "status numbered" || len(b.numbered) == 0 {
return "", fmt.Errorf("unexpected read: %v", args)
}
output := b.numbered[0]
b.numbered = b.numbered[1:]
return output, nil
}
func (b *scriptedBackend) Run(_ context.Context, command filter.NativeCommand) error {
b.writes = append(b.writes, command)
if b.failAt > 0 && len(b.writes) == b.failAt {
return errors.New("write failed")
}
return nil
}
func (f *fakeReader) Read(_ context.Context, args ...string) (string, error) {
f.calls = append(f.calls, append([]string(nil), args...))
key := stringsKey(args)
return f.outputs[key], f.errors[key]
}
func TestObserveNumberedIncomingIPv4PreservesGlobalPositions(t *testing.T) {
reader := &fakeReader{outputs: map[string]string{
"status numbered": numberedFixture,
"status verbose": "Status: active\nDefault: deny (incoming), allow (outgoing), deny (routed)\n",
}}
adapter := NewAdapterWithReader(reader)
snapshot, err := adapter.Observe(context.Background(), ufwScope(filter.FamilyIPv4))
if err != nil {
t.Fatalf("observe: %v", err)
}
if got, want := len(snapshot.Rules), 9; got != want {
t.Fatalf("expected %d IPv4 rules, got %d", want, got)
}
positions := make([]int, 0, len(snapshot.Rules))
for _, rule := range snapshot.Rules {
positions = append(positions, *rule.Locator.Position)
}
if !reflect.DeepEqual(positions, []int{1, 2, 3, 7, 8, 9, 10, 12, 13}) {
t.Fatalf("unexpected positions: %v", positions)
}
first := snapshot.Rules[0]
if first.ParseStatus != filter.ParseStatusSupported || first.Marker != "1panel-rule:rule-v4" ||
first.Rule.DestinationPort != "22" || first.Rule.Protocol != "tcp" || first.Rule.Action != filter.ActionAccept {
t.Fatalf("unexpected first rule: %#v", first)
}
second := snapshot.Rules[1]
if second.ParseStatus != filter.ParseStatusSupported || second.Rule.SourceAddress != "172.16.10.111/32" || second.Rule.Action != filter.ActionDrop {
t.Fatalf("unexpected source deny: %#v", second)
}
application := snapshot.Rules[2]
if application.ParseStatus != filter.ParseStatusOpaque || application.Rule.NativeKind != filter.NativeKindUFWApplication ||
application.Rule.Description != "OpenSSH" || application.Rule.Protocol != "" || application.Rule.Action != filter.ActionAccept {
t.Fatalf("UFW application profile was not identified: %#v", application)
}
eighth := snapshot.Rules[4]
if eighth.ParseStatus != filter.ParseStatusSupported || eighth.Rule.DestinationPort != "6000-6010" || eighth.Rule.Interface != "eth1" {
t.Fatalf("unexpected range rule: %#v", eighth)
}
barePort := snapshot.Rules[3]
if barePort.ParseStatus != filter.ParseStatusSupported || barePort.Rule.Protocol != "all" || barePort.Rule.DestinationPort != "53" {
t.Fatalf("unexpected bare-port rule: %#v", barePort)
}
for _, index := range []int{2, 7} {
if snapshot.Rules[index].ParseStatus != filter.ParseStatusOpaque {
t.Fatalf("expected rule at slice index %d to be opaque: %#v", index, snapshot.Rules[index])
}
}
multiPort := snapshot.Rules[5]
if multiPort.ParseStatus != filter.ParseStatusSupported || multiPort.Rule.Protocol != "tcp" || multiPort.Rule.DestinationPort != "22,80,443" {
t.Fatalf("multi-port display fields were not preserved: %#v", multiPort)
}
limited := snapshot.Rules[6]
if limited.ParseStatus != filter.ParseStatusPartial || limited.Rule.Protocol != "tcp" || limited.Rule.DestinationPort != "2222" || limited.Rule.Action != filter.ActionAccept {
t.Fatalf("limited rule display fields were not preserved: %#v", limited)
}
logged := snapshot.Rules[8]
if logged.ParseStatus != filter.ParseStatusPartial || logged.Rule.Protocol != "all" || logged.Rule.Interface != "eth0" {
t.Fatalf("logged rule display fields were not preserved: %#v", logged)
}
longApplication := snapshot.Rules[7]
if longApplication.Rule.NativeKind != filter.NativeKindUFWApplication ||
longApplication.Rule.Description != "a-very-long-application-profile-name" {
t.Fatalf("long UFW application profile was not preserved: %#v", longApplication)
}
if len(snapshot.Notices) != 0 {
t.Fatalf("active UFW default policy should not create a notice: %#v", snapshot.Notices)
}
if !reflect.DeepEqual(reader.calls, [][]string{{"status", "numbered"}}) {
t.Fatalf("unexpected commands: %#v", reader.calls)
}
}
func TestNativeDetailRunsUFWAppInfoWithCompleteProfileName(t *testing.T) {
info := "Profile: Nginx Full\nTitle: Web Server (Nginx, HTTP + HTTPS)\nDescription: Small, but very powerful and efficient web server\n\nPorts:\n 80,443/tcp"
reader := &fakeReader{outputs: map[string]string{"app info Nginx Full": info}}
got, err := NewAdapterWithReader(reader).NativeDetail(context.Background(), "Nginx Full", false)
if err != nil {
t.Fatalf("load UFW application detail: %v", err)
}
if got != info || !reflect.DeepEqual(reader.calls, [][]string{{"app", "info", "Nginx Full"}}) {
t.Fatalf("unexpected UFW application detail: output=%q calls=%#v", got, reader.calls)
}
}
func TestNativeDetailRejectsInvalidUFWProfileName(t *testing.T) {
if _, err := NewAdapterWithReader(&fakeReader{}).NativeDetail(context.Background(), "../OpenSSH", false); !errors.Is(err, filter.ErrInvalidRule) {
t.Fatalf("expected invalid UFW profile rejection, got %v", err)
}
}
func TestObserveNumberedIncomingIPv6KeepsFamilyGap(t *testing.T) {
reader := &fakeReader{outputs: map[string]string{
"status numbered": numberedFixture,
"status verbose": "Status: active\nDefault: deny (incoming), allow (outgoing), deny (routed)\n",
}}
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), ufwScope(filter.FamilyIPv6))
if err != nil {
t.Fatalf("observe: %v", err)
}
if len(snapshot.Rules) != 3 || *snapshot.Rules[0].Locator.Position != 4 || *snapshot.Rules[1].Locator.Position != 5 ||
*snapshot.Rules[2].Locator.Position != 11 {
t.Fatalf("unexpected IPv6 rules: %#v", snapshot.Rules)
}
first := snapshot.Rules[0]
if first.ParseStatus != filter.ParseStatusSupported || first.Rule.Description != "web v6" || first.Rule.DestinationPort != "443" {
t.Fatalf("unexpected IPv6 port rule: %#v", first)
}
second := snapshot.Rules[1]
if second.ParseStatus != filter.ParseStatusSupported || second.Rule.Interface != "eth0" ||
second.Rule.SourceAddress != "2001:db8::/64" || second.Rule.Action != filter.ActionReject {
t.Fatalf("unexpected IPv6 reject: %#v", second)
}
third := snapshot.Rules[2]
if third.ParseStatus != filter.ParseStatusSupported || third.Rule.SourceAddress != "2001:db8::/32" || third.Rule.DestinationPort != "25" {
t.Fatalf("unexpected explicit IPv6 rule without family suffix: %#v", third)
}
}
func TestObserveScopesReadsNumberedRulesOnce(t *testing.T) {
reader := &fakeReader{outputs: map[string]string{"status numbered": numberedFixture}}
snapshots, err := NewAdapterWithReader(reader).ObserveScopes(context.Background(), []filter.Scope{
ufwScope(filter.FamilyIPv4),
ufwScope(filter.FamilyIPv6),
})
if err != nil {
t.Fatalf("observe UFW families: %v", err)
}
if !reflect.DeepEqual(reader.calls, [][]string{{"status", "numbered"}}) {
t.Fatalf("UFW numbered rules were not read exactly once: %#v", reader.calls)
}
if len(snapshots) != 2 || snapshots[0].Scope.Family != filter.FamilyIPv4 || len(snapshots[0].Rules) != 9 ||
snapshots[1].Scope.Family != filter.FamilyIPv6 || len(snapshots[1].Rules) != 3 {
t.Fatalf("unexpected multi-family snapshots: %#v", snapshots)
}
}
func TestObserveBarePortRulesAreSupportedForBothFamilies(t *testing.T) {
output := `Status: active
[ 1] 22 ALLOW IN Anywhere
[ 2] 22 (v6) ALLOW IN Anywhere (v6)
`
for _, test := range []struct {
family filter.Family
position int
}{
{family: filter.FamilyIPv4, position: 1},
{family: filter.FamilyIPv6, position: 2},
} {
rules := parseNumberedRules(ufwScope(test.family), output)
if len(rules) != 1 {
t.Fatalf("expected one %s bare-port rule, got %#v", test.family, rules)
}
rule := rules[0]
if rule.ParseStatus != filter.ParseStatusSupported || rule.Rule.Protocol != "all" ||
rule.Rule.DestinationPort != "22" || rule.Locator.Position == nil || *rule.Locator.Position != test.position {
t.Fatalf("unexpected %s bare-port rule: %#v", test.family, rule)
}
}
}
func TestObserveInactiveUFWReturnsNoticeAndEmptyInventory(t *testing.T) {
reader := &fakeReader{outputs: map[string]string{
"status numbered": "Status: inactive\n",
"status verbose": "Status: inactive\n",
}}
snapshot, err := NewAdapterWithReader(reader).Observe(context.Background(), ufwScope(filter.FamilyIPv4))
if err != nil {
t.Fatalf("observe: %v", err)
}
if len(snapshot.Rules) != 0 || !hasNotice(snapshot.Notices, filter.ScopeNoticeManagedScopeInactive) {
t.Fatalf("unexpected inactive snapshot: %#v", snapshot)
}
}
func TestParseAnnotatedMultiPortRulesForBothFamilies(t *testing.T) {
output := `Status: active
[ 1] 80,443/tcp ALLOW IN Anywhere
[ 2] 137,138/udp (Samba) ALLOW IN Anywhere
[ 3] 80,443/tcp (v6) ALLOW IN Anywhere (v6)
[ 4] 137,138/udp (Samba (v6)) ALLOW IN Anywhere (v6)`
tests := []struct {
family filter.Family
positions []int
}{
{family: filter.FamilyIPv4, positions: []int{1, 2}},
{family: filter.FamilyIPv6, positions: []int{3, 4}},
}
for _, test := range tests {
t.Run(string(test.family), func(t *testing.T) {
rules := parseNumberedRules(ufwScope(test.family), output)
if len(rules) != 2 {
t.Fatalf("expected two %s rules, got %#v", test.family, rules)
}
for index, observed := range rules {
if observed.Locator.Position == nil || *observed.Locator.Position != test.positions[index] {
t.Fatalf("unexpected %s rule %d: %#v", test.family, index, observed)
}
}
if rules[0].ParseStatus != filter.ParseStatusSupported || rules[0].Rule.NativeKind != filter.NativeKindUFWRule ||
rules[0].Rule.Protocol != "tcp" || rules[0].Rule.DestinationPort != "80,443" ||
rules[0].Rule.Description != "" {
t.Fatalf("plain multi-port fields were not preserved: %#v", rules[0])
}
if rules[1].ParseStatus != filter.ParseStatusPartial || rules[1].Rule.NativeKind != filter.NativeKindUFWApplication ||
rules[1].Rule.Protocol != "udp" || rules[1].Rule.DestinationPort != "137,138" ||
rules[1].Rule.Description != "Samba" {
t.Fatalf("annotated application fields were not preserved: %#v", rules[1])
}
})
}
}
func TestObserveRejectsScopeOutsideIncoming(t *testing.T) {
adapter := NewAdapterWithReader(&fakeReader{})
_, err := adapter.Observe(context.Background(), filter.Scope{
Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, Chain: "outgoing", Direction: filter.Direction("output"),
})
if !errors.Is(err, filter.ErrInvalidScope) {
t.Fatalf("expected invalid output scope, got %v", err)
}
}
func TestCapabilitiesOnlyAdvertiseAtomicIncomingFamilies(t *testing.T) {
capabilities, err := NewAdapterWithReader(&fakeReader{}).Capabilities(context.Background())
if err != nil {
t.Fatalf("capabilities: %v", err)
}
if !capabilities.Marker || !capabilities.SupportsScope(ufwScope(filter.FamilyIPv4)) ||
!capabilities.SupportsScope(ufwScope(filter.FamilyIPv6)) ||
!capabilities.ExplicitPosition ||
capabilities.SupportsScope(filter.Scope{Provider: filter.ProviderUFW, Family: filter.FamilyIPv4, Chain: "outgoing", Direction: filter.Direction("output")}) {
t.Fatalf("unexpected capabilities: %#v", capabilities)
}
}
func TestPrepareRuleRejectsOptionsUFWCannotSynchronize(t *testing.T) {
rule := writableRule(filter.FamilyIPv4, "", "8080")
rule.SourcePort = "1024"
if _, err := NewAdapterWithReader(&fakeReader{}).PrepareRule(rule); !errors.Is(err, filter.ErrInvalidRule) {
t.Fatalf("unsupported source port passed UFW preview validation: %v", err)
}
}
func TestCompileCreateUsesFamilyExplicitFullSyntax(t *testing.T) {
tests := []struct {
name string
family filter.Family
address string
}{
{name: "IPv4", family: filter.FamilyIPv4, address: "0.0.0.0/0"},
{name: "IPv6", family: filter.FamilyIPv6, address: "::/0"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
snapshot := mustSnapshot(t, ufwScope(test.family), nil)
rule := writableRule(test.family, "new-rule", "8080")
plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeCreate, After: &rule, Append: true,
}})
if err != nil {
t.Fatalf("compile: %v", err)
}
want := []string{"allow", "in", "proto", "tcp", "from", test.address, "to", test.address, "port", "8080", "comment", "1panel-rule:new-rule"}
if !reflect.DeepEqual(plan.Rules[0].Commands[0].Args, want) {
t.Fatalf("unexpected command:\nwant: %#v\n got: %#v", want, plan.Rules[0].Commands[0].Args)
}
if got := plan.Rules[0].RollbackCommands[0].Args[:3]; !reflect.DeepEqual(got, []string{"--force", "delete", "allow"}) {
t.Fatalf("unexpected rollback prefix: %#v", got)
}
})
}
}
func TestCompileCreateBarePortUsesFamilyExplicitSyntaxWithoutProtocol(t *testing.T) {
tests := []struct {
family filter.Family
address string
}{
{family: filter.FamilyIPv4, address: "0.0.0.0/0"},
{family: filter.FamilyIPv6, address: "::/0"},
}
for _, test := range tests {
t.Run(string(test.family), func(t *testing.T) {
snapshot := mustSnapshot(t, ufwScope(test.family), nil)
rule := filter.FirewallRule{
UUID: "dns", Scope: ufwScope(test.family), NativeKind: filter.NativeKindUFWRule,
Protocol: "all", DestinationPort: "53", Action: filter.ActionAccept,
}
plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeCreate, After: &rule, Append: true,
}})
if err != nil {
t.Fatalf("compile bare-port create: %v", err)
}
want := []string{
"allow", "in", "from", test.address, "to", test.address,
"port", "53", "comment", "1panel-rule:dns",
}
if !reflect.DeepEqual(plan.Rules[0].Commands[0].Args, want) {
t.Fatalf("unexpected bare-port command:\nwant: %#v\n got: %#v", want, plan.Rules[0].Commands[0].Args)
}
})
}
}
func TestCompileCreateMultiportUsesNativeUFWPortSet(t *testing.T) {
snapshot := mustSnapshot(t, ufwScope(filter.FamilyIPv4), nil)
rule := writableRule(filter.FamilyIPv4, "web", "80,443,8080-8090")
plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeCreate, After: &rule, Append: true,
}})
if err != nil {
t.Fatalf("compile multiport create: %v", err)
}
wantPort := []string{"port", "80,443,8080:8090"}
args := plan.Rules[0].Commands[0].Args
found := false
for index := 0; index+1 < len(args); index++ {
if reflect.DeepEqual(args[index:index+2], wantPort) {
found = true
break
}
}
if !found {
t.Fatalf("UFW port set was not preserved: %#v", args)
}
}
func TestCompileCreateAppendsWithoutInvalidNextPosition(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
existing := parseNumberedRules(scope, "Status: active\n[ 8] 80/tcp ALLOW IN Anywhere\n")
snapshot := mustSnapshot(t, scope, existing)
order := int64(9)
rule := writableRule(filter.FamilyIPv4, "appended", "8080")
rule.OrderIndex = &order
plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeCreate,
After: &rule,
Append: true,
}})
if err != nil {
t.Fatalf("compile append: %v", err)
}
command := plan.Rules[0].Commands[0]
if len(command.Args) == 0 || command.Args[0] != "allow" {
t.Fatalf("append must not use invalid insert 9: %#v", command.Args)
}
if plan.Rules[0].Expected.Locator.Position == nil || *plan.Rules[0].Expected.Locator.Position != 9 {
t.Fatalf("append verification position was not preserved: %#v", plan.Rules[0].Expected.Locator)
}
}
func TestCompileLastRuleMutationUsesAppendForWriteAndRestore(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
observed := parseNumberedRules(scope, "Status: active\n[ 8] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n")
snapshot := mustSnapshot(t, scope, observed)
locator := observed[0].Locator
order := int64(8)
updated := writableRule(filter.FamilyIPv4, "managed", "443")
updated.OrderIndex = &order
updatePlan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeUpdate,
After: &updated,
Locator: &locator,
Append: true,
RestoreAtEnd: true,
}})
if err != nil {
t.Fatalf("compile last-rule update: %v", err)
}
if got := updatePlan.Rules[0]; got.Commands[1].Args[0] != "allow" || got.RollbackCommands[0].Args[0] != "allow" {
t.Fatalf("last-rule update must append both new and restored rules: %#v", got)
}
before := observed[0].Rule
before.UUID = "managed"
deletePlan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeDelete,
Before: &before,
Locator: &locator,
RestoreAtEnd: true,
}})
if err != nil {
t.Fatalf("compile last-rule delete: %v", err)
}
if got := deletePlan.Rules[0].RollbackCommands[0].Args; len(got) == 0 || got[0] != "allow" {
t.Fatalf("last-rule delete rollback must append: %#v", got)
}
}
func TestCompileAdoptUpdatesCommentWithoutChangingNumber(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
observed := parseNumberedRules(scope, "Status: active\n[ 7] Anywhere DENY IN 172.16.10.111 # imported\n")
if len(observed) != 1 {
t.Fatalf("expected one observed rule: %#v", observed)
}
snapshot := mustSnapshot(t, scope, observed)
rule := observed[0].Rule
rule.UUID = "adopted-rule"
locator := observed[0].Locator
plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeAdopt, After: &rule, Locator: &locator,
}})
if err != nil {
t.Fatalf("compile adoption: %v", err)
}
rulePlan := plan.Rules[0]
if len(rulePlan.Commands) != 1 || rulePlan.Commands[0].Args[0] != "deny" ||
rulePlan.Commands[0].Args[len(rulePlan.Commands[0].Args)-1] != "1panel-rule:adopted-rule" {
t.Fatalf("unexpected adoption command: %#v", rulePlan.Commands)
}
if rulePlan.RollbackCommands[0].Args[len(rulePlan.RollbackCommands[0].Args)-1] != "imported" ||
rulePlan.Expected.Locator.Position == nil || *rulePlan.Expected.Locator.Position != 7 {
t.Fatalf("adoption did not preserve comment and position: %#v", rulePlan)
}
}
func TestCompileUpdateAndDeleteRequireOwnedNumberedRule(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
observed := parseNumberedRules(scope, "Status: active\n[ 3] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n")
snapshot := mustSnapshot(t, scope, observed)
locator := observed[0].Locator
updated := writableRule(filter.FamilyIPv4, "managed", "443")
update, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeUpdate, After: &updated, Locator: &locator,
}})
if err != nil {
t.Fatalf("compile update: %v", err)
}
if got := update.Rules[0].Commands; len(got) != 2 || !reflect.DeepEqual(got[0].Args, []string{"--force", "delete", "3"}) ||
!reflect.DeepEqual(got[1].Args[:2], []string{"insert", "3"}) {
t.Fatalf("unexpected update commands: %#v", got)
}
before := observed[0].Rule
before.UUID = "managed"
deletePlan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeDelete, Before: &before, Locator: &locator,
}})
if err != nil {
t.Fatalf("compile delete: %v", err)
}
if !reflect.DeepEqual(deletePlan.Rules[0].Commands[0].Args, []string{"--force", "delete", "3"}) {
t.Fatalf("unexpected delete command: %#v", deletePlan.Rules[0].Commands)
}
external := observed[0]
external.Marker = ""
external.Rule.Description = "external"
externalSnapshot := mustSnapshot(t, scope, []filter.ObservedRule{external})
_, err = NewAdapterWithReader(&fakeReader{}).Compile(externalSnapshot, []filter.DesiredChange{{
Operation: filter.ChangeDelete, Before: &before, Locator: &locator,
}})
if err == nil {
t.Fatal("expected deletion of an external rule to be rejected")
}
}
func TestCompileUpdateUsesRequestedGlobalPosition(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
observed := parseNumberedRules(scope, "Status: active\n[ 3] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n")
snapshot := mustSnapshot(t, scope, observed)
locator := observed[0].Locator
updated := writableRule(filter.FamilyIPv4, "managed", "443")
target := int64(1)
updated.OrderIndex = &target
plan, err := NewAdapterWithReader(&fakeReader{}).Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeUpdate, Before: &observed[0].Rule, After: &updated, Locator: &locator,
}})
if err != nil {
t.Fatalf("compile positioned update: %v", err)
}
rulePlan := plan.Rules[0]
if len(rulePlan.Commands) != 2 || !reflect.DeepEqual(rulePlan.Commands[0].Args, []string{"--force", "delete", "3"}) ||
!reflect.DeepEqual(rulePlan.Commands[1].Args[:2], []string{"insert", "1"}) ||
rulePlan.Expected.Locator.Position == nil || *rulePlan.Expected.Locator.Position != 1 {
t.Fatalf("unexpected positioned update: %#v", rulePlan)
}
}
func TestCompileRejectsInactiveAndProtected(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
rule := writableRule(filter.FamilyIPv4, "rule", "8080")
inactive := mustSnapshot(t, scope, nil)
inactive.Notices = []filter.ScopeNotice{{Code: filter.ScopeNoticeManagedScopeInactive}}
_, err := NewAdapterWithReader(&fakeReader{}).Compile(inactive, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if !errors.Is(err, filter.ErrProviderUnavailable) {
t.Fatalf("expected inactive provider error, got %v", err)
}
observed := parseNumberedRules(scope, "Status: active\n[ 1] 8080/tcp ALLOW IN Anywhere\n")
observed[0].Protected = true
protected := mustSnapshot(t, scope, observed)
locator := observed[0].Locator
rule.UUID = "protected"
_, err = NewAdapterWithReader(&fakeReader{}).Compile(protected, []filter.DesiredChange{{Operation: filter.ChangeAdopt, After: &rule, Locator: &locator}})
if !errors.Is(err, filter.ErrProtectedRule) {
t.Fatalf("expected protected rule error, got %v", err)
}
_, err = NewAdapterWithReader(&fakeReader{}).Compile(mustSnapshot(t, scope, nil), []filter.DesiredChange{{Operation: filter.ChangeReorder, After: &rule}})
if !errors.Is(err, filter.ErrUnsupportedScope) {
t.Fatalf("expected unsupported standalone reorder error, got %v", err)
}
broadDeny := filter.FirewallRule{
UUID: "deny-all", Scope: scope, NativeKind: filter.NativeKindUFWRule, Protocol: "all", Action: filter.ActionDrop,
}
_, err = NewAdapterWithReader(&fakeReader{}).Compile(mustSnapshot(t, scope, nil), []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &broadDeny}})
if !errors.Is(err, filter.ErrLockoutRisk) {
t.Fatalf("expected broad deny lockout error, got %v", err)
}
}
func TestApplyVerifiesMarkerAcrossBothFamilies(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
snapshot := mustSnapshot(t, scope, nil)
rule := writableRule(filter.FamilyIPv4, "created", "8080")
backend := &scriptedBackend{numbered: []string{
"Status: active\n[ 1] 8080/tcp ALLOW IN Anywhere # 1panel-rule:created\n",
"Status: active\n",
}}
adapter := NewAdapterWithBackend(backend, backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile: %v", err)
}
if _, err = adapter.Apply(context.Background(), plan); err != nil {
t.Fatalf("apply: %v", err)
}
if len(backend.writes) != 1 {
t.Fatalf("unexpected writes: %#v", backend.writes)
}
}
func TestApplyVerifiesMultiportUpdate(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
beforeOutput := "Status: active\n[ 4] 4422,8088/tcp ALLOW IN Anywhere # 1panel-rule:managed\n"
observed := parseNumberedRules(scope, beforeOutput)
if len(observed) != 1 || observed[0].ParseStatus != filter.ParseStatusSupported {
t.Fatalf("unexpected existing multiport rule: %#v", observed)
}
snapshot := mustSnapshot(t, scope, observed)
updated := writableRule(filter.FamilyIPv4, "managed", "4422,8088,7944")
order := int64(4)
updated.OrderIndex = &order
backend := &scriptedBackend{numbered: []string{
"Status: active\n[ 4] 4422,8088,7944/tcp ALLOW IN Anywhere # 1panel-rule:managed\n",
"Status: active\n",
}}
adapter := NewAdapterWithBackend(backend, backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeUpdate, After: &updated, Locator: &observed[0].Locator,
}})
if err != nil {
t.Fatalf("compile multiport update: %v", err)
}
if _, err = adapter.Apply(context.Background(), plan); err != nil {
t.Fatalf("verify multiport update: %v", err)
}
if len(backend.writes) != 2 || backend.writes[1].Args[0] != "insert" || backend.writes[1].Args[1] != "4" {
t.Fatalf("unexpected multiport update writes: %#v", backend.writes)
}
}
func TestApplyCompensatesFamilyExpansion(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
snapshot := mustSnapshot(t, scope, nil)
rule := writableRule(filter.FamilyIPv4, "expanded", "8080")
backend := &scriptedBackend{numbered: []string{
"Status: active\n[ 1] 8080/tcp ALLOW IN Anywhere # 1panel-rule:expanded\n",
"Status: active\n[ 2] 8080/tcp (v6) ALLOW IN Anywhere (v6) # 1panel-rule:expanded\n",
}}
adapter := NewAdapterWithBackend(backend, backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile: %v", err)
}
if _, err = adapter.Apply(context.Background(), plan); err == nil {
t.Fatal("expected one-to-many verification failure")
}
if len(backend.writes) != 2 || backend.writes[1].Args[0] != "--force" || backend.writes[1].Args[1] != "delete" {
t.Fatalf("expected compensating delete, got %#v", backend.writes)
}
}
func TestRollbackReversesFullyAppliedUFWPlan(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
snapshot := mustSnapshot(t, scope, nil)
rule := writableRule(filter.FamilyIPv4, "rollback", "8080")
backend := &scriptedBackend{}
adapter := NewAdapterWithBackend(backend, backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile rollback plan: %v", err)
}
if err := adapter.Rollback(context.Background(), plan); err != nil {
t.Fatalf("rollback applied plan: %v", err)
}
if len(backend.writes) != 1 || !reflect.DeepEqual(backend.writes[0].Args[:2], []string{"--force", "delete"}) {
t.Fatalf("unexpected rollback writes: %#v", backend.writes)
}
}
func TestApplyCompensatesFailedUpdate(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
beforeOutput := "Status: active\n[ 3] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n"
observed := parseNumberedRules(scope, beforeOutput)
snapshot := mustSnapshot(t, scope, observed)
locator := observed[0].Locator
updated := writableRule(filter.FamilyIPv4, "managed", "443")
backend := &scriptedBackend{numbered: []string{
"Status: active\n[ 3] 443/tcp ALLOW IN Anywhere # 1panel-rule:managed\n",
}, failAt: 2}
adapter := NewAdapterWithBackend(backend, backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeUpdate, After: &updated, Locator: &locator}})
if err != nil {
t.Fatalf("compile: %v", err)
}
if _, err = adapter.Apply(context.Background(), plan); err == nil || !strings.Contains(err.Error(), "write failed") {
t.Fatalf("expected write failure, got %v", err)
}
if len(backend.writes) != 4 ||
!reflect.DeepEqual(backend.writes[2].Args[:3], []string{"--force", "delete", "allow"}) ||
!reflect.DeepEqual(backend.writes[3].Args[:2], []string{"insert", "3"}) {
t.Fatalf("expected possibly-applied new rule removal and old rule restore after failed update: %#v", backend.writes)
}
}
func TestApplyDoesNotCompensateFailedUpdateCommandWithoutSideEffect(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
beforeOutput := "Status: active\n[ 3] 80/tcp ALLOW IN Anywhere # 1panel-rule:managed\n"
observed := parseNumberedRules(scope, beforeOutput)
snapshot := mustSnapshot(t, scope, observed)
updated := writableRule(filter.FamilyIPv4, "managed", "443")
backend := &scriptedBackend{numbered: []string{"Status: active\n"}, failAt: 2}
adapter := NewAdapterWithBackend(backend, backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{
Operation: filter.ChangeUpdate, After: &updated, Locator: &observed[0].Locator,
}})
if err != nil {
t.Fatalf("compile: %v", err)
}
if _, err = adapter.Apply(context.Background(), plan); err == nil || !strings.Contains(err.Error(), "write failed") {
t.Fatalf("expected write failure, got %v", err)
}
if len(backend.writes) != 3 || !reflect.DeepEqual(backend.writes[2].Args[:2], []string{"insert", "3"}) {
t.Fatalf("expected only the successfully deleted old rule to be restored: %#v", backend.writes)
}
}
func TestApplyUsesCompiledNumberedSnapshotWithoutPreflight(t *testing.T) {
scope := ufwScope(filter.FamilyIPv4)
snapshot := mustSnapshot(t, scope, nil)
rule := writableRule(filter.FamilyIPv4, "stale", "8080")
backend := &scriptedBackend{numbered: []string{
"Status: active\n[ 1] 8080/tcp ALLOW IN Anywhere # 1panel-rule:stale\n",
"Status: active\n",
}}
adapter := NewAdapterWithBackend(backend, backend)
plan, err := adapter.Compile(snapshot, []filter.DesiredChange{{Operation: filter.ChangeCreate, After: &rule}})
if err != nil {
t.Fatalf("compile: %v", err)
}
if _, err = adapter.Apply(context.Background(), plan); err != nil {
t.Fatalf("apply compiled plan: %v", err)
}
if len(backend.writes) != 1 {
t.Fatalf("compiled plan was not written: %#v", backend.writes)
}
}
func ufwScope(family filter.Family) filter.Scope {
return filter.Scope{Provider: filter.ProviderUFW, Family: family, Chain: "incoming", Direction: filter.DirectionInput}
}
func writableRule(family filter.Family, uuid, port string) filter.FirewallRule {
return filter.FirewallRule{
UUID: uuid, Scope: ufwScope(family), NativeKind: filter.NativeKindUFWRule,
Protocol: "tcp", DestinationPort: port, Action: filter.ActionAccept,
}
}
func mustSnapshot(t *testing.T, scope filter.Scope, rules []filter.ObservedRule) filter.Snapshot {
t.Helper()
snapshot, err := filter.NewSnapshot(scope, rules)
if err != nil {
t.Fatalf("snapshot: %v", err)
}
return snapshot
}
func hasNotice(notices []filter.ScopeNotice, code filter.ScopeNoticeCode) bool {
for _, notice := range notices {
if notice.Code == code {
return true
}
}
return false
}
func stringsKey(values []string) string {
result := ""
for index, value := range values {
if index != 0 {
result += " "
}
result += value
}
return result
}
@@ -0,0 +1,307 @@
package runtime
import (
"context"
"errors"
"fmt"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
filterfirewalld "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/firewalld"
filteriptables "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/iptables"
filternftables "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/nftables"
filterufw "github.com/1Panel-dev/1Panel/agent/utils/firewall/filter/providers/ufw"
)
type SnapshotPolicy func(context.Context, filter.Snapshot) (filter.Snapshot, error)
type Engine struct {
adapter filter.Adapter
policy SnapshotPolicy
}
type Registry map[filter.Provider]*Engine
func NewRegistry(policy SnapshotPolicy) Registry {
return Registry{
filter.ProviderIptables: New(filteriptables.NewAdapter(), policy),
filter.ProviderNftables: New(filternftables.NewAdapter(), policy),
filter.ProviderFirewalld: New(filterfirewalld.NewAdapter(), policy),
filter.ProviderUFW: New(filterufw.NewAdapter(), policy),
}
}
func New(adapter filter.Adapter, policy SnapshotPolicy) *Engine {
return &Engine{adapter: adapter, policy: policy}
}
func (r Registry) Resolve(provider filter.Provider) (*Engine, error) {
engine, exists := r[provider]
if !exists || engine == nil || engine.adapter == nil {
return nil, fmt.Errorf("%w: %s", filter.ErrAdapterUnavailable, provider)
}
return engine, nil
}
func (r Registry) Providers() []filter.Provider {
providers := make([]filter.Provider, 0, len(r))
for provider := range r {
providers = append(providers, provider)
}
return providers
}
func (e *Engine) Provider() filter.Provider {
if e == nil || e.adapter == nil {
return ""
}
return e.adapter.Provider()
}
func (e *Engine) Observe(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) {
snapshot, err := e.adapter.Observe(ctx, scope)
if err != nil {
return filter.Snapshot{}, err
}
if e.policy == nil {
return snapshot, nil
}
return e.policy(ctx, snapshot)
}
func (e *Engine) ObserveScopes(ctx context.Context, scopes []filter.Scope) ([]filter.Snapshot, error) {
observer, ok := e.adapter.(filter.MultiScopeObserver)
if !ok {
return nil, fmt.Errorf("%w: %s multi-scope inventory", filter.ErrAdapterUnavailable, e.adapter.Provider())
}
snapshots, err := observer.ObserveScopes(ctx, scopes)
if err != nil {
return nil, err
}
if e.policy == nil {
return snapshots, nil
}
for index := range snapshots {
snapshots[index], err = e.policy(ctx, snapshots[index])
if err != nil {
return nil, err
}
}
return snapshots, nil
}
func (e *Engine) ObserveMutation(ctx context.Context, scope filter.Scope) (filter.Snapshot, error) {
snapshot, err := e.Observe(ctx, scope)
if err != nil {
return filter.Snapshot{}, err
}
for _, notice := range snapshot.Notices {
if notice.Code == filter.ScopeNoticeManagedScopeInactive || notice.Code == filter.ScopeNoticeManagedScopeMissing {
return filter.Snapshot{}, fmt.Errorf("%w: managed firewall scope is unavailable", filter.ErrProviderUnavailable)
}
}
return snapshot, nil
}
func (e *Engine) Prepare(rule filter.FirewallRule) (filter.FirewallRule, error) {
preparer, ok := e.adapter.(filter.RulePreparer)
if !ok {
return rule, nil
}
return preparer.PrepareRule(rule)
}
func (e *Engine) CheckRule(ctx context.Context, rule filter.FirewallRule) error {
checker, ok := e.adapter.(filter.RuleChecker)
if !ok {
return nil
}
return checker.CheckRule(ctx, rule)
}
func (e *Engine) CompileDesired(
ctx context.Context,
policyUUID string,
origin filter.RuleOrigin,
rules []filter.FirewallRule,
) ([]filter.DesiredRule, error) {
result := make([]filter.DesiredRule, 0, len(rules))
scopeOrdinals := make(map[string]int)
for _, rule := range rules {
prepared, err := e.Prepare(rule)
if err != nil {
return nil, err
}
if err = e.CheckRule(ctx, prepared); err != nil {
return nil, err
}
ruleKey, err := filter.RuleKey(prepared)
if err != nil {
return nil, err
}
scopeKey := prepared.Scope.Key()
ordinal := scopeOrdinals[scopeKey]
scopeOrdinals[scopeKey] = ordinal + 1
prepared.UUID = compiledRuleUUID(policyUUID, ruleKey, ordinal)
result = append(result, filter.DesiredRule{
UUID: policyUUID, Rule: prepared, RuleKey: ruleKey, Origin: origin,
Marker: "1panel-rule:" + prepared.UUID,
})
}
return result, nil
}
func (e *Engine) ValidatePosition(
ctx context.Context,
snapshot filter.Snapshot,
rule filter.FirewallRule,
target int64,
) error {
if rule.Scope.Provider == filter.ProviderUFW {
minimum, maximum := positionBounds(snapshot)
if target < minimum || target > maximum {
return fmt.Errorf(
"%w: target position %d is outside the %s range %d-%d",
filter.ErrInvalidRule, target, rule.Scope.Family, minimum, maximum,
)
}
return nil
}
maximum, err := e.MaxPosition(ctx, snapshot, rule)
if err != nil {
return err
}
if target > maximum {
return fmt.Errorf("%w: target position %d is out of range 1-%d", filter.ErrInvalidRule, target, maximum)
}
return nil
}
func (e *Engine) AppendPosition(ctx context.Context, snapshot filter.Snapshot, rule filter.FirewallRule) (int64, error) {
if rule.Scope.Family == filter.FamilyIPv4 {
return snapshotMaxPosition(snapshot) + 1, nil
}
maximum, err := e.MaxPosition(ctx, snapshot, rule)
if err != nil {
return 0, err
}
return maximum + 1, nil
}
func (e *Engine) MaxPosition(
ctx context.Context,
snapshot filter.Snapshot,
rule filter.FirewallRule,
) (int64, error) {
maximum := snapshotMaxPosition(snapshot)
if rule.Scope.Provider != filter.ProviderUFW {
return maximum, nil
}
relatedScope := rule.Scope
if relatedScope.Family == filter.FamilyIPv4 {
relatedScope.Family = filter.FamilyIPv6
} else {
relatedScope.Family = filter.FamilyIPv4
}
relatedSnapshot, err := e.ObserveMutation(ctx, relatedScope)
if err != nil {
return 0, err
}
if relatedMaximum := snapshotMaxPosition(relatedSnapshot); relatedMaximum > maximum {
maximum = relatedMaximum
}
return maximum, nil
}
func (e *Engine) NativeDetail(ctx context.Context, name string, permanent bool) (string, error) {
reader, ok := e.adapter.(filter.NativeDetailReader)
if !ok {
return "", fmt.Errorf("%w: native details for %s", filter.ErrAdapterUnavailable, e.Provider())
}
return reader.NativeDetail(ctx, name, permanent)
}
func (e *Engine) Capabilities(ctx context.Context) (filter.Capabilities, error) {
return e.adapter.Capabilities(ctx)
}
func (e *Engine) Execute(ctx context.Context, snapshot filter.Snapshot, changes []filter.DesiredChange) (filter.BackendPlan, filter.VerifyResult, error) {
plan, err := e.adapter.Compile(snapshot, changes)
if err != nil {
return filter.BackendPlan{}, filter.VerifyResult{}, err
}
result, err := e.adapter.Apply(ctx, plan)
if err != nil {
return plan, filter.VerifyResult{}, err
}
if result.Verification != nil {
if !result.Verification.Matched {
if rollbackErr := e.Rollback(ctx, plan); rollbackErr != nil {
return plan, *result.Verification, errors.Join(filter.ErrVerificationFailed, rollbackErr)
}
}
return plan, *result.Verification, nil
}
verification, err := e.adapter.Verify(ctx, plan)
if err != nil {
return plan, verification, e.rollback(ctx, plan, err)
}
if !verification.Matched {
if rollbackErr := e.Rollback(ctx, plan); rollbackErr != nil {
return plan, verification, errors.Join(filter.ErrVerificationFailed, rollbackErr)
}
}
return plan, verification, nil
}
func (e *Engine) Rollback(ctx context.Context, plan filter.BackendPlan) error {
rollbacker, ok := e.adapter.(filter.PlanRollbacker)
if !ok {
return fmt.Errorf("%w: provider %s does not support applied-plan rollback", filter.ErrAdapterUnavailable, e.adapter.Provider())
}
return rollbacker.Rollback(ctx, plan)
}
func (e *Engine) rollback(ctx context.Context, plan filter.BackendPlan, cause error) error {
if err := e.Rollback(ctx, plan); err != nil {
return errors.Join(cause, fmt.Errorf("rollback applied firewall plan: %w", err))
}
return cause
}
func positionBounds(snapshot filter.Snapshot) (int64, int64) {
minimum, maximum := int64(0), int64(0)
for _, observed := range snapshot.Rules {
if observed.Locator.Position == nil {
continue
}
position := int64(*observed.Locator.Position)
if minimum == 0 || position < minimum {
minimum = position
}
if position > maximum {
maximum = position
}
}
return minimum, maximum
}
func snapshotMaxPosition(snapshot filter.Snapshot) int64 {
var maximum int64
for _, observed := range snapshot.Rules {
if observed.Locator.Position != nil && int64(*observed.Locator.Position) > maximum {
maximum = int64(*observed.Locator.Position)
}
}
return maximum
}
func compiledRuleUUID(policyUUID, ruleKey string, scopeOrdinal int) string {
if scopeOrdinal == 0 {
return policyUUID
}
const suffixLength = 12
if len(ruleKey) > suffixLength {
ruleKey = ruleKey[:suffixLength]
}
return fmt.Sprintf("%s-%d-%s", policyUUID, scopeOrdinal+1, ruleKey)
}
-139
View File
@@ -1,139 +0,0 @@
package filter
import "testing"
func TestProtectSnapshotTreatsAllProtocolAsCoveringProtectedTransport(t *testing.T) {
scope := Scope{Provider: ProviderUFW, Family: FamilyIPv4, Chain: UFWInputChain, Direction: DirectionInput}
rules := []ObservedRule{
protectedPortTestRule(scope, "all", "22", 1),
protectedPortTestRule(scope, "udp", "22", 2),
protectedPortTestRule(scope, "all", "53", 3),
}
snapshot, err := NewSnapshot(scope, rules)
if err != nil {
t.Fatalf("create snapshot: %v", err)
}
protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Port: "22", Protocol: "tcp"}})
if err != nil {
t.Fatalf("protect snapshot: %v", err)
}
if !protected.Rules[0].Protected {
t.Fatal("all-protocol port 22 did not cover protected 22/tcp")
}
if protected.Rules[1].Protected {
t.Fatal("22/udp incorrectly matched protected 22/tcp")
}
if protected.Rules[2].Protected {
t.Fatal("unrelated all-protocol port was protected")
}
}
func TestProtectSnapshotMarksBareUFWPortForBothFamilies(t *testing.T) {
for _, family := range []Family{FamilyIPv4, FamilyIPv6} {
scope := Scope{Provider: ProviderUFW, Family: family, Chain: UFWInputChain, Direction: DirectionInput}
snapshot, err := NewSnapshot(scope, []ObservedRule{protectedPortTestRule(scope, "all", "22", 1)})
if err != nil {
t.Fatalf("create %s snapshot: %v", family, err)
}
protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Port: "22", Protocol: "tcp"}})
if err != nil {
t.Fatalf("protect %s snapshot: %v", family, err)
}
if !protected.Rules[0].Protected {
t.Fatalf("bare UFW port 22 was not protected for %s", family)
}
}
}
func TestProtectSnapshotMatchesPortSetsAndRanges(t *testing.T) {
scope := Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput}
rules := []ObservedRule{
protectedPortTestRule(scope, "tcp", "22,80,443", 1),
protectedPortTestRule(scope, "tcp", "8000-9000", 2),
protectedPortTestRule(scope, "udp", "22,443", 3),
}
snapshot, err := NewSnapshot(scope, rules)
if err != nil {
t.Fatalf("create snapshot: %v", err)
}
protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Port: "22", Protocol: "tcp"}, {Port: "8080", Protocol: "tcp"}})
if err != nil {
t.Fatalf("protect snapshot: %v", err)
}
if !protected.Rules[0].Protected || !protected.Rules[1].Protected || protected.Rules[2].Protected {
t.Fatalf("unexpected port-set protection: %#v", protected.Rules)
}
}
func TestProtectSnapshotRespectsConfiguredFamily(t *testing.T) {
for _, family := range []Family{FamilyIPv4, FamilyIPv6} {
scope := Scope{Provider: ProviderIptables, Family: family, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput}
snapshot, err := NewSnapshot(scope, []ObservedRule{protectedPortTestRule(scope, "tcp", "443", 1)})
if err != nil {
t.Fatalf("create %s snapshot: %v", family, err)
}
protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Family: "ipv6", Port: "443", Protocol: "tcp"}})
if err != nil {
t.Fatalf("protect %s snapshot: %v", family, err)
}
if protected.Rules[0].Protected != (family == FamilyIPv6) {
t.Fatalf("unexpected %s protection state: %#v", family, protected.Rules[0])
}
}
}
func TestProtectSnapshotDoesNotReclassifyManagedRule(t *testing.T) {
scope := Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput}
managed := protectedPortTestRule(scope, "all", "", 1)
managed.Marker = "1panel-rule:managed-rule"
snapshot, err := NewSnapshot(scope, []ObservedRule{managed})
if err != nil {
t.Fatalf("create snapshot: %v", err)
}
protected, err := ProtectSnapshot(snapshot, []PortWhitelist{{Port: "9999", Protocol: "tcp"}})
if err != nil {
t.Fatalf("protect snapshot: %v", err)
}
if protected.Rules[0].Protected {
t.Fatal("managed rule was reclassified as a protected system rule")
}
}
func TestManagedBroadAllowDoesNotBlockManagedRuleEdit(t *testing.T) {
scope := Scope{Provider: ProviderIptables, Family: FamilyIPv4, Table: "filter", Chain: IptablesInputChain, Direction: DirectionInput}
broadAllow := protectedPortTestRule(scope, "all", "", 1)
broadAllow.Marker = "1panel-rule:broad-allow"
target := protectedPortTestRule(scope, "tcp", "55101", 2)
target.Marker = "1panel-rule:edit-target"
snapshot, err := NewSnapshot(scope, []ObservedRule{broadAllow, target})
if err != nil {
t.Fatalf("create snapshot: %v", err)
}
snapshot, err = ProtectSnapshot(snapshot, []PortWhitelist{{Port: "9999", Protocol: "tcp"}})
if err != nil {
t.Fatalf("protect snapshot: %v", err)
}
after := target.Rule
after.Protocol = "udp"
after.SourceAddress = "198.51.100.21/32"
after.DestinationPort = "55113"
after.Action = ActionDrop
if err := GuardMutation(snapshot, target, after, "", PortWhitelist{Port: "9999", Protocol: "tcp"}); err != nil {
t.Fatalf("managed rule edit was blocked by another managed allow rule: %v", err)
}
}
func protectedPortTestRule(scope Scope, protocol, port string, position int) ObservedRule {
return ObservedRule{
Rule: FirewallRule{
Scope: scope, NativeKind: NativeKindUFWRule, Protocol: protocol,
DestinationPort: port, Action: ActionAccept,
},
Locator: Locator{
Provider: scope.Provider, ScopeKey: scope.Key(), Position: &position,
},
ParseStatus: ParseStatusSupported,
}
}
@@ -1,52 +0,0 @@
package forwarding
import (
"strings"
"github.com/1Panel-dev/1Panel/agent/constant"
)
const (
FamilyIPv4 = constant.FirewallFamilyIPv4
FamilyIPv6 = constant.FirewallFamilyIPv6
ChainPreRouting = "1PANEL_PREROUTING"
ChainPostRouting = "1PANEL_POSTROUTING"
ChainForward = "1PANEL_FORWARD"
ForwardFile = "1panel_forward.rules"
PreRoutingFile = "1panel_forward_pre.rules"
PostRoutingFile = "1panel_forward_post.rules"
)
type Rule struct {
Num string
Family string
Protocol string
Port string
TargetIP string
TargetPort string
Interface string
}
func (r Rule) Identity() string {
return strings.Join([]string{r.Family, r.Protocol, r.Port, r.TargetIP, r.TargetPort, r.Interface}, "\x00")
}
type OperationType string
const (
OperationAdd OperationType = "add"
OperationRemove OperationType = "remove"
)
type Adapter interface {
Name() string
List() ([]Rule, error)
Reconcile(rules []Rule) error
Enable() error
Cleanup() error
InitStatus() (bool, bool, error)
FamilyStatus(family string) (bool, bool, error)
Replay() error
}
@@ -1,20 +0,0 @@
package forwarding
import "testing"
func TestRuleIdentityIncludesEveryIdentityField(t *testing.T) {
base := Rule{Family: FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80", Interface: "eth0"}
variants := []Rule{
{Family: FamilyIPv6, Protocol: base.Protocol, Port: base.Port, TargetIP: base.TargetIP, TargetPort: base.TargetPort, Interface: base.Interface},
{Family: base.Family, Protocol: "udp", Port: base.Port, TargetIP: base.TargetIP, TargetPort: base.TargetPort, Interface: base.Interface},
{Family: base.Family, Protocol: base.Protocol, Port: "8081", TargetIP: base.TargetIP, TargetPort: base.TargetPort, Interface: base.Interface},
{Family: base.Family, Protocol: base.Protocol, Port: base.Port, TargetIP: "127.0.0.2", TargetPort: base.TargetPort, Interface: base.Interface},
{Family: base.Family, Protocol: base.Protocol, Port: base.Port, TargetIP: base.TargetIP, TargetPort: "81", Interface: base.Interface},
{Family: base.Family, Protocol: base.Protocol, Port: base.Port, TargetIP: base.TargetIP, TargetPort: base.TargetPort, Interface: "eth1"},
}
for _, variant := range variants {
if variant.Identity() == base.Identity() {
t.Fatalf("identity collision for %#v", variant)
}
}
}
@@ -0,0 +1,203 @@
package forwarding
import (
"errors"
"fmt"
"net/netip"
"strconv"
"strings"
"sync"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/utils/re"
)
var ErrRuleExists = errors.New("forwarding rule already exists")
const (
FamilyIPv4 = constant.FirewallFamilyIPv4
FamilyIPv6 = constant.FirewallFamilyIPv6
ChainPreRouting = "1PANEL_PREROUTING"
ChainPostRouting = "1PANEL_POSTROUTING"
ChainForward = "1PANEL_FORWARD"
ForwardFile = "1panel_forward.rules"
PreRoutingFile = "1panel_forward_pre.rules"
PostRoutingFile = "1panel_forward_post.rules"
)
type Rule struct {
Num string
Family string
Protocol string
Port string
TargetIP string
TargetPort string
Interface string
}
func (r Rule) Identity() string {
return strings.Join([]string{r.Family, r.Protocol, r.Port, r.TargetIP, r.TargetPort, r.Interface}, "\x00")
}
type OperationType string
const (
OperationAdd OperationType = "add"
OperationRemove OperationType = "remove"
)
type Adapter interface {
Name() string
List() ([]Rule, error)
Reconcile(rules []Rule) error
Enable() error
Cleanup() error
InitStatus() (bool, bool, error)
FamilyStatus(family string) (bool, bool, error)
Replay() error
}
type Status struct {
Name string
Version string
IsInit bool
IsBind bool
}
type RuntimeClient interface {
Version() (string, error)
}
type Manager struct {
adapter Adapter
runtime RuntimeClient
}
func NewManager(adapter Adapter, runtime RuntimeClient) *Manager {
return &Manager{adapter: adapter, runtime: runtime}
}
func (m *Manager) Status() (Status, error) {
status := Status{Name: m.adapter.Name(), Version: "-"}
var versionErr error
var initErr error
var wg sync.WaitGroup
wg.Add(1)
if m.runtime != nil {
wg.Add(1)
go func() {
defer wg.Done()
status.Version, versionErr = m.runtime.Version()
}()
}
go func() {
defer wg.Done()
status.IsInit, status.IsBind, initErr = m.adapter.InitStatus()
}()
wg.Wait()
return status, errors.Join(versionErr, initErr)
}
func (m *Manager) List(info, strategy string) ([]Rule, error) {
rules, err := m.adapter.List()
if err != nil {
return nil, err
}
if strategy != "" {
return []Rule{}, nil
}
filtered := make([]Rule, 0, len(rules))
for _, rule := range rules {
if info != "" && !strings.Contains(rule.Port, info) &&
!strings.Contains(rule.TargetPort, info) && !strings.Contains(rule.TargetIP, info) {
continue
}
filtered = append(filtered, rule)
}
return filtered, nil
}
func (m *Manager) Enable() error { return m.adapter.Enable() }
func (m *Manager) Reconcile(rules []Rule) error { return m.adapter.Reconcile(rules) }
func (m *Manager) Cleanup() error { return m.adapter.Cleanup() }
func (m *Manager) FamilyStatus(family string) (bool, bool, error) {
return m.adapter.FamilyStatus(family)
}
func (m *Manager) Replay() error { return m.adapter.Replay() }
func (m *Manager) Name() string { return m.adapter.Name() }
func NormalizeRule(rule Rule) (Rule, error) {
rule.Family = strings.ToLower(strings.TrimSpace(rule.Family))
if rule.Family == "" {
rule.Family = FamilyIPv4
}
if rule.Family != FamilyIPv4 && rule.Family != FamilyIPv6 {
return Rule{}, fmt.Errorf("unsupported forwarding family %q", rule.Family)
}
rule.Protocol = strings.ToLower(strings.TrimSpace(rule.Protocol))
if rule.Protocol != "tcp" && rule.Protocol != "udp" {
return Rule{}, fmt.Errorf("unsupported forwarding protocol %q", rule.Protocol)
}
var err error
if rule.Port, err = normalizeForwardPort(rule.Port); err != nil {
return Rule{}, fmt.Errorf("invalid forwarding port: %w", err)
}
if rule.TargetPort, err = normalizeForwardPort(rule.TargetPort); err != nil {
return Rule{}, fmt.Errorf("invalid forwarding target port: %w", err)
}
rule.TargetIP = strings.TrimSpace(rule.TargetIP)
if rule.TargetIP == "" || strings.EqualFold(rule.TargetIP, "localhost") {
if rule.Family == FamilyIPv6 {
rule.TargetIP = "::1"
} else {
rule.TargetIP = "127.0.0.1"
}
}
address, err := netip.ParseAddr(rule.TargetIP)
if err == nil {
address = address.Unmap()
}
if err != nil || (rule.Family == FamilyIPv4) != address.Is4() {
return Rule{}, fmt.Errorf("invalid %s forwarding target %q", rule.Family, rule.TargetIP)
}
rule.TargetIP = address.String()
rule.Interface = strings.TrimSpace(rule.Interface)
if rule.Interface == "all" || rule.Interface == "*" {
rule.Interface = ""
}
if rule.Interface != "" && !re.ForwardInterfaceRegex.MatchString(rule.Interface) {
return Rule{}, fmt.Errorf("invalid forwarding interface %q", rule.Interface)
}
return rule, nil
}
func normalizeForwardPort(value string) (string, error) {
parts := strings.Split(strings.TrimSpace(value), "-")
if len(parts) < 1 || len(parts) > 2 {
return "", fmt.Errorf("invalid port range %q", value)
}
ports := make([]int, len(parts))
for index, part := range parts {
port, err := strconv.Atoi(strings.TrimSpace(part))
if err != nil || port < 1 || port > 65535 {
return "", fmt.Errorf("invalid port %q", part)
}
ports[index] = port
}
if len(ports) == 2 {
if ports[0] > ports[1] {
return "", fmt.Errorf("descending port range %q", value)
}
if ports[0] != ports[1] {
return strconv.Itoa(ports[0]) + "-" + strconv.Itoa(ports[1]), nil
}
}
return strconv.Itoa(ports[0]), nil
}
@@ -1,83 +0,0 @@
package forwarding
import (
"errors"
"strings"
"sync"
)
var ErrRuleExists = errors.New("forwarding rule already exists")
type Status struct {
Name string
Version string
IsInit bool
IsBind bool
}
type RuntimeClient interface {
Version() (string, error)
}
type Manager struct {
adapter Adapter
runtime RuntimeClient
}
func NewManager(adapter Adapter, runtime RuntimeClient) *Manager {
return &Manager{adapter: adapter, runtime: runtime}
}
func (m *Manager) Status() (Status, error) {
status := Status{Name: m.adapter.Name(), Version: "-"}
var versionErr error
var initErr error
var wg sync.WaitGroup
wg.Add(1)
if m.runtime != nil {
wg.Add(1)
go func() {
defer wg.Done()
status.Version, versionErr = m.runtime.Version()
}()
}
go func() {
defer wg.Done()
status.IsInit, status.IsBind, initErr = m.adapter.InitStatus()
}()
wg.Wait()
return status, errors.Join(versionErr, initErr)
}
func (m *Manager) List(info, strategy string) ([]Rule, error) {
rules, err := m.adapter.List()
if err != nil {
return nil, err
}
if strategy != "" {
return []Rule{}, nil
}
filtered := make([]Rule, 0, len(rules))
for _, rule := range rules {
if info != "" && !strings.Contains(rule.Port, info) &&
!strings.Contains(rule.TargetPort, info) && !strings.Contains(rule.TargetIP, info) {
continue
}
filtered = append(filtered, rule)
}
return filtered, nil
}
func (m *Manager) Enable() error { return m.adapter.Enable() }
func (m *Manager) Reconcile(rules []Rule) error { return m.adapter.Reconcile(rules) }
func (m *Manager) Cleanup() error { return m.adapter.Cleanup() }
func (m *Manager) FamilyStatus(family string) (bool, bool, error) {
return m.adapter.FamilyStatus(family)
}
func (m *Manager) Replay() error { return m.adapter.Replay() }
func (m *Manager) Name() string { return m.adapter.Name() }
@@ -1,329 +0,0 @@
package providers
import (
"os"
"reflect"
"strings"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/iptables_helper"
)
type commandCall struct {
name string
args []string
}
func commandKey(name string, args ...string) string {
return strings.Join(append([]string{name}, args...), " ")
}
func TestForwardingAdapterFactoryContract(t *testing.T) {
for _, name := range []string{"iptables", "nftables"} {
adapter, err := New(name)
if err != nil {
t.Fatalf("%s: %v", name, err)
}
if adapter.Name() != name {
t.Fatalf("got adapter %q want %q", adapter.Name(), name)
}
if name == "nftables" {
if _, ok := adapter.(*nftablesAdapter); !ok {
t.Fatalf("nftables must use its native forwarding adapter, got %T", adapter)
}
} else if _, ok := adapter.(*iptablesNATAdapter); !ok {
t.Fatalf("%s must use the iptables forwarding adapter, got %T", name, adapter)
}
}
for _, name := range []string{"firewalld", "ufw", "unknown"} {
if _, err := New(name); err == nil {
t.Fatalf("unsupported forwarding provider %q must be rejected", name)
}
}
}
type backendCall struct {
method string
table string
args []string
}
type fakeIptablesBackend struct {
calls []backendCall
stdout map[string]string
err error
ipv6 bool
}
func (f *fakeIptablesBackend) IPv6Available() bool { return f.ipv6 }
func (f *fakeIptablesBackend) Run(table string, args ...string) error {
f.calls = append(f.calls, backendCall{method: "run", table: table, args: append([]string(nil), args...)})
return f.err
}
func (f *fakeIptablesBackend) RunWithStd(table string, args ...string) (string, error) {
f.calls = append(f.calls, backendCall{method: "stdout", table: table, args: append([]string(nil), args...)})
return f.stdout[commandKey(table, args...)], f.err
}
func (f *fakeIptablesBackend) RunIPv6(table string, args ...string) error {
f.calls = append(f.calls, backendCall{method: "run6", table: table, args: append([]string(nil), args...)})
return f.err
}
func (f *fakeIptablesBackend) RunIPv6WithStd(table string, args ...string) (string, error) {
f.calls = append(f.calls, backendCall{method: "stdout6", table: table, args: append([]string(nil), args...)})
return f.stdout["ipv6 "+commandKey(table, args...)], f.err
}
func (f *fakeIptablesBackend) AddChainWithAppend(table, parentChain, chain string) error {
f.calls = append(f.calls, backendCall{method: "add-chain", table: table, args: []string{parentChain, chain}})
return f.err
}
func (f *fakeIptablesBackend) AddIPv6ChainWithAppend(table, parentChain, chain string) error {
f.calls = append(f.calls, backendCall{method: "add-chain6", table: table, args: []string{parentChain, chain}})
return f.err
}
func (f *fakeIptablesBackend) Restore(family, input string) error {
f.calls = append(f.calls, backendCall{method: "restore", table: family, args: []string{input}})
return f.err
}
func (f *fakeIptablesBackend) LoadRulesFromFile(table, chain, fileName string) error {
f.calls = append(f.calls, backendCall{method: "load", table: table, args: []string{chain, fileName}})
return f.err
}
func (f *fakeIptablesBackend) LoadIPv6RulesFromFile(table, chain, fileName string) error {
f.calls = append(f.calls, backendCall{method: "load6", table: table, args: []string{chain, fileName}})
return f.err
}
func TestIptablesReconcileRebuildsOwnedChains(t *testing.T) {
backend := &fakeIptablesBackend{stdout: map[string]string{}}
adapter := &iptablesNATAdapter{provider: "iptables", backend: backend, system: &fakeForwardingSystem{}}
if err := adapter.Reconcile([]forwarding.Rule{{
Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80",
}}); err != nil {
t.Fatal(err)
}
wantScript := "*nat\n" +
"-F " + forwarding.ChainPreRouting + "\n" +
"-F " + forwarding.ChainPostRouting + "\n" +
"-A " + forwarding.ChainPreRouting + " -p tcp --dport 8080 -j REDIRECT --to-port 80\n" +
"COMMIT\n" +
"*filter\n" +
"-F " + forwarding.ChainForward + "\n" +
"COMMIT\n"
want := []backendCall{{method: "restore", table: forwarding.FamilyIPv4, args: []string{wantScript}}}
if !reflect.DeepEqual(backend.calls, want) {
t.Fatalf("reconcile transcript changed\ngot %#v\nwant %#v", backend.calls, want)
}
}
func TestIptablesReconcileBatchesIPv4AndIPv6Separately(t *testing.T) {
backend := &fakeIptablesBackend{stdout: map[string]string{}, ipv6: true}
adapter := &iptablesNATAdapter{provider: "iptables", backend: backend, system: &fakeForwardingSystem{}}
if err := adapter.Reconcile([]forwarding.Rule{{
Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "8443", TargetIP: "2001:db8::20", TargetPort: "443",
}}); err != nil {
t.Fatal(err)
}
if len(backend.calls) != 2 || backend.calls[0].method != "restore" || backend.calls[0].table != forwarding.FamilyIPv4 ||
backend.calls[1].method != "restore" || backend.calls[1].table != forwarding.FamilyIPv6 {
t.Fatalf("expected one restore call per family, got %#v", backend.calls)
}
if !strings.Contains(backend.calls[1].args[0], "--to-destination [2001:db8::20]:443") {
t.Fatalf("unexpected IPv6 restore script:\n%s", backend.calls[1].args[0])
}
}
func TestIptablesForwardLifecycleUsesSingleRestoreScript(t *testing.T) {
script := buildIptablesForwardLifecycleScript(map[string]string{
iptables_helper.NatTab: "-N " + forwarding.ChainPreRouting + "\n-A PREROUTING -j " + forwarding.ChainPreRouting,
iptables_helper.FilterTab: "",
}, true)
if strings.Count(script, "*nat\n") != 1 || strings.Count(script, "*filter\n") != 1 || strings.Count(script, "COMMIT\n") != 2 {
t.Fatalf("unexpected lifecycle restore transaction:\n%s", script)
}
if strings.Contains(script, "-N "+forwarding.ChainPreRouting+"\n") ||
!strings.Contains(script, "-N "+forwarding.ChainPostRouting+"\n") ||
!strings.Contains(script, "-A FORWARD -j "+forwarding.ChainForward+"\n") {
t.Fatalf("lifecycle restore did not preserve/create the expected chains:\n%s", script)
}
}
func TestIptablesForwardCleanupUsesSingleRestoreScript(t *testing.T) {
script := buildIptablesForwardLifecycleScript(map[string]string{
iptables_helper.NatTab: strings.Join([]string{
"-N " + forwarding.ChainPreRouting,
"-A PREROUTING -j " + forwarding.ChainPreRouting,
"-N " + forwarding.ChainPostRouting,
"-A POSTROUTING -j " + forwarding.ChainPostRouting,
}, "\n"),
iptables_helper.FilterTab: "-N " + forwarding.ChainForward + "\n-A FORWARD -j " + forwarding.ChainForward,
}, false)
for _, line := range []string{
"-D PREROUTING -j " + forwarding.ChainPreRouting,
"-F " + forwarding.ChainPostRouting,
"-X " + forwarding.ChainForward,
} {
if !strings.Contains(script, line+"\n") {
t.Fatalf("cleanup restore is missing %q:\n%s", line, script)
}
}
}
type fileWrite struct {
name string
data string
}
type fakeForwardingSystem struct {
reads map[string][]byte
writes []fileWrite
runs []commandCall
}
func (f *fakeForwardingSystem) ReadFile(name string) ([]byte, error) {
data, ok := f.reads[name]
if !ok {
return nil, os.ErrNotExist
}
return data, nil
}
func (f *fakeForwardingSystem) WriteFile(name string, data []byte, _ os.FileMode) error {
f.writes = append(f.writes, fileWrite{name: name, data: string(data)})
return nil
}
func (f *fakeForwardingSystem) RunWithOptionalSudo(name string, args ...string) error {
f.runs = append(f.runs, commandCall{name: name, args: append([]string(nil), args...)})
return nil
}
func TestIptablesNATEnableReplayAndStatusContract(t *testing.T) {
natStatus := strings.Join([]string{
"-N THIRD_PARTY_DNAT",
"-A PREROUTING -j THIRD_PARTY_DNAT",
"-N " + forwarding.ChainPreRouting,
"-N " + forwarding.ChainPostRouting,
"-A PREROUTING -j " + forwarding.ChainPreRouting,
"-A POSTROUTING -j " + forwarding.ChainPostRouting,
}, "\n")
filterStatus := "-N " + forwarding.ChainForward + "\n-A FORWARD -j " + forwarding.ChainForward + "\n-A FORWARD -j DOCKER-USER\n"
backend := &fakeIptablesBackend{stdout: map[string]string{
"nat -S": natStatus,
"filter -S": filterStatus,
}}
system := &fakeForwardingSystem{reads: map[string][]byte{
"/proc/sys/net/ipv4/ip_forward": []byte("1\n"),
"/etc/sysctl.conf": []byte("net.ipv4.tcp_syncookies = 1\n"),
}}
adapter := &iptablesNATAdapter{provider: "iptables", backend: backend, system: system}
if err := adapter.Enable(); err != nil {
t.Fatal(err)
}
if len(system.writes) != 2 || system.writes[0].name != "/proc/sys/net/ipv4/ip_forward" ||
!strings.Contains(system.writes[1].data, "net.ipv4.ip_forward = 1") {
t.Fatalf("sysctl writes changed: %#v", system.writes)
}
if !reflect.DeepEqual(system.runs, []commandCall{{name: "sysctl", args: []string{"-p"}}}) {
t.Fatalf("sysctl transcript changed: %#v", system.runs)
}
init, bind, err := adapter.InitStatus()
if err != nil {
t.Fatal(err)
}
if !init || !bind {
t.Fatalf("expected initialized and bound, got %v %v", init, bind)
}
system.reads["/proc/sys/net/ipv4/ip_forward"] = []byte("0\n")
init, bind, err = adapter.InitStatus()
if err != nil {
t.Fatal(err)
}
if !init || bind {
t.Fatalf("disabled IP forwarding must preserve initialization without reporting a binding, got %v %v", init, bind)
}
backend.calls = nil
if err := adapter.Replay(); err != nil {
t.Fatal(err)
}
wantLoads := []backendCall{
{method: "load", table: iptables_helper.FilterTab, args: []string{forwarding.ChainForward, forwarding.ForwardFile}},
{method: "load", table: iptables_helper.NatTab, args: []string{forwarding.ChainPreRouting, forwarding.PreRoutingFile}},
{method: "load", table: iptables_helper.NatTab, args: []string{forwarding.ChainPostRouting, forwarding.PostRoutingFile}},
}
if !reflect.DeepEqual(backend.calls, wantLoads) {
t.Fatalf("replay transcript changed: %#v", backend.calls)
}
}
func TestIptablesListParsingContract(t *testing.T) {
stdout := strings.Join([]string{
"1 0 0 DNAT tcp -- eth0 * 0.0.0.0/0 0.0.0.0/0 tcp dpt:8080 to:10.0.0.2:80",
"2 0 0 REDIRECT udp -- * * 0.0.0.0/0 0.0.0.0/0 udp dpts:9000:9001 redir ports 53",
}, "\n")
rules := parseIptablesRules(stdout, forwarding.FamilyIPv4)
want := []forwarding.Rule{
{Num: "1", Family: forwarding.FamilyIPv4, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80", Interface: "eth0"},
{Num: "2", Family: forwarding.FamilyIPv4, Protocol: "udp", Port: "9000-9001", TargetIP: "127.0.0.1", TargetPort: "53", Interface: "*"},
}
if !reflect.DeepEqual(rules, want) {
t.Fatalf("got %#v want %#v", rules, want)
}
}
func TestIptablesListParsingAllowsExtraIPv6MatchColumns(t *testing.T) {
stdout := strings.Join([]string{
"3 0 0 REDIRECT tcp -- * * ::/0 ::/0 tcp dpt:55204 ctstate NEW redir ports 80",
"4 0 0 REDIRECT tcp -- * * ::/0 ::/0 tcp dpt:55205 redir ports",
}, "\n")
rules := parseIptablesRules(stdout, forwarding.FamilyIPv6)
want := []forwarding.Rule{{
Num: "3", Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "55204",
TargetIP: "::1", TargetPort: "80", Interface: "*",
}}
if !reflect.DeepEqual(rules, want) {
t.Fatalf("got %#v want %#v", rules, want)
}
}
func TestIptablesIPv6NATContract(t *testing.T) {
stdout := "1 0 0 DNAT tcp -- eth0 * ::/0 ::/0 tcp dpt:8443 to:[2001:db8::20]:443"
rules := parseIptablesRules(stdout, forwarding.FamilyIPv6)
if len(rules) != 1 || rules[0].Family != forwarding.FamilyIPv6 || rules[0].TargetIP != "2001:db8::20" || rules[0].TargetPort != "443" {
t.Fatalf("unexpected parsed IPv6 rules: %#v", rules)
}
}
func TestEnableIPv4ForwardingReplacesDisabledSetting(t *testing.T) {
content := strings.Join([]string{
"# net.ipv4.ip_forward = 0",
"net.ipv4.ip_forward=0",
"net.ipv4.tcp_syncookies = 1",
}, "\n")
want := strings.Join([]string{
"# net.ipv4.ip_forward = 0",
"net.ipv4.ip_forward = 1",
"net.ipv4.tcp_syncookies = 1",
"",
}, "\n")
if got := enableIPv4Forwarding(content); got != want {
t.Fatalf("got %q want %q", got, want)
}
}
func TestEnableForwardingSysctlsAddsIPv6(t *testing.T) {
content := "net.ipv4.ip_forward = 0\nnet.ipv6.conf.all.forwarding=0\n"
want := "net.ipv4.ip_forward = 1\nnet.ipv6.conf.all.forwarding = 1\n"
if got := enableForwardingSysctls(content, true); got != want {
t.Fatalf("got %q want %q", got, want)
}
}
@@ -134,7 +134,7 @@ func (l *iptablesNATAdapter) Reconcile(rules []forwarding.Rule) error {
forwarding.FamilyIPv6: nil,
}
for _, rule := range rules {
normalized, err := NormalizeRule(rule)
normalized, err := forwarding.NormalizeRule(rule)
if err != nil {
return err
}
@@ -185,7 +185,7 @@ func rebuildNftForwardCommands(rules []forwarding.Rule) ([][]string, error) {
}
}
for _, rule := range rules {
normalized, err := NormalizeRule(rule)
normalized, err := forwarding.NormalizeRule(rule)
if err != nil {
return nil, err
}
@@ -1,137 +0,0 @@
package providers
import (
"strings"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding"
)
func TestNftForwardNamingContract(t *testing.T) {
if nftForwardTable != "nft_1panel_forward" {
t.Fatalf("unexpected nftables forwarding table %q", nftForwardTable)
}
if nftForwardFile != "1panel_forward.nft" {
t.Fatalf("unexpected nftables forwarding rules file %q", nftForwardFile)
}
wantChains := map[string]string{
forwarding.ChainPreRouting: "NFT_1PANEL_PREROUTING",
forwarding.ChainPostRouting: "NFT_1PANEL_POSTROUTING",
forwarding.ChainForward: "NFT_1PANEL_FORWARD",
}
for logical, want := range wantChains {
if got := nftForwardChain(logical); got != want {
t.Fatalf("nftForwardChain(%q) = %q, want %q", logical, got, want)
}
}
}
func TestNormalizeNftForwardRuleRejectsScriptTokens(t *testing.T) {
base := forwarding.Rule{Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80"}
tests := []struct {
name string
mutate func(*forwarding.Rule)
}{
{name: "source port", mutate: func(rule *forwarding.Rule) { rule.Port = "8080\nflush ruleset" }},
{name: "target address", mutate: func(rule *forwarding.Rule) { rule.TargetIP = "10.0.0.2; flush ruleset" }},
{name: "target port", mutate: func(rule *forwarding.Rule) { rule.TargetPort = "80; flush ruleset" }},
{name: "interface", mutate: func(rule *forwarding.Rule) { rule.Interface = "eth0;flush" }},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
rule := base
test.mutate(&rule)
if _, err := NormalizeRule(rule); err == nil {
t.Fatal("expected invalid nftables forwarding rule")
}
})
}
}
func TestNormalizeForwardRuleTreatsWildcardInterfaceAsAll(t *testing.T) {
rule, err := NormalizeRule(forwarding.Rule{
Protocol: "udp", Port: "9000-9001", TargetPort: "53", Interface: "*",
})
if err != nil {
t.Fatalf("normalize wildcard forwarding interface: %v", err)
}
if rule.Interface != "" {
t.Fatalf("wildcard forwarding interface normalized to %q, want empty", rule.Interface)
}
}
func TestRebuildNftForwardCommandsUseArguments(t *testing.T) {
commands, err := rebuildNftForwardCommands([]forwarding.Rule{{
Protocol: "tcp", Port: "08080", TargetIP: "10.0.0.2", TargetPort: "080", Interface: "eth0",
}})
if err != nil {
t.Fatalf("build commands: %v", err)
}
if len(commands) != 10 {
t.Fatalf("unexpected command count: %d", len(commands))
}
joined := strings.Join(commands[6], " ")
if !strings.Contains(joined, "add rule ip nft_1panel_forward NFT_1PANEL_PREROUTING") ||
!strings.Contains(joined, `iifname "eth0"`) || !strings.Contains(joined, "dport 8080") ||
!strings.Contains(joined, "dnat to 10.0.0.2:80") {
t.Fatalf("unexpected prerouting command: %q", joined)
}
}
func TestNftForwardCommandsBuildSingleBatchScript(t *testing.T) {
commands, err := rebuildNftForwardCommands([]forwarding.Rule{{
Protocol: "tcp", Port: "8080", TargetIP: "127.0.0.1", TargetPort: "80",
}})
if err != nil {
t.Fatal(err)
}
script, err := nftCommandsScript(commands)
if err != nil {
t.Fatal(err)
}
if lines := strings.Count(script, "\n"); lines != len(commands) {
t.Fatalf("batch script has %d lines, want %d:\n%s", lines, len(commands), script)
}
if !strings.Contains(script, "flush chain ip nft_1panel_forward NFT_1PANEL_PREROUTING\n") ||
!strings.Contains(script, "add rule ip nft_1panel_forward NFT_1PANEL_PREROUTING meta l4proto tcp tcp dport 8080 redirect to :80") {
t.Fatalf("unexpected nftables batch script:\n%s", script)
}
if _, err := nftCommandsScript([][]string{{"add", "rule\nflush ruleset"}}); err == nil {
t.Fatal("batch script accepted a newline token")
}
}
func TestRebuildNftIPv6ForwardCommands(t *testing.T) {
commands, err := rebuildNftForwardCommands([]forwarding.Rule{{
Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "8443", TargetIP: "2001:db8::20", TargetPort: "443", Interface: "eth0",
}})
if err != nil {
t.Fatalf("build IPv6 commands: %v", err)
}
if len(commands) != 10 {
t.Fatalf("unexpected command count: %d", len(commands))
}
preRouting := strings.Join(commands[6], " ")
if !strings.Contains(preRouting, "add rule ip6 nft_1panel_forward NFT_1PANEL_PREROUTING") ||
!strings.Contains(preRouting, "dnat to [2001:db8::20]:443") {
t.Fatalf("unexpected IPv6 prerouting command: %q", preRouting)
}
forward := strings.Join(commands[8], " ")
if !strings.Contains(forward, "ip6 daddr 2001:db8::20") {
t.Fatalf("unexpected IPv6 forward command: %q", forward)
}
}
func TestNormalizeForwardRuleEnforcesAddressFamily(t *testing.T) {
ipv6, err := NormalizeRule(forwarding.Rule{Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "8080", TargetIP: "2001:db8::2", TargetPort: "80"})
if err != nil || ipv6.TargetIP != "2001:db8::2" {
t.Fatalf("normalize IPv6 rule = %#v, %v", ipv6, err)
}
if _, err := NormalizeRule(forwarding.Rule{Family: forwarding.FamilyIPv6, Protocol: "tcp", Port: "8080", TargetIP: "10.0.0.2", TargetPort: "80"}); err == nil {
t.Fatal("expected an IPv4 target to be rejected for an IPv6 rule")
}
loopback, err := NormalizeRule(forwarding.Rule{Family: forwarding.FamilyIPv6, Protocol: "udp", Port: "5353", TargetPort: "53"})
if err != nil || loopback.TargetIP != "::1" {
t.Fatalf("IPv6 loopback normalization = %#v, %v", loopback, err)
}
}
@@ -1,82 +0,0 @@
package providers
import (
"fmt"
"net/netip"
"regexp"
"strconv"
"strings"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding"
)
var forwardInterfacePattern = regexp.MustCompile(`^[A-Za-z0-9_.:@-]{1,15}$`)
func NormalizeRule(rule forwarding.Rule) (forwarding.Rule, error) {
rule.Family = strings.ToLower(strings.TrimSpace(rule.Family))
if rule.Family == "" {
rule.Family = forwarding.FamilyIPv4
}
if rule.Family != forwarding.FamilyIPv4 && rule.Family != forwarding.FamilyIPv6 {
return forwarding.Rule{}, fmt.Errorf("unsupported forwarding family %q", rule.Family)
}
rule.Protocol = strings.ToLower(strings.TrimSpace(rule.Protocol))
if rule.Protocol != "tcp" && rule.Protocol != "udp" {
return forwarding.Rule{}, fmt.Errorf("unsupported forwarding protocol %q", rule.Protocol)
}
var err error
if rule.Port, err = normalizeForwardPort(rule.Port); err != nil {
return forwarding.Rule{}, fmt.Errorf("invalid forwarding port: %w", err)
}
if rule.TargetPort, err = normalizeForwardPort(rule.TargetPort); err != nil {
return forwarding.Rule{}, fmt.Errorf("invalid forwarding target port: %w", err)
}
rule.TargetIP = strings.TrimSpace(rule.TargetIP)
if rule.TargetIP == "" || strings.EqualFold(rule.TargetIP, "localhost") {
if rule.Family == forwarding.FamilyIPv6 {
rule.TargetIP = "::1"
} else {
rule.TargetIP = "127.0.0.1"
}
}
address, err := netip.ParseAddr(rule.TargetIP)
if err == nil {
address = address.Unmap()
}
if err != nil || (rule.Family == forwarding.FamilyIPv4) != address.Is4() {
return forwarding.Rule{}, fmt.Errorf("invalid %s forwarding target %q", rule.Family, rule.TargetIP)
}
rule.TargetIP = address.String()
rule.Interface = strings.TrimSpace(rule.Interface)
if rule.Interface == "all" || rule.Interface == "*" {
rule.Interface = ""
}
if rule.Interface != "" && !forwardInterfacePattern.MatchString(rule.Interface) {
return forwarding.Rule{}, fmt.Errorf("invalid forwarding interface %q", rule.Interface)
}
return rule, nil
}
func normalizeForwardPort(value string) (string, error) {
parts := strings.Split(strings.TrimSpace(value), "-")
if len(parts) < 1 || len(parts) > 2 {
return "", fmt.Errorf("invalid port range %q", value)
}
ports := make([]int, len(parts))
for index, part := range parts {
port, err := strconv.Atoi(strings.TrimSpace(part))
if err != nil || port < 1 || port > 65535 {
return "", fmt.Errorf("invalid port %q", part)
}
ports[index] = port
}
if len(ports) == 2 {
if ports[0] > ports[1] {
return "", fmt.Errorf("descending port range %q", value)
}
if ports[0] != ports[1] {
return strconv.Itoa(ports[0]) + "-" + strconv.Itoa(ports[1]), nil
}
}
return strconv.Itoa(ports[0]), nil
}
@@ -1,12 +0,0 @@
package iptables_helper
import "testing"
func TestHasBaseChainBinding(t *testing.T) {
if hasBaseChainBinding("-P INPUT ACCEPT") {
t.Fatal("input policy was treated as a 1Panel binding")
}
if !hasBaseChainBinding("-A INPUT -j " + BasicAfterChain) {
t.Fatal("partial iptables binding was not detected")
}
}
@@ -1,214 +0,0 @@
package iptables_helper
import (
"errors"
"os"
"path/filepath"
"strings"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall"
)
func TestBuildBaseChainsRestoreScriptBatchesPersistedRules(t *testing.T) {
dir := t.TempDir()
files := map[string]string{
BasicBeforeFileName: "-A " + BasicBeforeChain + " -i lo -j ACCEPT\n",
BasicFileName: "-A " + BasicChain + " -p tcp --dport 8080 -j ACCEPT\n",
BasicAfterFileName: "-A " + BasicAfterChain + " -p tcp -j DROP\n",
}
for name, content := range files {
if err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o600); err != nil {
t.Fatal(err)
}
}
script, err := buildBaseChainsRestoreScript(dir, "9443", false)
if err != nil {
t.Fatal(err)
}
for _, expected := range []string{
"*filter\n",
"-F " + BasicBeforeChain + "\n",
"-F " + BasicChain + "\n",
"-F " + BasicAfterChain + "\n",
files[BasicBeforeFileName],
files[BasicFileName],
files[BasicAfterFileName],
"-A " + BasicBeforeChain + " -p tcp -m tcp --dport 9443 -j ACCEPT\n",
"COMMIT\n",
} {
if !strings.Contains(script, expected) {
t.Fatalf("restore script does not contain %q:\n%s", expected, script)
}
}
if strings.Count(script, "COMMIT\n") != 1 {
t.Fatalf("restore script is not a single batch:\n%s", script)
}
}
func TestBuildBaseChainsRestoreScriptUsesIPv6Files(t *testing.T) {
dir := t.TempDir()
want := "-A " + BasicChain + " -p ipv6-icmp -j ACCEPT\n"
if err := os.WriteFile(filepath.Join(dir, IPv6FileName(BasicFileName)), []byte(want), 0o600); err != nil {
t.Fatal(err)
}
script, err := buildBaseChainsRestoreScript(dir, "9443", true)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(script, want) {
t.Fatalf("IPv6 persisted rule was not restored:\n%s", script)
}
}
func TestBuildRequiredPortsRestoreScriptBatchesAddsAndDeletes(t *testing.T) {
desired := []firewall.PortWhitelist{{Protocol: "tcp", Port: "22"}, {Protocol: "udp", Port: "53"}}
before := []FilterRules{
{Chain: BasicBeforeChain, Protocol: "tcp", DstPort: "22", Strategy: "accept"},
{Chain: BasicBeforeChain, Protocol: "tcp", DstPort: "22", Strategy: "accept"},
{Chain: BasicBeforeChain, Protocol: "tcp", DstPort: "80", Strategy: "accept"},
}
after := []FilterRules{{Chain: BasicAfterChain, Protocol: "udp", DstPort: "5353", Strategy: "accept"}}
script := buildRequiredPortsRestoreScript(desired, before, after, "", "", true)
for _, line := range []string{
"-D 1PANEL_BASIC_BEFORE -p tcp -m tcp --dport 22 -j ACCEPT",
"-D 1PANEL_BASIC_BEFORE -p tcp -m tcp --dport 80 -j ACCEPT",
"-D 1PANEL_BASIC_AFTER -p udp -m udp --dport 5353 -j ACCEPT",
"-A 1PANEL_BASIC_BEFORE -p udp -m udp --dport 53 -j ACCEPT",
"-A 1PANEL_BASIC_AFTER -p tcp -j DROP",
"-A 1PANEL_BASIC_AFTER -p udp -j DROP",
} {
if !strings.Contains(script, line+"\n") {
t.Fatalf("batch script is missing %q:\n%s", line, script)
}
}
if got := strings.Count(script, "--dport 22 -j ACCEPT"); got != 1 {
t.Fatalf("duplicate desired port was not removed exactly once: count=%d\n%s", got, script)
}
if !strings.HasPrefix(script, "*filter\n") || !strings.HasSuffix(script, "COMMIT\n") {
t.Fatalf("invalid restore transaction:\n%s", script)
}
}
func TestBuildRequiredPortsRestoreScriptSkipsUnchangedState(t *testing.T) {
desired := []firewall.PortWhitelist{{Protocol: "tcp", Port: "22"}}
before := []FilterRules{{Chain: BasicBeforeChain, Protocol: "tcp", DstPort: "22", Strategy: "accept"}}
if script := buildRequiredPortsRestoreScript(desired, before, nil, "", "", false); script != "" {
t.Fatalf("unchanged required ports generated a restore transaction:\n%s", script)
}
}
func TestBuildBaseChainBindingsRestoreScriptRebindsInOneTransaction(t *testing.T) {
output := strings.Join([]string{
"-A INPUT -j " + BasicBeforeChain,
"-A INPUT -j external",
"-A INPUT -j " + BasicAfterChain,
}, "\n")
script := buildBaseChainBindingsRestoreScript(output, true)
for _, line := range []string{
"-D INPUT -j " + BasicBeforeChain,
"-D INPUT -j " + BasicAfterChain,
"-I INPUT 1 -j " + BasicBeforeChain,
"-I INPUT 2 -j " + BasicChain,
"-I INPUT 3 -j " + BasicAfterChain,
} {
if !strings.Contains(script, line+"\n") {
t.Fatalf("binding batch is missing %q:\n%s", line, script)
}
}
if strings.Contains(script, "external") || strings.Count(script, "COMMIT\n") != 1 {
t.Fatalf("binding batch modified an external rule or is not atomic:\n%s", script)
}
}
func TestLoadInitStatusUsesIPv6BaselineWithoutIPv4TerminalRules(t *testing.T) {
output := strings.Join([]string{
"-N " + BasicBeforeChain,
"-N " + BasicChain,
"-N " + BasicAfterChain,
"-A " + BasicBeforeChain + " -i lo -m comment --comment \"Loopback Whitelist\" -j ACCEPT",
"-A " + BasicBeforeChain + " -m conntrack --ctstate RELATED,ESTABLISHED -m comment --comment \"ESTABLISHED Whitelist\" -j ACCEPT",
"-A " + InputChain + " -j " + BasicBeforeChain,
"-A " + InputChain + " -j " + BasicChain,
"-A " + InputChain + " -j " + BasicAfterChain,
}, "\n")
runner := func(string, ...string) (string, error) { return output, nil }
initialized, bound, err := loadInitStatus("base", runner, false)
if err != nil || !initialized || !bound {
t.Fatalf("IPv6 baseline status = initialized:%v bound:%v err:%v", initialized, bound, err)
}
initialized, bound, err = loadInitStatus("base", runner, true)
if err != nil || initialized || bound {
t.Fatalf("IPv4 status ignored missing terminal rules: initialized:%v bound:%v err:%v", initialized, bound, err)
}
}
func TestRepairIPv6BaseChainsOnlyBindsInitializedChains(t *testing.T) {
bindCalls, ensureCalls := 0, 0
err := repairIPv6BaseChains(true, false, func() error {
bindCalls++
return nil
}, func() error {
ensureCalls++
return nil
})
if err != nil {
t.Fatalf("repair initialized IPv6 base chains: %v", err)
}
if bindCalls != 1 || ensureCalls != 0 {
t.Fatalf("repair calls = bind:%d ensure:%d, want bind:1 ensure:0", bindCalls, ensureCalls)
}
}
func TestRepairIPv6BaseChainsRebuildsMissingChains(t *testing.T) {
bindCalls, ensureCalls := 0, 0
err := repairIPv6BaseChains(false, false, func() error {
bindCalls++
return nil
}, func() error {
ensureCalls++
return nil
})
if err != nil {
t.Fatalf("repair missing IPv6 base chains: %v", err)
}
if bindCalls != 0 || ensureCalls != 1 {
t.Fatalf("repair calls = bind:%d ensure:%d, want bind:0 ensure:1", bindCalls, ensureCalls)
}
}
func TestRepairIPv6BaseChainsLeavesHealthyChainsUnchanged(t *testing.T) {
bindCalls, ensureCalls := 0, 0
err := repairIPv6BaseChains(true, true, func() error {
bindCalls++
return nil
}, func() error {
ensureCalls++
return nil
})
if err != nil {
t.Fatalf("repair healthy IPv6 base chains: %v", err)
}
if bindCalls != 0 || ensureCalls != 0 {
t.Fatalf("repair calls = bind:%d ensure:%d, want no operation", bindCalls, ensureCalls)
}
}
func TestRepairIPv6BaseChainsPropagatesOperationErrors(t *testing.T) {
wantBindErr := errors.New("bind failed")
if err := repairIPv6BaseChains(true, false, func() error { return wantBindErr }, func() error { return nil }); !errors.Is(err, wantBindErr) {
t.Fatalf("bind error = %v, want %v", err, wantBindErr)
}
wantEnsureErr := errors.New("ensure failed")
if err := repairIPv6BaseChains(false, false, func() error { return nil }, func() error { return wantEnsureErr }); !errors.Is(err, wantEnsureErr) {
t.Fatalf("ensure error = %v, want %v", err, wantEnsureErr)
}
}
func TestLoadFamilyInitStatusRejectsUnknownFamily(t *testing.T) {
initialized, bound, err := LoadFamilyInitStatus("inet", "base")
if err == nil || initialized || bound {
t.Fatalf("unknown family status = initialized:%v bound:%v err:%v", initialized, bound, err)
}
}
-105
View File
@@ -1,105 +0,0 @@
package lifecycle
import (
"fmt"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle/providers"
)
type Client interface {
Name() string
Start() error
Stop() error
Restart() error
Status() (bool, error)
Version() (string, error)
}
// Resetter restores a service-backed firewall to its installation defaults.
// Implementations must leave the firewall disabled after a successful reset.
type Resetter interface {
Reset() error
}
func NewClient() (Client, error) {
runtime, err := DetectRuntime()
if err != nil {
return nil, err
}
return NewClientFor(runtime.Provider)
}
func NewClientFor(provider string) (Client, error) {
switch provider {
case "firewalld":
if !which("firewalld") {
return nil, fmt.Errorf("firewalld is not installed")
}
return providers.NewFirewalld()
case "ufw":
if !which("ufw") {
return nil, fmt.Errorf("ufw is not installed")
}
return providers.NewUFW()
case "iptables":
commands, err := ResolveIptablesCommands()
if err != nil {
return nil, err
}
return providers.NewIptables(commands.IPv4)
case "nftables":
if !which("nft") {
return nil, fmt.Errorf("nftables is not installed")
}
return providers.NewNftables()
default:
return nil, fmt.Errorf("unsupported firewall provider: %s", provider)
}
}
func InstalledProviders() []string {
providers := make([]string, 0, 4)
if which("firewalld") {
providers = append(providers, ProviderFirewalld)
}
if which("ufw") {
providers = append(providers, ProviderUFW)
}
if _, err := ResolveIptablesCommands(); err == nil {
providers = append(providers, ProviderIptables)
}
if which("nft") {
providers = append(providers, ProviderNftables)
}
return providers
}
func NewNetfilterClients() ([]Client, error) {
clients := make([]Client, 0, 2)
if which("nft") {
client, err := providers.NewNftables()
if err != nil {
return nil, err
}
clients = append(clients, client)
}
if commands, err := ResolveIptablesCommands(); err == nil {
client, err := providers.NewIptables(commands.IPv4)
if err != nil {
return nil, err
}
clients = append(clients, client)
}
if len(clients) == 0 {
return nil, fmt.Errorf("no supported forwarding backend detected (iptables/iptables-nft/nftables)")
}
return clients, nil
}
func DetectProvider() (string, error) {
runtime, err := DetectRuntime()
if err != nil {
return "", err
}
return runtime.Provider, nil
}
@@ -1,61 +0,0 @@
package lifecycle
import "testing"
func TestNewNetfilterClientsIgnoreHostFirewallService(t *testing.T) {
original := which
t.Cleanup(func() { which = original })
tests := []struct {
name string
commands map[string]bool
want []string
wantErr bool
}{
{
name: "firewalld host with iptables",
commands: map[string]bool{
"firewalld": true, "iptables": true, "iptables-restore": true,
},
want: []string{ProviderIptables},
},
{
name: "ufw host with iptables nft",
commands: map[string]bool{
"ufw": true, "iptables-nft": true, "iptables-nft-restore": true, "nft": true,
},
want: []string{ProviderNftables, ProviderIptables},
},
{
name: "firewalld host with native nft",
commands: map[string]bool{"firewalld": true, "nft": true},
want: []string{ProviderNftables},
},
{
name: "service without netfilter command backend",
commands: map[string]bool{"firewalld": true},
wantErr: true,
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
which = func(name string) bool { return test.commands[name] }
clients, err := NewNetfilterClients()
if (err != nil) != test.wantErr {
t.Fatalf("NewNetfilterClients() error = %v, wantErr %v", err, test.wantErr)
}
if test.wantErr {
return
}
if len(clients) != len(test.want) {
t.Fatalf("NewNetfilterClients() returned %d clients, want %d", len(clients), len(test.want))
}
for index, client := range clients {
if client.Name() != test.want[index] {
t.Fatalf("NewNetfilterClients()[%d] = %q, want %q", index, client.Name(), test.want[index])
}
}
})
}
}
+230
View File
@@ -0,0 +1,230 @@
package lifecycle
import (
"errors"
"fmt"
"sync"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/utils/cmd"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/lifecycle/providers"
)
const (
ProviderFirewalld = constant.FirewallProviderFirewalld
ProviderUFW = constant.FirewallProviderUFW
ProviderIptables = constant.FirewallProviderIptables
ProviderNftables = constant.FirewallProviderNftables
)
var ErrNotInstalled = errors.New("is not installed")
type IptablesCommands struct {
IPv4 string
IPv6 string
Restore4 string
Restore6 string
}
func (c IptablesCommands) IPv6Available() bool {
return c.IPv6 != "" && c.Restore6 != ""
}
type Runtime struct {
Provider string
Iptables IptablesCommands
}
var which = cmd.Which
func DetectRuntime() (Runtime, error) {
hasFirewalld := which("firewalld")
hasUFW := which("ufw")
if hasFirewalld && hasUFW {
return Runtime{}, errors.New("it is detected that the system has both firewalld and ufw services. To avoid conflicts, please uninstall and try again")
}
if hasFirewalld {
return Runtime{Provider: ProviderFirewalld}, nil
}
if hasUFW {
return Runtime{Provider: ProviderUFW}, nil
}
if commands, ok := detectIptablesCommands(""); ok {
return Runtime{Provider: ProviderIptables, Iptables: commands}, nil
}
if commands, ok := detectIptablesCommands("-nft"); ok {
return Runtime{Provider: ProviderIptables, Iptables: commands}, nil
}
if which("nft") {
return Runtime{Provider: ProviderNftables}, nil
}
return Runtime{}, errors.New("no system firewall service detected (firewalld/ufw/iptables/iptables-nft/nft), please check and try again")
}
func detectIptablesCommands(suffix string) (IptablesCommands, bool) {
ipv4 := "iptables" + suffix
restore4 := "iptables" + suffix + "-restore"
if !which(ipv4) || !which(restore4) {
return IptablesCommands{}, false
}
commands := IptablesCommands{IPv4: ipv4, Restore4: restore4}
ipv6 := "ip6tables" + suffix
restore6 := "ip6tables" + suffix + "-restore"
if which(ipv6) && which(restore6) {
commands.IPv6 = ipv6
commands.Restore6 = restore6
}
return commands, true
}
func ResolveIptablesCommands() (IptablesCommands, error) {
if commands, ok := detectIptablesCommands(""); ok {
return commands, nil
}
if commands, ok := detectIptablesCommands("-nft"); ok {
return commands, nil
}
return IptablesCommands{}, fmt.Errorf("no complete iptables command family is available")
}
type Client interface {
Name() string
Start() error
Stop() error
Restart() error
Status() (bool, error)
Version() (string, error)
}
// Resetter restores a service-backed firewall to its installation defaults.
// Implementations must leave the firewall disabled after a successful reset.
type Resetter interface {
Reset() error
}
func NewClient() (Client, error) {
runtime, err := DetectRuntime()
if err != nil {
return nil, err
}
return NewClientFor(runtime.Provider)
}
func NewClientFor(provider string) (Client, error) {
switch provider {
case "firewalld":
if !which("firewalld") {
return nil, fmt.Errorf("firewalld %w", ErrNotInstalled)
}
return providers.NewFirewalld()
case "ufw":
if !which("ufw") {
return nil, fmt.Errorf("ufw %w", ErrNotInstalled)
}
return providers.NewUFW()
case "iptables":
commands, err := ResolveIptablesCommands()
if err != nil {
return nil, err
}
return providers.NewIptables(commands.IPv4)
case "nftables":
if !which("nft") {
return nil, fmt.Errorf("nftables %w", ErrNotInstalled)
}
return providers.NewNftables()
default:
return nil, fmt.Errorf("unsupported firewall provider: %s", provider)
}
}
func InstalledProviders() []string {
providers := make([]string, 0, 4)
if which("firewalld") {
providers = append(providers, ProviderFirewalld)
}
if which("ufw") {
providers = append(providers, ProviderUFW)
}
if _, err := ResolveIptablesCommands(); err == nil {
providers = append(providers, ProviderIptables)
}
if which("nft") {
providers = append(providers, ProviderNftables)
}
return providers
}
func NewNetfilterClients() ([]Client, error) {
clients := make([]Client, 0, 2)
if which("nft") {
client, err := providers.NewNftables()
if err != nil {
return nil, err
}
clients = append(clients, client)
}
if commands, err := ResolveIptablesCommands(); err == nil {
client, err := providers.NewIptables(commands.IPv4)
if err != nil {
return nil, err
}
clients = append(clients, client)
}
if len(clients) == 0 {
return nil, fmt.Errorf("no supported forwarding backend detected (iptables/iptables-nft/nftables)")
}
return clients, nil
}
func DetectProvider() (string, error) {
runtime, err := DetectRuntime()
if err != nil {
return "", err
}
return runtime.Provider, nil
}
type State struct {
Name string
IsActive bool
}
type Status struct {
State
Version string
}
func LoadState(client Client) (State, error) {
state := State{Name: client.Name()}
var err error
state.IsActive, err = client.Status()
return state, err
}
func LoadStatus(client Client) (Status, error) {
status := Status{Version: "-"}
var state State
var version string
var stateErr, versionErr error
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
state, stateErr = LoadState(client)
}()
go func() {
defer wg.Done()
version, versionErr = client.Version()
}()
wg.Wait()
status.State = state
if stateErr != nil {
return status, errors.Join(stateErr, versionErr)
}
if !status.IsActive {
return status, nil
}
status.Version = version
return status, versionErr
}
@@ -1,48 +0,0 @@
package lifecycle
import (
"errors"
"testing"
)
type operatorTestClient struct {
name string
started bool
stopped bool
}
func (f *operatorTestClient) Name() string { return f.name }
func (f *operatorTestClient) Start() error { f.started = true; return nil }
func (f *operatorTestClient) Stop() error { f.stopped = true; return nil }
func (f *operatorTestClient) Restart() error { return nil }
func (f *operatorTestClient) Status() (bool, error) { return true, nil }
func (f *operatorTestClient) Version() (string, error) { return "test", nil }
func TestOperatorDelegatesLifecycleAndPreparesStart(t *testing.T) {
client := &operatorTestClient{name: "iptables"}
operator := NewOperator(client)
prepared := false
if err := operator.Operate("start", false, func(got Client) error {
prepared = got == client
return nil
}); err != nil {
t.Fatal(err)
}
if !client.started || !prepared {
t.Fatalf("start was not fully coordinated: started=%v prepared=%v", client.started, prepared)
}
}
func TestOperatorKeepsFirewallRunningWhenPostStartPreparationFails(t *testing.T) {
client := &operatorTestClient{name: "firewalld"}
wantErr := errors.New("sync accepted ports")
err := NewOperator(client).Operate(OperationStart, false, func(Client) error {
return wantErr
})
if !errors.Is(err, wantErr) {
t.Fatalf("start returned error %v, want %v", err, wantErr)
}
if !client.started || client.stopped {
t.Fatalf("post-start failure changed firewall state: started=%v stopped=%v", client.started, client.stopped)
}
}
@@ -1,43 +0,0 @@
package lifecycle
import (
"os"
"path/filepath"
"testing"
)
func TestDetectProvider(t *testing.T) {
tests := []struct {
name string
executables []string
want string
wantErr bool
}{
{name: "none", wantErr: true},
{name: "iptables", executables: []string{"iptables", "iptables-restore"}, want: "iptables"},
{name: "iptables-nft", executables: []string{"iptables-nft", "iptables-nft-restore"}, want: "iptables"},
{name: "nftables", executables: []string{"nft"}, want: "nftables"},
{name: "ufw", executables: []string{"iptables", "iptables-restore", "ufw"}, want: "ufw"},
{name: "firewalld", executables: []string{"iptables", "iptables-restore", "firewalld"}, want: "firewalld"},
{name: "conflict", executables: []string{"firewalld", "ufw"}, wantErr: true},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
directory := t.TempDir()
for _, executable := range test.executables {
name := filepath.Join(directory, executable)
if err := os.WriteFile(name, []byte("#!/bin/sh\n"), 0755); err != nil {
t.Fatal(err)
}
}
t.Setenv("PATH", directory)
got, err := DetectProvider()
if (err != nil) != test.wantErr {
t.Fatalf("DetectProvider() error = %v, wantErr %v", err, test.wantErr)
}
if got != test.want {
t.Fatalf("DetectProvider() = %q, want %q", got, test.want)
}
})
}
}
@@ -1,107 +0,0 @@
package providers
import (
"errors"
"os"
"path/filepath"
"testing"
)
func TestFirewalldStoppedRecognizesNormalInactiveResult(t *testing.T) {
for _, test := range []struct {
stdout string
err error
}{
{stdout: "not running\n"},
{err: errors.New("stderr: FirewallD is not running, exit status 252")},
} {
if !firewalldStopped(test.stdout, test.err) {
t.Fatalf("expected stopped result for stdout=%q err=%v", test.stdout, test.err)
}
}
}
func TestFirewalldStoppedKeepsUnexpectedFailures(t *testing.T) {
if firewalldStopped("", errors.New("permission denied")) {
t.Fatal("unexpected command failures must not be treated as an inactive firewall")
}
}
func TestReplaceFirewalldConfigCreatesCleanConfigurationAndBackup(t *testing.T) {
root := t.TempDir()
configDir := filepath.Join(root, "firewalld")
backupDir := filepath.Join(root, "firewalld.backup")
originalZone := filepath.Join(configDir, "zones", "custom.xml")
if err := os.MkdirAll(filepath.Dir(originalZone), 0750); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(originalZone, []byte("custom"), 0600); err != nil {
t.Fatal(err)
}
prepared := false
rollback, err := replaceFirewalldConfig(
configDir,
backupDir,
func(path string) error {
prepared = path == configDir
return nil
},
func() error {
for _, directory := range firewalldConfigSubdirectories {
if info, err := os.Stat(filepath.Join(configDir, directory)); err != nil || !info.IsDir() {
t.Fatalf("expected clean %s directory, info=%v err=%v", directory, info, err)
}
}
if _, err := os.Stat(originalZone); !os.IsNotExist(err) {
t.Fatalf("custom zone must not remain in clean configuration: %v", err)
}
return nil
},
)
if err != nil {
t.Fatalf("replace firewalld configuration: %v", err)
}
if !prepared {
t.Fatal("expected clean configuration to be prepared")
}
if content, err := os.ReadFile(filepath.Join(backupDir, "zones", "custom.xml")); err != nil || string(content) != "custom" {
t.Fatalf("expected original configuration in backup, content=%q err=%v", content, err)
}
if err := rollback(); err != nil {
t.Fatalf("rollback firewalld configuration: %v", err)
}
if content, err := os.ReadFile(originalZone); err != nil || string(content) != "custom" {
t.Fatalf("expected original configuration after rollback, content=%q err=%v", content, err)
}
}
func TestReplaceFirewalldConfigRollsBackValidationFailure(t *testing.T) {
root := t.TempDir()
configDir := filepath.Join(root, "firewalld")
backupDir := filepath.Join(root, "firewalld.backup")
originalConfig := filepath.Join(configDir, "firewalld.conf")
if err := os.Mkdir(configDir, 0750); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(originalConfig, []byte("DefaultZone=custom\n"), 0600); err != nil {
t.Fatal(err)
}
rollback, err := replaceFirewalldConfig(
configDir,
backupDir,
nil,
func() error { return errors.New("invalid defaults") },
)
if err == nil || rollback != nil {
t.Fatalf("expected validation failure with automatic rollback, hasRollback=%t err=%v", rollback != nil, err)
}
if content, readErr := os.ReadFile(originalConfig); readErr != nil || string(content) != "DefaultZone=custom\n" {
t.Fatalf("expected original configuration after failed validation, content=%q err=%v", content, readErr)
}
if _, statErr := os.Stat(backupDir); !os.IsNotExist(statErr) {
t.Fatalf("backup must be restored after failed validation: %v", statErr)
}
}
-84
View File
@@ -1,84 +0,0 @@
package lifecycle
import (
"errors"
"fmt"
"github.com/1Panel-dev/1Panel/agent/constant"
"github.com/1Panel-dev/1Panel/agent/utils/cmd"
)
const (
ProviderFirewalld = constant.FirewallProviderFirewalld
ProviderUFW = constant.FirewallProviderUFW
ProviderIptables = constant.FirewallProviderIptables
ProviderNftables = constant.FirewallProviderNftables
)
type IptablesCommands struct {
IPv4 string
IPv6 string
Restore4 string
Restore6 string
}
func (c IptablesCommands) IPv6Available() bool {
return c.IPv6 != "" && c.Restore6 != ""
}
type Runtime struct {
Provider string
Iptables IptablesCommands
}
var which = cmd.Which
func DetectRuntime() (Runtime, error) {
hasFirewalld := which("firewalld")
hasUFW := which("ufw")
if hasFirewalld && hasUFW {
return Runtime{}, errors.New("it is detected that the system has both firewalld and ufw services. To avoid conflicts, please uninstall and try again")
}
if hasFirewalld {
return Runtime{Provider: ProviderFirewalld}, nil
}
if hasUFW {
return Runtime{Provider: ProviderUFW}, nil
}
if commands, ok := detectIptablesCommands(""); ok {
return Runtime{Provider: ProviderIptables, Iptables: commands}, nil
}
if commands, ok := detectIptablesCommands("-nft"); ok {
return Runtime{Provider: ProviderIptables, Iptables: commands}, nil
}
if which("nft") {
return Runtime{Provider: ProviderNftables}, nil
}
return Runtime{}, errors.New("no system firewall service detected (firewalld/ufw/iptables/iptables-nft/nft), please check and try again")
}
func detectIptablesCommands(suffix string) (IptablesCommands, bool) {
ipv4 := "iptables" + suffix
restore4 := "iptables" + suffix + "-restore"
if !which(ipv4) || !which(restore4) {
return IptablesCommands{}, false
}
commands := IptablesCommands{IPv4: ipv4, Restore4: restore4}
ipv6 := "ip6tables" + suffix
restore6 := "ip6tables" + suffix + "-restore"
if which(ipv6) && which(restore6) {
commands.IPv6 = ipv6
commands.Restore6 = restore6
}
return commands, true
}
func ResolveIptablesCommands() (IptablesCommands, error) {
if commands, ok := detectIptablesCommands(""); ok {
return commands, nil
}
if commands, ok := detectIptablesCommands("-nft"); ok {
return commands, nil
}
return IptablesCommands{}, fmt.Errorf("no complete iptables command family is available")
}
@@ -1,41 +0,0 @@
package lifecycle
import "testing"
func TestDetectRuntimePriority(t *testing.T) {
original := which
t.Cleanup(func() { which = original })
tests := []struct {
name string
commands map[string]bool
provider string
executable string
}{
{name: "default iptables", commands: map[string]bool{"iptables": true, "iptables-restore": true, "iptables-nft": true, "iptables-nft-restore": true, "nft": true}, provider: ProviderIptables, executable: "iptables"},
{name: "explicit iptables nft", commands: map[string]bool{"iptables-nft": true, "iptables-nft-restore": true, "nft": true}, provider: ProviderIptables, executable: "iptables-nft"},
{name: "native nft", commands: map[string]bool{"nft": true}, provider: ProviderNftables},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
which = func(name string) bool { return test.commands[name] }
runtime, err := DetectRuntime()
if err != nil {
t.Fatalf("detect: %v", err)
}
if runtime.Provider != test.provider || runtime.Iptables.IPv4 != test.executable {
t.Fatalf("unexpected runtime: %#v", runtime)
}
})
}
}
func TestDetectRuntimeRequiresRestoreCommand(t *testing.T) {
original := which
t.Cleanup(func() { which = original })
which = func(name string) bool { return name == "iptables" || name == "nft" }
runtime, err := DetectRuntime()
if err != nil || runtime.Provider != ProviderNftables {
t.Fatalf("incomplete iptables family should fall back to nft: runtime=%#v err=%v", runtime, err)
}
}
-50
View File
@@ -1,50 +0,0 @@
package lifecycle
import (
"errors"
"sync"
)
type State struct {
Name string
IsActive bool
}
type Status struct {
State
Version string
}
func LoadState(client Client) (State, error) {
state := State{Name: client.Name()}
var err error
state.IsActive, err = client.Status()
return state, err
}
func LoadStatus(client Client) (Status, error) {
status := Status{Version: "-"}
var state State
var version string
var stateErr, versionErr error
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
state, stateErr = LoadState(client)
}()
go func() {
defer wg.Done()
version, versionErr = client.Version()
}()
wg.Wait()
status.State = state
if stateErr != nil {
return status, errors.Join(stateErr, versionErr)
}
if !status.IsActive {
return status, nil
}
status.Version = version
return status, versionErr
}
@@ -1,55 +0,0 @@
package lifecycle
import (
"errors"
"testing"
)
type statusTestClient struct {
active bool
statusErr error
versionErr error
versioned *bool
}
func (statusTestClient) Name() string { return "ufw" }
func (statusTestClient) Start() error { return nil }
func (statusTestClient) Stop() error { return nil }
func (statusTestClient) Restart() error { return nil }
func (c statusTestClient) Status() (bool, error) { return c.active, c.statusErr }
func (c statusTestClient) Version() (string, error) {
if c.versioned != nil {
*c.versioned = true
}
return "1.0", c.versionErr
}
func TestLoadStatusAggregatesClientState(t *testing.T) {
status, err := LoadStatus(statusTestClient{active: true})
if err != nil {
t.Fatal(err)
}
if status.Name != "ufw" || status.Version != "1.0" || !status.IsActive {
t.Fatalf("unexpected status: %#v", status)
}
}
func TestLoadStatusReturnsClientErrors(t *testing.T) {
statusErr := errors.New("status failed")
versionErr := errors.New("version failed")
_, err := LoadStatus(statusTestClient{active: true, statusErr: statusErr, versionErr: versionErr})
if !errors.Is(err, statusErr) || !errors.Is(err, versionErr) {
t.Fatalf("got %v, want joined status and version errors", err)
}
}
func TestLoadStatusIgnoresVersionFailureForInactiveFirewall(t *testing.T) {
versioned := false
status, err := LoadStatus(statusTestClient{versionErr: errors.New("firewall is stopped"), versioned: &versioned})
if err != nil {
t.Fatal(err)
}
if status.Name != "ufw" || status.Version != "-" || status.IsActive || !versioned {
t.Fatalf("unexpected inactive status: %#v versioned=%v", status, versioned)
}
}
@@ -1,66 +0,0 @@
package nftables_helper
import (
"fmt"
"strings"
"time"
"github.com/1Panel-dev/1Panel/agent/utils/cmd"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
const (
TableName = "nft_1panel_filter"
InputChain = "NFT_1PANEL_INPUT"
BasicBeforeChain = "NFT_1PANEL_BASIC_BEFORE"
BasicChain = "NFT_1PANEL_BASIC"
BasicAfterChain = "NFT_1PANEL_BASIC_AFTER"
)
func TableFamily(family filter.Family) string {
if family == filter.FamilyIPv6 {
return "ip6"
}
return "ip"
}
func BasicChains() []string {
return []string{BasicBeforeChain, BasicChain, BasicAfterChain}
}
func run(args ...string) (string, error) {
return cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudoAndStdout("nft", args...)
}
func runCommand(args ...string) error {
return cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudo("nft", args...)
}
func runBatch(commands ...[]string) error {
script, err := buildBatchScript(commands...)
if err != nil || script == "" {
return err
}
manager := cmd.NewCommandMgr(
cmd.WithTimeout(60*time.Second),
cmd.WithStdin(strings.NewReader(script)),
)
return manager.RunWithOptionalSudo("nft", "-f", "-")
}
func buildBatchScript(commands ...[]string) (string, error) {
var script strings.Builder
for _, command := range commands {
if len(command) == 0 {
continue
}
for _, token := range command {
if strings.ContainsAny(token, "\r\n") {
return "", fmt.Errorf("invalid newline in nftables batch command")
}
}
script.WriteString(strings.Join(command, " "))
script.WriteByte('\n')
}
return script.String(), nil
}
@@ -1,51 +0,0 @@
package nftables_helper
import (
"strings"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
func TestNativeNames(t *testing.T) {
tests := []struct {
family filter.Family
tableFamily string
}{
{family: filter.FamilyIPv4, tableFamily: "ip"},
{family: filter.FamilyIPv6, tableFamily: "ip6"},
}
for _, test := range tests {
if got := TableFamily(test.family); got != test.tableFamily {
t.Fatalf("TableFamily(%s) = %q, want %q", test.family, got, test.tableFamily)
}
}
wantChains := []string{"NFT_1PANEL_BASIC_BEFORE", "NFT_1PANEL_BASIC", "NFT_1PANEL_BASIC_AFTER"}
for index, got := range BasicChains() {
if got != wantChains[index] {
t.Fatalf("BasicChains()[%d] = %q, want %q", index, got, wantChains[index])
}
}
}
func TestBuildBatchScriptCombinesCommands(t *testing.T) {
script, err := buildBatchScript(
[]string{"flush", "chain", "ip", TableName, BasicBeforeChain},
[]string{"add", "rule", "ip", TableName, BasicBeforeChain, "tcp", "dport", "443", "accept"},
[]string{"delete", "rule", "ip6", TableName, BasicBeforeChain, "handle", "12"},
)
if err != nil {
t.Fatal(err)
}
if strings.Count(script, "\n") != 3 || !strings.Contains(script, "tcp dport 443 accept\n") ||
!strings.Contains(script, "delete rule ip6 "+TableName+" "+BasicBeforeChain+" handle 12\n") {
t.Fatalf("unexpected nftables batch script:\n%s", script)
}
}
func TestBuildBatchScriptRejectsNewline(t *testing.T) {
if _, err := buildBatchScript([]string{"add", "rule", "ip", TableName, BasicChain, "unsafe\nrule"}); err == nil {
t.Fatal("expected newline validation error")
}
}
@@ -1,45 +0,0 @@
package nftables_helper
import (
"reflect"
"testing"
"github.com/1Panel-dev/1Panel/agent/utils/firewall"
)
func TestHasBaseChainBinding(t *testing.T) {
if hasBaseChainBinding(`chain NFT_1PANEL_INPUT { policy accept; }`) {
t.Fatal("empty input chain was treated as bound")
}
if !hasBaseChainBinding(`jump NFT_1PANEL_BASIC`) {
t.Fatal("partial nftables binding was not detected")
}
}
func TestRequiredPortChangesPreserveExistingRules(t *testing.T) {
existing := []requiredPortRule{
{Key: "22/tcp", Handle: "12"},
{Key: "22/tcp", Handle: "13"},
{Key: "80/tcp", Handle: "14"},
}
desired := []firewall.PortWhitelist{{Port: "22", Protocol: "tcp"}, {Port: "53", Protocol: "udp"}}
missing, stale := requiredPortChanges(existing, desired)
if want := []firewall.PortWhitelist{{Port: "53", Protocol: "udp"}}; !reflect.DeepEqual(missing, want) {
t.Fatalf("missing=%#v, want %#v", missing, want)
}
if want := []string{"13", "14"}; !reflect.DeepEqual(stale, want) {
t.Fatalf("stale=%#v, want %#v", stale, want)
}
}
func TestRequiredPortRules(t *testing.T) {
output := `
tcp dport 22 accept comment "1Panel Port Whitelist" # handle 12
tcp dport 80 accept comment "external" # handle 13
udp dport 53 accept comment "1Panel Port Whitelist" # handle 14
`
want := []requiredPortRule{{Key: "22/tcp", Handle: "12"}, {Key: "53/udp", Handle: "14"}}
if got := requiredPortRules(output); !reflect.DeepEqual(got, want) {
t.Fatalf("rules=%#v, want %#v", got, want)
}
}
@@ -1,11 +1,70 @@
package nftables_helper
import (
"fmt"
"strings"
"time"
"github.com/1Panel-dev/1Panel/agent/utils/cmd"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
const (
TableName = "nft_1panel_filter"
InputChain = "NFT_1PANEL_INPUT"
BasicBeforeChain = "NFT_1PANEL_BASIC_BEFORE"
BasicChain = "NFT_1PANEL_BASIC"
BasicAfterChain = "NFT_1PANEL_BASIC_AFTER"
)
func TableFamily(family filter.Family) string {
if family == filter.FamilyIPv6 {
return "ip6"
}
return "ip"
}
func BasicChains() []string {
return []string{BasicBeforeChain, BasicChain, BasicAfterChain}
}
func run(args ...string) (string, error) {
return cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudoAndStdout("nft", args...)
}
func runCommand(args ...string) error {
return cmd.NewCommandMgr(cmd.WithTimeout(60*time.Second)).RunWithOptionalSudo("nft", args...)
}
func runBatch(commands ...[]string) error {
script, err := buildBatchScript(commands...)
if err != nil || script == "" {
return err
}
manager := cmd.NewCommandMgr(
cmd.WithTimeout(60*time.Second),
cmd.WithStdin(strings.NewReader(script)),
)
return manager.RunWithOptionalSudo("nft", "-f", "-")
}
func buildBatchScript(commands ...[]string) (string, error) {
var script strings.Builder
for _, command := range commands {
if len(command) == 0 {
continue
}
for _, token := range command {
if strings.ContainsAny(token, "\r\n") {
return "", fmt.Errorf("invalid newline in nftables batch command")
}
}
script.WriteString(strings.Join(command, " "))
script.WriteByte('\n')
}
return script.String(), nil
}
func LoadInitStatus(tab string) (bool, bool, error) {
if tab != "base" {
return false, false, nil
+91
View File
@@ -3,6 +3,7 @@ package firewall
import (
"encoding/json"
"fmt"
"sort"
"strconv"
"strings"
@@ -193,3 +194,93 @@ func PortWhitelistKey(item PortWhitelist) string {
}
return key
}
type SystemPort struct {
Family string
Port string
Protocol string
}
func RuleForSystemPort(provider filter.Provider, port SystemPort) filter.FirewallRule {
scope := filter.Scope{Provider: provider, Direction: filter.DirectionInput}
family := filter.Family(strings.ToLower(strings.TrimSpace(port.Family)))
switch provider {
case filter.ProviderIptables, filter.ProviderNftables:
if family != filter.FamilyIPv6 {
family = filter.FamilyIPv4
}
scope.Family, scope.Table = family, "filter"
case filter.ProviderFirewalld:
if family != filter.FamilyIPv4 && family != filter.FamilyIPv6 {
family = filter.FamilyInet
}
scope.Family, scope.Zone = family, filter.FirewalldInputZone
case filter.ProviderUFW:
if family != filter.FamilyIPv6 {
family = filter.FamilyIPv4
}
scope.Family = family
}
return filter.FirewallRule{
Scope: scope, Protocol: port.Protocol, DestinationPort: port.Port,
Action: filter.ActionAccept, Description: "1Panel managed accepted port",
}
}
func NormalizeSystemPorts(ports []SystemPort) (map[string]SystemPort, error) {
result := make(map[string]SystemPort, len(ports))
for _, port := range ports {
normalized, err := filter.NormalizeRule(RuleForSystemPort(filter.ProviderIptables, port))
if err != nil {
return nil, err
}
family := strings.ToLower(strings.TrimSpace(port.Family))
if family != "" {
family = string(normalized.Scope.Family)
}
item := SystemPort{Family: family, Port: normalized.DestinationPort, Protocol: normalized.Protocol}
result[SystemPortKey(item)] = item
}
return result, nil
}
func SystemPortKey(port SystemPort) string {
key := LegacySystemPortKey(port)
if family := strings.ToLower(strings.TrimSpace(port.Family)); family != "" {
return family + "/" + key
}
return key
}
func LegacySystemPortKey(port SystemPort) string {
return strings.ToLower(strings.TrimSpace(port.Protocol)) + "/" + strings.TrimSpace(port.Port)
}
func SortedSystemPortKeys(ports map[string]SystemPort) []string {
keys := make([]string, 0, len(ports))
for key := range ports {
keys = append(keys, key)
}
sort.Strings(keys)
return keys
}
func ContainsPort(ports []PortWhitelist, target PortWhitelist) bool {
for _, port := range ports {
familyMatches := port.Family == "" || target.Family == "" || port.Family == target.Family
if familyMatches && port.Port == target.Port && port.Protocol == target.Protocol {
return true
}
}
return false
}
func ExcludePorts(ports, excluded []PortWhitelist) []PortWhitelist {
result := make([]PortWhitelist, 0, len(ports))
for _, port := range ports {
if !ContainsPort(excluded, port) {
result = append(result, port)
}
}
return result
}
@@ -1,72 +0,0 @@
package firewall
import "testing"
func TestParsePortWhitelistLegacyAndStructured(t *testing.T) {
legacy, err := ParsePortWhitelist("80/tcp,53/udp")
if err != nil {
t.Fatalf("parse legacy whitelist: %v", err)
}
if len(legacy) != 2 || legacy[0] != (PortWhitelist{Family: "ipv4", Port: "80", Protocol: "tcp"}) {
t.Fatalf("unexpected legacy whitelist: %#v", legacy)
}
structured, err := ParsePortWhitelist(`[
{"family":"ipv6","protocol":"TCP","port":"8080:8090"},
{"family":"ipv4","protocol":"udp","port":"53"}
]`)
if err != nil {
t.Fatalf("parse structured whitelist: %v", err)
}
want := []PortWhitelist{
{Family: "ipv6", Port: "8080-8090", Protocol: "tcp"},
{Family: "ipv4", Port: "53", Protocol: "udp"},
}
if len(structured) != len(want) {
t.Fatalf("unexpected structured whitelist: %#v", structured)
}
for index := range want {
if structured[index] != want[index] {
t.Fatalf("rule %d = %#v, want %#v", index, structured[index], want[index])
}
}
}
func TestParsePortWhitelistRejectsInvalidRule(t *testing.T) {
for _, value := range []string{
`[{"family":"inet","protocol":"tcp","port":"80"}]`,
`[{"family":"ipv4","protocol":"icmp","port":"80"}]`,
`[{"family":"ipv4","protocol":"tcp","port":"9000-8000"}]`,
`[{"family":"ipv6","protocol":"udp","port":"65536"}]`,
`[{"family":"ipv4","protocol":"tcp","port":"8000-8100"},{"family":"ipv4","protocol":"tcp","port":"8080"}]`,
} {
if _, err := ParsePortWhitelist(value); err == nil {
t.Fatalf("expected %q to fail", value)
}
}
}
func TestParsePortWhitelistKeepsFamiliesDistinct(t *testing.T) {
rules, err := ParsePortWhitelist(`[
{"family":"ipv4","protocol":"tcp","port":"443"},
{"family":"ipv6","protocol":"tcp","port":"443"},
{"family":"ipv4","protocol":"tcp","port":"443"}
]`)
if err != nil {
t.Fatalf("parse whitelist: %v", err)
}
if len(rules) != 2 {
t.Fatalf("expected family-specific deduplication, got %#v", rules)
}
}
func TestNormalizePortWhitelistPrefersFamilyNeutralRequiredRule(t *testing.T) {
rules := NormalizePortWhitelist([]PortWhitelist{
{Family: "ipv4", Port: "22", Protocol: "tcp"},
{Family: "ipv6", Port: "22", Protocol: "tcp"},
{Port: "22", Protocol: "tcp"},
})
if len(rules) != 1 || rules[0].Family != "" {
t.Fatalf("family-neutral rule should replace family-specific duplicates: %#v", rules)
}
}
+120
View File
@@ -0,0 +1,120 @@
package sync
type Status string
const (
StatusReady Status = "ready"
StatusExisting Status = "existing"
StatusRemove Status = "remove"
StatusBlocked Status = "blocked"
)
type Outcome string
const (
OutcomeApplied Outcome = "applied"
OutcomeSkipped Outcome = "skipped"
OutcomeRemoved Outcome = "removed"
OutcomeFailed Outcome = "failed"
)
type ReasonCode string
const (
ReasonInvalidPolicy ReasonCode = "invalid_policy"
ReasonAlreadyExists ReasonCode = "already_exists_in_target"
ReasonOnlyExistsInTarget ReasonCode = "only_exists_in_target"
ReasonManagedOnlyInTarget ReasonCode = "managed_only_exists_in_target"
ReasonUnsafeRemoval ReasonCode = "unsafe_managed_rule_removal"
)
func ReasonMessage(code ReasonCode) string {
switch code {
case ReasonAlreadyExists:
return "rule already exists in target backend"
case ReasonOnlyExistsInTarget:
return "rule exists only in target backend"
case ReasonManagedOnlyInTarget:
return "managed rule exists only in target backend"
case ReasonUnsafeRemoval:
return "managed runtime rule cannot be safely removed"
default:
return ""
}
}
type Desired[T any, P any] struct {
Value T
Payload P
Err error
}
type Item[P any] struct {
Payload P
Status Status
ReasonCode ReasonCode
Reason string
}
func Diff[T any, P any](desired []Desired[T, P], actual []T, key func(T) string, actualPayload func(T) P) []Item[P] {
items := make([]Item[P], 0, len(desired)+len(actual))
actualByKey := make(map[string][]int, len(actual))
for index, value := range actual {
actualByKey[key(value)] = append(actualByKey[key(value)], index)
}
matched := make([]bool, len(actual))
for _, candidate := range desired {
item := Item[P]{Payload: candidate.Payload}
switch {
case candidate.Err != nil:
item.Status, item.ReasonCode, item.Reason = StatusBlocked, ReasonInvalidPolicy, candidate.Err.Error()
default:
match := unmatchedIndex(actualByKey[key(candidate.Value)], matched)
if match >= 0 {
matched[match] = true
item.Status, item.ReasonCode = StatusExisting, ReasonAlreadyExists
item.Reason = ReasonMessage(item.ReasonCode)
} else {
item.Status = StatusReady
}
}
items = append(items, item)
}
for index, value := range actual {
if matched[index] {
continue
}
items = append(items, Item[P]{
Payload: actualPayload(value), Status: StatusRemove, ReasonCode: ReasonOnlyExistsInTarget,
Reason: ReasonMessage(ReasonOnlyExistsInTarget),
})
}
return items
}
func StatesEqual[T any](left, right []T, key func(T) string) bool {
if len(left) != len(right) {
return false
}
counts := make(map[string]int, len(left))
for _, value := range left {
counts[key(value)]++
}
for _, value := range right {
valueKey := key(value)
if counts[valueKey] == 0 {
return false
}
counts[valueKey]--
}
return true
}
func unmatchedIndex(indices []int, matched []bool) int {
for _, index := range indices {
if !matched[index] {
return index
}
}
return -1
}
+141
View File
@@ -0,0 +1,141 @@
package sync
import (
"slices"
"strings"
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
)
func SupportsManagedOrder(provider filter.Provider) bool {
return provider == filter.ProviderIptables || provider == filter.ProviderNftables || provider == filter.ProviderUFW
}
func ManagedOrderDrift(snapshot filter.Snapshot, desiredMarkers []string) (map[string]struct{}, bool) {
if !SupportsManagedOrder(snapshot.Scope.Provider) || len(desiredMarkers) < 2 {
return nil, true
}
expected := make(map[string]struct{}, len(desiredMarkers))
for _, marker := range desiredMarkers {
expected[marker] = struct{}{}
}
actual := make([]string, 0, len(desiredMarkers))
segments := make(map[string]int, len(desiredMarkers))
segment := 0
for _, observed := range snapshot.Rules {
_, wanted := expected[observed.Marker]
if wanted {
actual = append(actual, observed.Marker)
if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque {
segment++
segments[observed.Marker] = segment
segment++
} else {
segments[observed.Marker] = segment
}
continue
}
if strings.HasPrefix(observed.Marker, "1panel-rule:") &&
!observed.Protected && observed.ParseStatus != filter.ParseStatusOpaque {
continue
}
segment++
}
desired := make([]string, 0, len(actual))
for _, marker := range desiredMarkers {
if _, exists := segments[marker]; exists {
desired = append(desired, marker)
}
}
if slices.Equal(actual, desired) {
return nil, true
}
drifted := make(map[string]struct{}, len(desired))
for index := range desired {
if actual[index] != desired[index] {
drifted[actual[index]] = struct{}{}
drifted[desired[index]] = struct{}{}
}
}
feasible, previousSegment := true, -1
for _, marker := range desired {
if segments[marker] < previousSegment {
feasible = false
break
}
previousSegment = segments[marker]
}
return drifted, feasible
}
func InsertionPosition(snapshot filter.Snapshot, desiredMarkers []string, targetMarker string) (int64, bool) {
if !SupportsManagedOrder(snapshot.Scope.Provider) {
return 0, false
}
targetIndex := slices.Index(desiredMarkers, targetMarker)
if targetIndex < 0 {
return 0, false
}
for index := targetIndex - 1; index >= 0; index-- {
if _, position, exists := ObservedByMarker(snapshot, desiredMarkers[index]); exists {
return int64(position + 1), true
}
}
for index := targetIndex + 1; index < len(desiredMarkers); index++ {
if _, position, exists := ObservedByMarker(snapshot, desiredMarkers[index]); exists {
return int64(position), true
}
}
return 0, false
}
func NextManagedOrderChange(snapshot filter.Snapshot, desiredMarkers []string) (string, int, bool, error) {
expected := make(map[string]struct{}, len(desiredMarkers))
for _, marker := range desiredMarkers {
expected[marker] = struct{}{}
}
actual := make([]string, 0, len(desiredMarkers))
positions := make([]int, 0, len(desiredMarkers))
for index, observed := range snapshot.Rules {
if _, exists := expected[observed.Marker]; !exists {
continue
}
actual = append(actual, observed.Marker)
position := index + 1
if observed.Locator.Position != nil {
position = *observed.Locator.Position
}
positions = append(positions, position)
}
if len(actual) != len(desiredMarkers) {
return "", 0, false, filter.ErrRuleStale
}
for index := range desiredMarkers {
if actual[index] != desiredMarkers[index] {
return desiredMarkers[index], positions[index], false, nil
}
}
return "", 0, true, nil
}
func ObservedByMarker(snapshot filter.Snapshot, marker string) (filter.ObservedRule, int, bool) {
for index, observed := range snapshot.Rules {
if observed.Marker == marker {
position := index + 1
if observed.Locator.Position != nil {
position = *observed.Locator.Position
}
return observed, position, true
}
}
return filter.ObservedRule{}, 0, false
}
func ObservedRule(observed filter.ObservedRule) filter.FirewallRule {
rule := observed.Rule
if rule.UUID == "" && strings.HasPrefix(observed.Marker, "1panel-rule:") {
rule.UUID = strings.TrimSpace(strings.TrimPrefix(observed.Marker, "1panel-rule:"))
}
return rule
}
+1
View File
@@ -5,4 +5,5 @@ import "regexp"
var (
UFWNumberedRulePrefixRegex = regexp.MustCompile(`^\s*\[\s*([0-9]+)\]\s+(.+?)\s*$`)
UFWNumberedRuleRegex = regexp.MustCompile(`^\s*\[\s*([0-9]+)\]\s+(.+?)\s+(ALLOW|DENY|REJECT|LIMIT)(?:\s+(IN|OUT|FWD))?\s+(.+?)\s*$`)
ForwardInterfaceRegex = regexp.MustCompile(`^[A-Za-z0-9_.:@-]{1,15}$`)
)
-24
View File
@@ -1,24 +0,0 @@
package migration
import "testing"
func TestCoreMigrationsRegisterFirewallMenuUpgrade(t *testing.T) {
migrations := coreMigrations()
seen := make(map[string]int, len(migrations))
foundFirewallMenu := false
for index, migration := range migrations {
if migration == nil || migration.ID == "" {
t.Fatalf("invalid migration at index %d: %#v", index, migration)
}
if previous, exists := seen[migration.ID]; exists {
t.Fatalf("duplicate migration ID %q at indexes %d and %d", migration.ID, previous, index)
}
seen[migration.ID] = index
if migration.ID == "20260819-update-firewall-menu-path" {
foundFirewallMenu = true
}
}
if !foundFirewallMenu {
t.Fatal("firewall menu path migration is not registered")
}
}
@@ -1,114 +0,0 @@
package migrations
import (
"encoding/json"
"path/filepath"
"testing"
"github.com/1Panel-dev/1Panel/core/app/dto"
"github.com/1Panel-dev/1Panel/core/app/model"
"github.com/1Panel-dev/1Panel/core/init/migration/helper"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func TestUpdateFirewallMenuPathMigratesCustomizedMenu(t *testing.T) {
db := newCoreFirewallMigrationTestDB(t)
menus := []dto.ShowMenu{{
ID: "custom-parent", Label: "Custom", Children: []dto.ShowMenu{
{ID: "74", Label: "renamed-firewall", Title: "custom.title", Path: "/hosts/firewall/port", Sort: 987, IsShow: false},
{ID: "other", Label: "Other", Path: "/hosts/firewall/port", Sort: 123, IsShow: true},
},
}}
seedCoreHideMenu(t, db, menus)
if err := UpdateFirewallMenuPath.Migrate(db); err != nil {
t.Fatal(err)
}
after := loadCoreHideMenu(t, db)
firewall := after[0].Children[0]
if firewall.Path != "/hosts/firewall/rules" || firewall.Title != "custom.title" || firewall.Sort != 987 || firewall.IsShow {
t.Fatalf("migration did not preserve customized firewall menu: %#v", firewall)
}
if other := after[0].Children[1]; other.Path != "/hosts/firewall/port" {
t.Fatalf("unrelated menu was changed: %#v", other)
}
}
func TestUpdateFirewallMenuPathSupportsLegacyLabelAndIsIdempotent(t *testing.T) {
db := newCoreFirewallMigrationTestDB(t)
menus := []dto.ShowMenu{{Children: []dto.ShowMenu{
{ID: "legacy-id", Label: "FirewallPort", Path: "/hosts/firewall/port"},
{ID: "74", Label: "FirewallPort", Path: "/hosts/firewall/rules"},
}}}
seedCoreHideMenu(t, db, menus)
for i := 0; i < 2; i++ {
if err := UpdateFirewallMenuPath.Migrate(db); err != nil {
t.Fatalf("migrate firewall menu on pass %d: %v", i+1, err)
}
}
after := loadCoreHideMenu(t, db)
for _, menu := range after[0].Children {
if menu.Path != "/hosts/firewall/rules" {
t.Fatalf("legacy firewall menu path remained after migration: %#v", menu)
}
}
}
func TestDefaultMenuUsesFirewallV2Route(t *testing.T) {
var menus []dto.ShowMenu
if err := json.Unmarshal([]byte(helper.LoadMenus()), &menus); err != nil {
t.Fatal(err)
}
for _, parent := range menus {
for _, child := range parent.Children {
if child.ID == "74" {
if child.Path != "/hosts/firewall/rules" {
t.Fatalf("default firewall path = %q", child.Path)
}
return
}
}
}
t.Fatal("default firewall menu was not found")
}
func newCoreFirewallMigrationTestDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "migration.db")), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatal(err)
}
return db
}
func seedCoreHideMenu(t *testing.T, db *gorm.DB, menus []dto.ShowMenu) {
t.Helper()
value, err := json.Marshal(menus)
if err != nil {
t.Fatal(err)
}
if err := db.Create(&model.Setting{Key: "HideMenu", Value: string(value)}).Error; err != nil {
t.Fatal(err)
}
}
func loadCoreHideMenu(t *testing.T, db *gorm.DB) []dto.ShowMenu {
t.Helper()
var setting model.Setting
if err := db.Where("key = ?", "HideMenu").First(&setting).Error; err != nil {
t.Fatal(err)
}
var menus []dto.ShowMenu
if err := json.Unmarshal([]byte(setting.Value), &menus); err != nil {
t.Fatal(err)
}
return menus
}
+1
View File
@@ -277,6 +277,7 @@ export namespace Firewall {
forwardRule?: RuleForward;
dockerRule?: DockerGuardEndpoint;
status: RuleSyncStatus;
reasonCode?: string;
reason?: string;
}
@@ -147,7 +147,7 @@
</template>
<el-table-column :label="$t('firewall.ruleSyncReason')" min-width="220" show-overflow-tooltip>
<template #default="{ row }">
{{ reasonText(row.reason) }}
{{ reasonText(row.reasonCode, row.reason) }}
</template>
</el-table-column>
</ComplexTable>
@@ -296,7 +296,16 @@ const syncReasonKeys: Record<string, string> = {
'protected firewall rule cannot be modified': 'protectedRule',
};
const reasonText = (reason?: string) => {
const syncReasonCodeKeys: Record<string, string> = {
already_exists_in_target: 'alreadyExistsInTarget',
only_exists_in_target: 'onlyInTarget',
managed_only_exists_in_target: 'managedOnlyInTarget',
unsafe_managed_rule_removal: 'managedRuntimeCannotRemove',
};
const reasonText = (reasonCode?: string, reason?: string) => {
const codedReasonKey = reasonCode ? syncReasonCodeKeys[reasonCode] : undefined;
if (codedReasonKey) return i18n.global.t(`firewall.ruleSyncReasonDetail.${codedReasonKey}`);
if (!reason) return '-';
const reasonKey = syncReasonKeys[reason];
if (reasonKey) return i18n.global.t(`firewall.ruleSyncReasonDetail.${reasonKey}`);