From abfdaa4dde1c14a0c23f8cd7e105e550378e860e Mon Sep 17 00:00:00 2001 From: ssongliu Date: Tue, 22 Sep 2026 14:41:05 +0800 Subject: [PATCH] fix: allow updated SSH ports through firewall before restart --- agent/app/service/firewall.go | 41 +---------------- agent/app/service/firewall_utils.go | 69 ++++++++++++++++++----------- agent/app/service/ssh.go | 2 +- agent/utils/terminal/attachment.go | 2 +- 4 files changed, 46 insertions(+), 68 deletions(-) diff --git a/agent/app/service/firewall.go b/agent/app/service/firewall.go index 8e75f0f240c3..e1eff330b18d 100644 --- a/agent/app/service/firewall.go +++ b/agent/app/service/firewall.go @@ -2,7 +2,6 @@ package service import ( "context" - "encoding/json" "errors" "fmt" "slices" @@ -99,45 +98,7 @@ func (s *FirewallService) UpdatePanelPort(ctx context.Context, oldPort, port uin if oldPort == port { return nil } - firewallWhitelistMu.Lock() - defer firewallWhitelistMu.Unlock() - entries, err := loadFirewallPortWhiteList() - if err != nil { - return err - } - panelPorts := make([]firewall.PortWhitelist, 0) - for index := range entries { - if entries[index].Type != firewall.PortWhitelistTypePanel { - continue - } - entries[index].Port = strconv.Itoa(int(port)) - panelPorts = append(panelPorts, entries[index]) - } - if len(panelPorts) == 0 { - panelPorts = append(panelPorts, firewall.PortWhitelist{ - Type: firewall.PortWhitelistTypePanel, Port: strconv.Itoa(int(port)), Protocol: "tcp", - Sources: []string{"0.0.0.0/0", "::/0"}, - }) - entries = append(entries, panelPorts...) - } - entries, err = firewall.ValidatePortWhitelist(entries) - if err != nil { - return err - } - value, err := json.Marshal(entries) - if err != nil { - return err - } - client, err := s.baseClient() - if err != nil && !errors.Is(err, lifecycle.ErrNotInstalled) { - return err - } - if err == nil { - if _, err := s.syncPortWhitelist(ctx, filter.Provider(client.Name()), panelPorts); err != nil { - return err - } - } - return settingRepo.UpdateOrCreate(constant.FirewallPortWhiteList, string(value)) + return s.updateSystemAccessPortWhitelist(ctx, firewall.PortWhitelistTypePanel, []string{strconv.Itoa(int(port))}) } func (s *FirewallService) LoadBaseInfo(chainGroup string) (dto.FirewallSubsystemStatus, error) { diff --git a/agent/app/service/firewall_utils.go b/agent/app/service/firewall_utils.go index f27cbe83ee0d..e87dd0093ab6 100644 --- a/agent/app/service/firewall_utils.go +++ b/agent/app/service/firewall_utils.go @@ -371,39 +371,56 @@ func (s *FirewallService) syncPortWhitelist(ctx context.Context, provider filter return created, errors.Join(failures...) } -func updateSystemAccessPortWhitelist(ctx context.Context, serviceType string, ports []string) error { +func (s *FirewallService) updateSystemAccessPortWhitelist(ctx context.Context, serviceType string, ports []string) error { + if len(ports) == 0 { + return fmt.Errorf("firewall whitelist %s requires a port", serviceType) + } firewallWhitelistMu.Lock() defer firewallWhitelistMu.Unlock() - firewallRuleMutationMu.Lock() - defer firewallRuleMutationMu.Unlock() - return global.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - entries, err := loadPortWhitelistSetting(tx) - if err != nil { - return err + entries, err := loadFirewallPortWhiteList() + if err != nil { + return err + } + servicePorts := make([]firewall.PortWhitelist, 0) + for index := range entries { + if entries[index].Type != serviceType { + continue } - for index := range entries { - if entries[index].Type != serviceType { - continue - } - if len(ports) == 0 { - return fmt.Errorf("firewall whitelist %s requires a port", serviceType) + entries[index].Port = ports[0] + servicePorts = append(servicePorts, entries[index]) + } + if len(servicePorts) == 0 { + servicePorts = append(servicePorts, firewall.PortWhitelist{ + Type: serviceType, Port: ports[0], Protocol: "tcp", + Sources: []string{"0.0.0.0/0", "::/0"}, + }) + entries = append(entries, servicePorts...) + } + entries, err = firewall.ValidatePortWhitelist(entries) + if err != nil { + return err + } + value, err := json.Marshal(entries) + if err != nil { + return err + } + client, err := s.baseClient() + if err != nil && !errors.Is(err, lifecycle.ErrNotInstalled) { + return err + } + if err == nil { + allowances := make([]firewall.PortWhitelist, 0, len(servicePorts)*len(ports)) + for _, rule := range servicePorts { + for _, port := range ports { + rule.Port = port + allowances = append(allowances, rule) } - entries[index].Port = ports[0] - } - entries, err = firewall.ValidatePortWhitelist(entries) - if err != nil { - return err - } - value, err := json.Marshal(entries) - if err != nil { - return err } - err = tx.Where("key = ?", constant.FirewallPortWhiteList).Assign(map[string]interface{}{"value": string(value)}).FirstOrCreate(&model.Setting{Key: constant.FirewallPortWhiteList}).Error - if err != nil { + if _, err := s.syncPortWhitelist(ctx, filter.Provider(client.Name()), allowances); err != nil { return err } - return nil - }) + } + return settingRepo.UpdateOrCreate(constant.FirewallPortWhiteList, string(value)) } func loadPortWhitelistSetting(db *gorm.DB) ([]firewall.PortWhitelist, error) { diff --git a/agent/app/service/ssh.go b/agent/app/service/ssh.go index 7436246d9ee8..d69ef9dfefe4 100644 --- a/agent/app/service/ssh.go +++ b/agent/app/service/ssh.go @@ -228,7 +228,7 @@ func (u *SSHService) Update(req dto.SSHUpdate) error { return err } if req.Key == "Port" { - if err := updateSystemAccessPortWhitelist(context.Background(), firewall.PortWhitelistTypeSSH, splitSSHPorts(req.NewValue)); err != nil { + if err := newFirewallService().updateSystemAccessPortWhitelist(context.Background(), firewall.PortWhitelistTypeSSH, splitSSHPorts(req.NewValue)); err != nil { if restoreErr := rewriteSSHManagedDirectives(sshPath, "Port", buildSSHDirectiveLines("Port", oldPortValue)); restoreErr != nil { return fmt.Errorf("save SSH whitelist: %w; restore SSH configuration: %v", err, restoreErr) } diff --git a/agent/utils/terminal/attachment.go b/agent/utils/terminal/attachment.go index 132405c8b7f2..cb25a0bb01a5 100644 --- a/agent/utils/terminal/attachment.go +++ b/agent/utils/terminal/attachment.go @@ -19,7 +19,7 @@ import ( const ( pingInterval = 30 * time.Second pongWait = 75 * time.Second - writeWait = 5 * time.Second + writeWait = 30 * time.Second ) var errAttachmentClosed = errors.New("terminal attachment is closed")