mirror of
https://github.com/1Panel-dev/1Panel.git
synced 2026-09-22 08:00:53 +00:00
refactor: simplify firewall service structure (#13646)
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
}
|
||||
@@ -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 = ¤tPosition
|
||||
} else if *after.OrderIndex != currentPosition {
|
||||
if err := validatePositionTarget(ctx, runtime, snapshot, before.Rule, *after.OrderIndex); err != nil {
|
||||
if err := runtime.ValidatePosition(ctx, snapshot, before.Rule, *after.OrderIndex); err != nil {
|
||||
return preparedManagedUpdate{}, err
|
||||
}
|
||||
}
|
||||
@@ -1630,9 +1528,10 @@ func (s *FirewallService) selectedProviderForStoredRule(
|
||||
if s.selectedProvider != nil {
|
||||
return s.selectedProvider(ctx)
|
||||
}
|
||||
if len(s.adapters) == 1 {
|
||||
for provider := range s.adapters {
|
||||
return provider, nil
|
||||
if s.adapters != nil {
|
||||
providers := s.adapters.Providers()
|
||||
if len(providers) == 1 {
|
||||
return providers[0], nil
|
||||
}
|
||||
}
|
||||
return "", fmt.Errorf("%w: selected provider is unavailable", filter.ErrProviderUnavailable)
|
||||
@@ -1649,7 +1548,7 @@ func (s *FirewallService) executeManagedMutation(ctx context.Context, request ma
|
||||
}
|
||||
appendRule, restoreAtEnd := false, false
|
||||
if after.Scope.Provider == filter.ProviderUFW && request.AdapterOperation == filter.ChangeUpdate {
|
||||
maxPosition, maxErr := maxPositionForRule(ctx, request.Runtime, request.Snapshot, after)
|
||||
maxPosition, maxErr := request.Runtime.MaxPosition(ctx, request.Snapshot, after)
|
||||
if maxErr != nil {
|
||||
return maxErr
|
||||
}
|
||||
@@ -1910,71 +1809,39 @@ func (s *FirewallService) adoptExternalSystemPort(ctx context.Context, port dto.
|
||||
}
|
||||
|
||||
func systemPortRule(provider filter.Provider, port dto.FirewallSystemPort) filter.FirewallRule {
|
||||
scope := filter.Scope{Provider: provider, Direction: filter.DirectionInput}
|
||||
family := filter.Family(strings.ToLower(strings.TrimSpace(port.Family)))
|
||||
switch provider {
|
||||
case filter.ProviderIptables, filter.ProviderNftables:
|
||||
if family != filter.FamilyIPv6 {
|
||||
family = filter.FamilyIPv4
|
||||
}
|
||||
scope.Family = family
|
||||
scope.Table = "filter"
|
||||
case filter.ProviderFirewalld:
|
||||
if family != filter.FamilyIPv4 && family != filter.FamilyIPv6 {
|
||||
family = filter.FamilyInet
|
||||
}
|
||||
scope.Family = family
|
||||
scope.Zone = filter.FirewalldInputZone
|
||||
case filter.ProviderUFW:
|
||||
if family != filter.FamilyIPv6 {
|
||||
family = filter.FamilyIPv4
|
||||
}
|
||||
scope.Family = family
|
||||
}
|
||||
return filter.FirewallRule{
|
||||
Scope: scope, Protocol: port.Protocol, DestinationPort: port.Port,
|
||||
Action: filter.ActionAccept, Description: "1Panel managed accepted port",
|
||||
}
|
||||
return firewall.RuleForSystemPort(provider, firewall.SystemPort(port))
|
||||
}
|
||||
|
||||
func normalizeSystemPorts(ports []dto.FirewallSystemPort) (map[string]dto.FirewallSystemPort, error) {
|
||||
result := make(map[string]dto.FirewallSystemPort, len(ports))
|
||||
domainPorts := make([]firewall.SystemPort, 0, len(ports))
|
||||
for _, port := range ports {
|
||||
normalized, err := filter.NormalizeRule(systemPortRule(filter.ProviderIptables, port))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
family := strings.ToLower(strings.TrimSpace(port.Family))
|
||||
if family != "" {
|
||||
family = string(normalized.Scope.Family)
|
||||
}
|
||||
item := dto.FirewallSystemPort{
|
||||
Family: family, Port: normalized.DestinationPort, Protocol: normalized.Protocol,
|
||||
}
|
||||
result[systemPortKey(item)] = item
|
||||
domainPorts = append(domainPorts, firewall.SystemPort(port))
|
||||
}
|
||||
normalized, err := firewall.NormalizeSystemPorts(domainPorts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make(map[string]dto.FirewallSystemPort, len(normalized))
|
||||
for key, port := range normalized {
|
||||
result[key] = dto.FirewallSystemPort(port)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func systemPortKey(port dto.FirewallSystemPort) string {
|
||||
key := legacySystemPortKey(port)
|
||||
if family := strings.ToLower(strings.TrimSpace(port.Family)); family != "" {
|
||||
return family + "/" + key
|
||||
}
|
||||
return key
|
||||
return firewall.SystemPortKey(firewall.SystemPort(port))
|
||||
}
|
||||
|
||||
func legacySystemPortKey(port dto.FirewallSystemPort) string {
|
||||
return strings.ToLower(strings.TrimSpace(port.Protocol)) + "/" + strings.TrimSpace(port.Port)
|
||||
return firewall.LegacySystemPortKey(firewall.SystemPort(port))
|
||||
}
|
||||
|
||||
func sortedSystemPortKeys(ports map[string]dto.FirewallSystemPort) []string {
|
||||
keys := make([]string, 0, len(ports))
|
||||
for key := range ports {
|
||||
keys = append(keys, key)
|
||||
domainPorts := make(map[string]firewall.SystemPort, len(ports))
|
||||
for key, port := range ports {
|
||||
domainPorts[key] = firewall.SystemPort(port)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
return keys
|
||||
return firewall.SortedSystemPortKeys(domainPorts)
|
||||
}
|
||||
|
||||
func firewallRuleModelForCreate(rule filter.FirewallRule, request dto.FirewallRuleCreateItem, origin string) (model.FirewallRule, error) {
|
||||
@@ -2162,7 +2029,7 @@ func mergeFirewallInventory(
|
||||
}
|
||||
|
||||
func desiredFirewallRuleFromModel(stored model.FirewallRule) (filter.DesiredRule, error) {
|
||||
rules, err := firewallPolicyRulesForProvider(stored, filter.ProviderIptables)
|
||||
rules, err := stored.RulesForProvider(filter.ProviderIptables)
|
||||
if err != nil {
|
||||
return filter.DesiredRule{}, err
|
||||
}
|
||||
@@ -2188,7 +2055,7 @@ func (s *FirewallService) compileStoredFirewallRules(
|
||||
stored model.FirewallRule,
|
||||
target filter.Provider,
|
||||
) ([]filter.DesiredRule, error) {
|
||||
rules, err := firewallPolicyRulesForProvider(stored, target)
|
||||
rules, err := stored.RulesForProvider(target)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -2196,41 +2063,7 @@ func (s *FirewallService) compileStoredFirewallRules(
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make([]filter.DesiredRule, 0, len(rules))
|
||||
scopeOrdinals := make(map[string]int)
|
||||
for _, rule := range rules {
|
||||
rule, err = runtime.Prepare(rule)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = runtime.CheckRule(ctx, rule); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ruleKey, keyErr := filter.RuleKey(rule)
|
||||
if keyErr != nil {
|
||||
return nil, keyErr
|
||||
}
|
||||
scopeKey := rule.Scope.Key()
|
||||
ordinal := scopeOrdinals[scopeKey]
|
||||
scopeOrdinals[scopeKey] = ordinal + 1
|
||||
rule.UUID = compiledFirewallRuleUUID(stored.UUID, ruleKey, ordinal)
|
||||
result = append(result, filter.DesiredRule{
|
||||
UUID: stored.UUID, Rule: rule, RuleKey: ruleKey, Origin: filter.RuleOrigin(stored.Origin),
|
||||
Marker: "1panel-rule:" + rule.UUID,
|
||||
})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func compiledFirewallRuleUUID(policyUUID, ruleKey string, scopeOrdinal int) string {
|
||||
if scopeOrdinal == 0 {
|
||||
return policyUUID
|
||||
}
|
||||
const suffixLength = 12
|
||||
if len(ruleKey) > suffixLength {
|
||||
ruleKey = ruleKey[:suffixLength]
|
||||
}
|
||||
return fmt.Sprintf("%s-%d-%s", policyUUID, scopeOrdinal+1, ruleKey)
|
||||
return runtime.CompileDesired(ctx, stored.UUID, filter.RuleOrigin(stored.Origin), rules)
|
||||
}
|
||||
|
||||
func (s *FirewallService) desiredFirewallRulesForScope(
|
||||
@@ -2238,7 +2071,7 @@ func (s *FirewallService) desiredFirewallRulesForScope(
|
||||
stored []model.FirewallRule,
|
||||
scope filter.Scope,
|
||||
) ([]filter.DesiredRule, error) {
|
||||
sortFirewallPolicies(stored, scope.Provider)
|
||||
model.SortFirewallRules(stored, scope.Provider)
|
||||
desired := make([]filter.DesiredRule, 0, len(stored))
|
||||
for _, record := range stored {
|
||||
compiled, err := s.compileStoredFirewallRules(ctx, record, scope.Provider)
|
||||
@@ -2266,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)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -20,6 +19,7 @@ import (
|
||||
"github.com/1Panel-dev/1Panel/agent/utils/firewall/docker_guard"
|
||||
"github.com/1Panel-dev/1Panel/agent/utils/firewall/filter"
|
||||
"github.com/1Panel-dev/1Panel/agent/utils/firewall/forwarding"
|
||||
firewallsync "github.com/1Panel-dev/1Panel/agent/utils/firewall/sync"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
@@ -29,9 +29,9 @@ var (
|
||||
)
|
||||
|
||||
const (
|
||||
firewallRuleSyncReady = "ready"
|
||||
firewallRuleSyncExisting = "existing"
|
||||
firewallRuleSyncBlocked = "blocked"
|
||||
firewallRuleSyncReady = firewallsync.StatusReady
|
||||
firewallRuleSyncExisting = firewallsync.StatusExisting
|
||||
firewallRuleSyncBlocked = firewallsync.StatusBlocked
|
||||
)
|
||||
|
||||
func firewallSyncSubsystem(value string) string {
|
||||
@@ -42,13 +42,13 @@ func firewallSyncSubsystem(value string) string {
|
||||
return value
|
||||
}
|
||||
|
||||
type firewallRuleSyncOutcome string
|
||||
type firewallRuleSyncOutcome = firewallsync.Outcome
|
||||
|
||||
const (
|
||||
firewallRuleSyncApplied firewallRuleSyncOutcome = "applied"
|
||||
firewallRuleSyncSkipped firewallRuleSyncOutcome = "skipped"
|
||||
firewallRuleSyncRemoved firewallRuleSyncOutcome = "removed"
|
||||
firewallRuleSyncFailed firewallRuleSyncOutcome = "failed"
|
||||
firewallRuleSyncApplied = firewallsync.OutcomeApplied
|
||||
firewallRuleSyncSkipped = firewallsync.OutcomeSkipped
|
||||
firewallRuleSyncRemoved = firewallsync.OutcomeRemoved
|
||||
firewallRuleSyncFailed = firewallsync.OutcomeFailed
|
||||
)
|
||||
|
||||
type firewallRuleSyncEntry struct {
|
||||
@@ -146,7 +146,7 @@ func (s *FirewallService) loadFirewallRuleSyncPlan(
|
||||
expectedMarkers[scopeKey+"\x00"+entry.desired.Marker] = struct{}{}
|
||||
}
|
||||
|
||||
scopes := firewallRuleSyncScopes(target)
|
||||
scopes := filter.ManagedInputScopes(target)
|
||||
if hasCompileErrors {
|
||||
scopes = scopesWithFirewallSyncCandidates(scopes, entriesByScope)
|
||||
}
|
||||
@@ -203,14 +203,18 @@ func (s *FirewallService) loadFirewallRuleSyncPlan(
|
||||
source: model.FirewallRule{UUID: strings.TrimPrefix(observed.Marker, "1panel-rule:")},
|
||||
rule: observed.Rule, remove: ©,
|
||||
}
|
||||
status, reason := "remove", "managed rule exists only in target backend"
|
||||
status := firewallsync.StatusRemove
|
||||
reasonCode := firewallsync.ReasonManagedOnlyInTarget
|
||||
reason := firewallsync.ReasonMessage(reasonCode)
|
||||
if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque {
|
||||
status, reason = firewallRuleSyncBlocked, "managed runtime rule cannot be safely removed"
|
||||
status = firewallRuleSyncBlocked
|
||||
reasonCode = firewallsync.ReasonUnsafeRemoval
|
||||
reason = firewallsync.ReasonMessage(reasonCode)
|
||||
entry.err = filter.ErrProtectedRule
|
||||
}
|
||||
entry.item = dto.FirewallRuleSyncItem{
|
||||
SourceUUID: entry.source.UUID, Rule: &entry.rule,
|
||||
Status: status, Reason: reason,
|
||||
Status: status, ReasonCode: reasonCode, Reason: reason,
|
||||
}
|
||||
entries = append(entries, entry)
|
||||
scopeEntries = append(scopeEntries, entry)
|
||||
@@ -232,76 +236,20 @@ func (s *FirewallService) loadFirewallRuleSyncPlan(
|
||||
}, nil
|
||||
}
|
||||
|
||||
func firewallProviderHasOrderedManagedRules(provider filter.Provider) bool {
|
||||
return provider == filter.ProviderIptables || provider == filter.ProviderNftables || provider == filter.ProviderUFW
|
||||
}
|
||||
|
||||
func planFirewallManagedOrder(
|
||||
snapshot filter.Snapshot,
|
||||
entries []*firewallRuleSyncEntry,
|
||||
) {
|
||||
if !firewallProviderHasOrderedManagedRules(snapshot.Scope.Provider) {
|
||||
return
|
||||
}
|
||||
byMarker := make(map[string]*firewallRuleSyncEntry, len(entries))
|
||||
desiredMarkers := make([]string, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
if entry.remove == nil && entry.err == nil && entry.desired.Marker != "" {
|
||||
byMarker[entry.desired.Marker] = entry
|
||||
desiredMarkers = append(desiredMarkers, entry.desired.Marker)
|
||||
}
|
||||
}
|
||||
if len(byMarker) < 2 {
|
||||
return
|
||||
}
|
||||
|
||||
actual := make([]string, 0, len(byMarker))
|
||||
segments := make(map[string]int, len(byMarker))
|
||||
segment := 0
|
||||
for _, observed := range snapshot.Rules {
|
||||
_, expected := byMarker[observed.Marker]
|
||||
if expected {
|
||||
actual = append(actual, observed.Marker)
|
||||
if observed.Protected || observed.ParseStatus == filter.ParseStatusOpaque {
|
||||
segment++
|
||||
segments[observed.Marker] = segment
|
||||
segment++
|
||||
} else {
|
||||
segments[observed.Marker] = segment
|
||||
}
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(observed.Marker, "1panel-rule:") &&
|
||||
!observed.Protected && observed.ParseStatus != filter.ParseStatusOpaque {
|
||||
continue
|
||||
}
|
||||
segment++
|
||||
}
|
||||
|
||||
desired := make([]string, 0, len(actual))
|
||||
for _, entry := range entries {
|
||||
marker := entry.desired.Marker
|
||||
if _, exists := segments[marker]; exists {
|
||||
desired = append(desired, marker)
|
||||
}
|
||||
}
|
||||
if slices.Equal(actual, desired) {
|
||||
return
|
||||
}
|
||||
drifted := make(map[string]struct{}, len(desired))
|
||||
for index := range desired {
|
||||
if actual[index] != desired[index] {
|
||||
drifted[actual[index]] = struct{}{}
|
||||
drifted[desired[index]] = struct{}{}
|
||||
}
|
||||
}
|
||||
feasible, previousSegment := true, -1
|
||||
for _, marker := range desired {
|
||||
if segments[marker] < previousSegment {
|
||||
feasible = false
|
||||
break
|
||||
}
|
||||
previousSegment = segments[marker]
|
||||
}
|
||||
for _, marker := range desired {
|
||||
drifted, feasible := firewallsync.ManagedOrderDrift(snapshot, desiredMarkers)
|
||||
for _, marker := range desiredMarkers {
|
||||
if _, exists := drifted[marker]; !exists {
|
||||
continue
|
||||
}
|
||||
@@ -529,7 +477,7 @@ func (s *FirewallService) loadStoredFirewallRuleSyncCandidates(
|
||||
if err != nil {
|
||||
return "", nil, false, err
|
||||
}
|
||||
sortFirewallPolicies(stored, selected)
|
||||
model.SortFirewallRules(stored, selected)
|
||||
entries := make([]*firewallRuleSyncEntry, 0, len(stored))
|
||||
hasCompileErrors := false
|
||||
for _, record := range stored {
|
||||
@@ -558,137 +506,12 @@ func scopesWithFirewallSyncCandidates(scopes []filter.Scope, candidates map[stri
|
||||
return result
|
||||
}
|
||||
|
||||
func firewallRuleSyncScopes(provider filter.Provider) []filter.Scope {
|
||||
base := filter.Scope{Provider: provider, Direction: filter.DirectionInput}
|
||||
switch provider {
|
||||
case filter.ProviderIptables, filter.ProviderNftables:
|
||||
result := make([]filter.Scope, 0, 6)
|
||||
for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} {
|
||||
for _, chain := range []string{filter.BasicBeforeChain, filter.IptablesInputChain, filter.BasicAfterChain} {
|
||||
scope := base
|
||||
scope.Family, scope.Table, scope.Chain = family, "filter", chain
|
||||
result = append(result, scope)
|
||||
}
|
||||
}
|
||||
return result
|
||||
case filter.ProviderFirewalld:
|
||||
base.Family, base.Zone = filter.FamilyInet, filter.FirewalldInputZone
|
||||
return []filter.Scope{base}
|
||||
case filter.ProviderUFW:
|
||||
result := make([]filter.Scope, 0, 2)
|
||||
for _, family := range []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6} {
|
||||
scope := base
|
||||
scope.Family, scope.Chain = family, filter.UFWInputChain
|
||||
result = append(result, scope)
|
||||
}
|
||||
return result
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func firewallPolicyRulesForProvider(stored model.FirewallRule, provider filter.Provider) ([]filter.FirewallRule, error) {
|
||||
if stored.CompatibilityError != "" {
|
||||
return nil, fmt.Errorf("%w: %s", filter.ErrUnsupportedScope, stored.CompatibilityError)
|
||||
}
|
||||
connectionStates := make([]string, 0)
|
||||
if stored.ConnectionStates != "" {
|
||||
connectionStates = strings.Split(stored.ConnectionStates, ",")
|
||||
}
|
||||
base := filter.FirewallRule{
|
||||
Protocol: stored.Protocol, SourceAddress: stored.SourceAddress, SourcePort: stored.SourcePort,
|
||||
DestinationAddress: stored.DestinationAddress, DestinationPort: stored.DestinationPort,
|
||||
Interface: stored.Interface, ConnectionStates: connectionStates,
|
||||
Action: filter.Action(stored.Action), Description: stored.Description,
|
||||
}
|
||||
if provider == filter.ProviderFirewalld {
|
||||
base.Priority = stored.Priority
|
||||
}
|
||||
families := []filter.Family{filter.Family(stored.Family)}
|
||||
if provider != filter.ProviderFirewalld && len(families) == 1 && families[0] == filter.FamilyInet {
|
||||
hasIPv4, hasIPv6 := firewallRuleAddressFamilies(base)
|
||||
switch {
|
||||
case hasIPv4 && hasIPv6:
|
||||
return nil, fmt.Errorf("%w: inet policy contains both IPv4 and IPv6 addresses", filter.ErrUnsupportedScope)
|
||||
case hasIPv6 || strings.EqualFold(base.Protocol, "icmpv6"):
|
||||
families = []filter.Family{filter.FamilyIPv6}
|
||||
case hasIPv4:
|
||||
families = []filter.Family{filter.FamilyIPv4}
|
||||
default:
|
||||
families = []filter.Family{filter.FamilyIPv4, filter.FamilyIPv6}
|
||||
}
|
||||
}
|
||||
result := make([]filter.FirewallRule, 0, len(families))
|
||||
for _, family := range families {
|
||||
rule := base
|
||||
rule.Scope = filter.Scope{Provider: provider, Family: family, Direction: filter.DirectionInput}
|
||||
switch provider {
|
||||
case filter.ProviderIptables, filter.ProviderNftables:
|
||||
rule.Scope.Table, rule.Scope.Chain = "filter", filter.IptablesInputChain
|
||||
case filter.ProviderFirewalld:
|
||||
rule.Scope.Zone = filter.FirewalldInputZone
|
||||
case filter.ProviderUFW:
|
||||
rule.Scope.Chain = filter.UFWInputChain
|
||||
default:
|
||||
return nil, fmt.Errorf("%w: unsupported firewall provider %q", filter.ErrProviderUnavailable, provider)
|
||||
}
|
||||
expanded, err := filter.ExpandAtomicRules(rule)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, expanded...)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func sortFirewallPolicies(stored []model.FirewallRule, provider filter.Provider) {
|
||||
sort.SliceStable(stored, func(i, j int) bool {
|
||||
left, right := stored[i], stored[j]
|
||||
if provider == filter.ProviderFirewalld {
|
||||
switch {
|
||||
case left.Priority == nil && right.Priority != nil:
|
||||
return false
|
||||
case left.Priority != nil && right.Priority == nil:
|
||||
return true
|
||||
case left.Priority != nil && right.Priority != nil && *left.Priority != *right.Priority:
|
||||
return *left.Priority < *right.Priority
|
||||
}
|
||||
} else {
|
||||
switch {
|
||||
case left.Sequence == nil && right.Sequence != nil:
|
||||
return false
|
||||
case left.Sequence != nil && right.Sequence == nil:
|
||||
return true
|
||||
case left.Sequence != nil && right.Sequence != nil && *left.Sequence != *right.Sequence:
|
||||
return *left.Sequence < *right.Sequence
|
||||
}
|
||||
}
|
||||
return left.UUID < right.UUID
|
||||
})
|
||||
}
|
||||
|
||||
func firewallRuleAddressFamilies(rule filter.FirewallRule) (bool, bool) {
|
||||
hasIPv4, hasIPv6 := false, false
|
||||
for _, address := range []string{rule.SourceAddress, rule.DestinationAddress} {
|
||||
address = strings.TrimSpace(address)
|
||||
if address == "" {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(address, ":") {
|
||||
hasIPv6 = true
|
||||
} else {
|
||||
hasIPv4 = true
|
||||
}
|
||||
}
|
||||
return hasIPv4, hasIPv6
|
||||
}
|
||||
|
||||
func (s *FirewallService) classifyFirewallRuleSyncCandidate(
|
||||
clientIP string,
|
||||
snapshot filter.Snapshot,
|
||||
entry *firewallRuleSyncEntry,
|
||||
item filter.InventoryItem,
|
||||
) (string, string) {
|
||||
) (firewallsync.Status, string) {
|
||||
switch item.Match {
|
||||
case filter.InventoryMatchExact:
|
||||
return firewallRuleSyncExisting, "rule already matches database policy"
|
||||
@@ -818,64 +641,24 @@ func buildDatabaseSyncPlan[T any](
|
||||
key func(T) string,
|
||||
actualItem func(T) dto.FirewallRuleSyncItem,
|
||||
) databaseSyncPlan {
|
||||
items := make([]dto.FirewallRuleSyncItem, 0, len(desired)+len(actual))
|
||||
actualByKey := make(map[string][]int, len(actual))
|
||||
for index, value := range actual {
|
||||
actualByKey[key(value)] = append(actualByKey[key(value)], index)
|
||||
}
|
||||
matched := make([]bool, len(actual))
|
||||
candidates := make([]firewallsync.Desired[T, dto.FirewallRuleSyncItem], 0, len(desired))
|
||||
for _, candidate := range desired {
|
||||
item := candidate.item
|
||||
switch {
|
||||
case candidate.err != nil:
|
||||
item.Status, item.Reason = firewallRuleSyncBlocked, candidate.err.Error()
|
||||
default:
|
||||
match := unmatchedDatabaseSyncIndex(actualByKey[key(candidate.value)], matched)
|
||||
if match >= 0 {
|
||||
matched[match] = true
|
||||
item.Status, item.Reason = firewallRuleSyncExisting, "rule already exists in target backend"
|
||||
} else {
|
||||
item.Status = firewallRuleSyncReady
|
||||
}
|
||||
}
|
||||
items = append(items, item)
|
||||
candidates = append(candidates, firewallsync.Desired[T, dto.FirewallRuleSyncItem]{
|
||||
Value: candidate.value, Payload: candidate.item, Err: candidate.err,
|
||||
})
|
||||
}
|
||||
for index, value := range actual {
|
||||
if matched[index] {
|
||||
continue
|
||||
}
|
||||
item := actualItem(value)
|
||||
item.Status, item.Reason = "remove", "rule exists only in target backend"
|
||||
diff := firewallsync.Diff(candidates, actual, key, actualItem)
|
||||
items := make([]dto.FirewallRuleSyncItem, 0, len(diff))
|
||||
for _, diffItem := range diff {
|
||||
item := diffItem.Payload
|
||||
item.Status, item.ReasonCode, item.Reason = diffItem.Status, diffItem.ReasonCode, diffItem.Reason
|
||||
items = append(items, item)
|
||||
}
|
||||
return databaseSyncPlan{subsystem: subsystem, target: target, items: items}
|
||||
}
|
||||
|
||||
func unmatchedDatabaseSyncIndex(indices []int, matched []bool) int {
|
||||
for _, index := range indices {
|
||||
if !matched[index] {
|
||||
return index
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
func databaseSyncStatesEqual[T any](left, right []T, key func(T) string) bool {
|
||||
if len(left) != len(right) {
|
||||
return false
|
||||
}
|
||||
counts := make(map[string]int, len(left))
|
||||
for _, value := range left {
|
||||
counts[key(value)]++
|
||||
}
|
||||
for _, value := range right {
|
||||
valueKey := key(value)
|
||||
if counts[valueKey] == 0 {
|
||||
return false
|
||||
}
|
||||
counts[valueKey]--
|
||||
}
|
||||
return true
|
||||
return firewallsync.StatesEqual(left, right, key)
|
||||
}
|
||||
|
||||
func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview {
|
||||
@@ -894,7 +677,7 @@ func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview {
|
||||
case firewallRuleSyncBlocked:
|
||||
result.Blocked++
|
||||
result.Total++
|
||||
case "remove":
|
||||
case firewallsync.StatusRemove:
|
||||
result.Removed++
|
||||
}
|
||||
}
|
||||
@@ -904,7 +687,7 @@ func (p databaseSyncPlan) preview() dto.FirewallRuleSyncPreview {
|
||||
func (p databaseSyncPlan) baseResult() dto.FirewallRuleSyncResult {
|
||||
result := dto.FirewallRuleSyncResult{Subsystem: p.subsystem, TargetProvider: p.target}
|
||||
for _, item := range p.items {
|
||||
if item.Status != "remove" {
|
||||
if item.Status != firewallsync.StatusRemove {
|
||||
result.Total++
|
||||
}
|
||||
}
|
||||
@@ -921,7 +704,7 @@ func (p databaseSyncPlan) completedResult() dto.FirewallRuleSyncResult {
|
||||
result.Skipped++
|
||||
case firewallRuleSyncBlocked:
|
||||
appendDatabaseSyncFailure(&result, item, errors.New(item.Reason))
|
||||
case "remove":
|
||||
case firewallsync.StatusRemove:
|
||||
result.Removed++
|
||||
}
|
||||
}
|
||||
@@ -1116,10 +899,10 @@ func (s *DockerPortGuardService) syncRules(
|
||||
}
|
||||
plan := buildDockerDatabaseSyncPlan(filter.Provider(target), policies, targetPolicies)
|
||||
result, reconcileErr := plan.reconcile(func() error {
|
||||
if err := reconcileDockerGuardSyncTarget(target, runtimePolicies, targetRuntime); err != nil {
|
||||
if err := docker_guard.ReconcileTarget(target, runtimePolicies, targetRuntime); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := verifyDockerGuardRuleSync(targetRuntime, runtimePolicies); err != nil {
|
||||
if err := docker_guard.Verify(targetRuntime, runtimePolicies); err != nil {
|
||||
return err
|
||||
}
|
||||
if len(policies) == 0 {
|
||||
@@ -1153,7 +936,7 @@ func buildDockerDatabaseSyncPlan(
|
||||
})
|
||||
}
|
||||
return buildDatabaseSyncPlan(
|
||||
"docker", target, desired, actual, dockerGuardPolicySyncKey,
|
||||
"docker", target, desired, actual, docker_guard.PolicySyncKey,
|
||||
func(policy docker_guard.Policy) dto.FirewallRuleSyncItem {
|
||||
return dto.FirewallRuleSyncItem{SourceUUID: policy.UUID, DockerRule: dockerGuardRuntimeRuleSyncDTO(policy)}
|
||||
},
|
||||
@@ -1178,7 +961,7 @@ func (r *firewallScopeReconciler) reconcile() {
|
||||
continue
|
||||
}
|
||||
switch entry.item.Status {
|
||||
case "remove":
|
||||
case firewallsync.StatusRemove:
|
||||
removes = append(removes, entry)
|
||||
case firewallRuleSyncReady:
|
||||
switch entry.match {
|
||||
@@ -1217,7 +1000,7 @@ func (r *firewallScopeReconciler) applyGroups(entries []*firewallRuleSyncEntry)
|
||||
if len(entries) == 0 {
|
||||
return
|
||||
}
|
||||
batch := r.runtime.adapter.Provider() == filter.ProviderIptables || r.runtime.adapter.Provider() == filter.ProviderNftables
|
||||
batch := r.runtime.Provider() == filter.ProviderIptables || r.runtime.Provider() == filter.ProviderNftables
|
||||
if batch {
|
||||
r.apply(entries)
|
||||
return
|
||||
@@ -1237,7 +1020,7 @@ func (r *firewallScopeReconciler) apply(entries []*firewallRuleSyncEntry) {
|
||||
continue
|
||||
}
|
||||
if !changed {
|
||||
if entry.item.Status == "remove" {
|
||||
if entry.item.Status == firewallsync.StatusRemove {
|
||||
entry.outcome = firewallRuleSyncRemoved
|
||||
} else {
|
||||
entry.outcome = firewallRuleSyncSkipped
|
||||
@@ -1261,7 +1044,7 @@ func (r *firewallScopeReconciler) apply(entries []*firewallRuleSyncEntry) {
|
||||
return
|
||||
}
|
||||
for _, entry := range active {
|
||||
if entry.item.Status == "remove" {
|
||||
if entry.item.Status == firewallsync.StatusRemove {
|
||||
entry.outcome = firewallRuleSyncRemoved
|
||||
} else {
|
||||
entry.outcome = firewallRuleSyncApplied
|
||||
@@ -1294,7 +1077,7 @@ func (r *firewallScopeReconciler) restoreOrder() {
|
||||
|
||||
changed := false
|
||||
for step := 0; step < len(desiredMarkers); step++ {
|
||||
marker, position, converged, err := nextFirewallManagedOrderChange(r.snapshot, desiredMarkers)
|
||||
marker, position, converged, err := firewallsync.NextManagedOrderChange(r.snapshot, desiredMarkers)
|
||||
if err != nil {
|
||||
r.failOrder(reorderEntries, err)
|
||||
return
|
||||
@@ -1309,18 +1092,18 @@ func (r *firewallScopeReconciler) restoreOrder() {
|
||||
}
|
||||
return
|
||||
}
|
||||
observed, _, exists := firewallRuleSyncObservedByMarker(r.snapshot, marker)
|
||||
observed, _, exists := firewallsync.ObservedByMarker(r.snapshot, marker)
|
||||
if !exists {
|
||||
r.failOrder(reorderEntries, filter.ErrRuleStale)
|
||||
return
|
||||
}
|
||||
after := firewallRuleSyncObservedRule(observed)
|
||||
after := firewallsync.ObservedRule(observed)
|
||||
target := int64(position)
|
||||
after.OrderIndex = &target
|
||||
before := firewallRuleSyncObservedRule(observed)
|
||||
before := firewallsync.ObservedRule(observed)
|
||||
locator := observed.Locator
|
||||
operation := filter.ChangeReorder
|
||||
if r.runtime.adapter.Provider() == filter.ProviderUFW {
|
||||
if r.runtime.Provider() == filter.ProviderUFW {
|
||||
operation = filter.ChangeUpdate
|
||||
}
|
||||
_, verification, executeErr := r.runtime.Execute(r.ctx, r.snapshot, []filter.DesiredChange{{
|
||||
@@ -1339,59 +1122,6 @@ func (r *firewallScopeReconciler) restoreOrder() {
|
||||
r.failOrder(reorderEntries, fmt.Errorf("%w: managed rule order did not converge", filter.ErrVerificationFailed))
|
||||
}
|
||||
|
||||
func nextFirewallManagedOrderChange(
|
||||
snapshot filter.Snapshot,
|
||||
desiredMarkers []string,
|
||||
) (string, int, bool, error) {
|
||||
expected := make(map[string]struct{}, len(desiredMarkers))
|
||||
for _, marker := range desiredMarkers {
|
||||
expected[marker] = struct{}{}
|
||||
}
|
||||
actual := make([]string, 0, len(desiredMarkers))
|
||||
positions := make([]int, 0, len(desiredMarkers))
|
||||
for index, observed := range snapshot.Rules {
|
||||
if _, exists := expected[observed.Marker]; !exists {
|
||||
continue
|
||||
}
|
||||
actual = append(actual, observed.Marker)
|
||||
position := index + 1
|
||||
if observed.Locator.Position != nil {
|
||||
position = *observed.Locator.Position
|
||||
}
|
||||
positions = append(positions, position)
|
||||
}
|
||||
if len(actual) != len(desiredMarkers) {
|
||||
return "", 0, false, filter.ErrRuleStale
|
||||
}
|
||||
for index := range desiredMarkers {
|
||||
if actual[index] != desiredMarkers[index] {
|
||||
return desiredMarkers[index], positions[index], false, nil
|
||||
}
|
||||
}
|
||||
return "", 0, true, nil
|
||||
}
|
||||
|
||||
func firewallRuleSyncObservedByMarker(snapshot filter.Snapshot, marker string) (filter.ObservedRule, int, bool) {
|
||||
for index, observed := range snapshot.Rules {
|
||||
if observed.Marker == marker {
|
||||
position := index + 1
|
||||
if observed.Locator.Position != nil {
|
||||
position = *observed.Locator.Position
|
||||
}
|
||||
return observed, position, true
|
||||
}
|
||||
}
|
||||
return filter.ObservedRule{}, 0, false
|
||||
}
|
||||
|
||||
func firewallRuleSyncObservedRule(observed filter.ObservedRule) filter.FirewallRule {
|
||||
rule := observed.Rule
|
||||
if rule.UUID == "" && strings.HasPrefix(observed.Marker, "1panel-rule:") {
|
||||
rule.UUID = strings.TrimSpace(strings.TrimPrefix(observed.Marker, "1panel-rule:"))
|
||||
}
|
||||
return rule
|
||||
}
|
||||
|
||||
func (r *firewallScopeReconciler) failOrder(entries []*firewallRuleSyncEntry, err error) {
|
||||
for _, entry := range entries {
|
||||
if entry.outcome != firewallRuleSyncFailed {
|
||||
@@ -1416,7 +1146,7 @@ func firewallRuleSyncChange(
|
||||
if observed.Protected {
|
||||
return filter.DesiredChange{}, false, filter.ErrProtectedRule
|
||||
}
|
||||
before := firewallRuleSyncObservedRule(observed)
|
||||
before := firewallsync.ObservedRule(observed)
|
||||
locator := observed.Locator
|
||||
return filter.DesiredChange{
|
||||
Operation: filter.ChangeDelete, Before: &before, Locator: &locator,
|
||||
@@ -1452,7 +1182,7 @@ func firewallRuleSyncChange(
|
||||
if err := filter.GuardMutation(snapshot, *item.Observed, entry.rule, clientIP, protectedPorts...); err != nil {
|
||||
return filter.DesiredChange{}, false, err
|
||||
}
|
||||
before := firewallRuleSyncObservedRule(*item.Observed)
|
||||
before := firewallsync.ObservedRule(*item.Observed)
|
||||
locator := item.Observed.Locator
|
||||
change.Operation = filter.ChangeUpdate
|
||||
change.Before = &before
|
||||
@@ -1468,29 +1198,16 @@ func firewallRuleSyncInsertionPosition(
|
||||
entries []*firewallRuleSyncEntry,
|
||||
target *firewallRuleSyncEntry,
|
||||
) *int64 {
|
||||
if !firewallProviderHasOrderedManagedRules(snapshot.Scope.Provider) {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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,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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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))))
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+59
@@ -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
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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}$`)
|
||||
)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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}`);
|
||||
|
||||
Reference in New Issue
Block a user