mirror of
https://github.com/MengMengCode/CLICD.git
synced 2026-08-08 14:34:49 +08:00
FIX #20
This commit is contained in:
@@ -57,7 +57,11 @@ type trafficStats struct {
|
|||||||
portTotalCounts map[int]int
|
portTotalCounts map[int]int
|
||||||
udpDestCounts map[int]map[string]int
|
udpDestCounts map[int]map[string]int
|
||||||
udpTotalCounts map[int]int
|
udpTotalCounts map[int]int
|
||||||
|
udpDestTotalCounts map[string]int
|
||||||
synSentByDst map[string]int
|
synSentByDst map[string]int
|
||||||
|
tcpSynDestPorts map[string]map[int]int
|
||||||
|
tcpSynPortDestCounts map[int]map[string]int
|
||||||
|
tcpSynPortTotalCounts map[int]int
|
||||||
}
|
}
|
||||||
|
|
||||||
var scanner *SecurityScanner
|
var scanner *SecurityScanner
|
||||||
@@ -239,6 +243,10 @@ func newTrafficStats() *trafficStats {
|
|||||||
udpDestCounts: make(map[int]map[string]int),
|
udpDestCounts: make(map[int]map[string]int),
|
||||||
udpTotalCounts: make(map[int]int),
|
udpTotalCounts: make(map[int]int),
|
||||||
synSentByDst: make(map[string]int),
|
synSentByDst: make(map[string]int),
|
||||||
|
udpDestTotalCounts: make(map[string]int),
|
||||||
|
tcpSynDestPorts: make(map[string]map[int]int),
|
||||||
|
tcpSynPortDestCounts: make(map[int]map[string]int),
|
||||||
|
tcpSynPortTotalCounts: make(map[int]int),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -264,52 +272,64 @@ func (ts *trafficStats) add(conn connEntry) {
|
|||||||
}
|
}
|
||||||
ts.udpDestCounts[conn.dstPort][conn.dstIP]++
|
ts.udpDestCounts[conn.dstPort][conn.dstIP]++
|
||||||
ts.udpTotalCounts[conn.dstPort]++
|
ts.udpTotalCounts[conn.dstPort]++
|
||||||
|
ts.udpDestTotalCounts[conn.dstIP]++
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if conn.state == "SYN_SENT" {
|
if conn.proto == "tcp" && conn.state == "SYN_SENT" {
|
||||||
ts.totalSynSent++
|
ts.totalSynSent++
|
||||||
ts.synSentByDst[conn.dstIP]++
|
ts.synSentByDst[conn.dstIP]++
|
||||||
|
if conn.dstPort > 0 {
|
||||||
|
if ts.tcpSynDestPorts[conn.dstIP] == nil {
|
||||||
|
ts.tcpSynDestPorts[conn.dstIP] = make(map[int]int)
|
||||||
|
}
|
||||||
|
ts.tcpSynDestPorts[conn.dstIP][conn.dstPort]++
|
||||||
|
if ts.tcpSynPortDestCounts[conn.dstPort] == nil {
|
||||||
|
ts.tcpSynPortDestCounts[conn.dstPort] = make(map[string]int)
|
||||||
|
}
|
||||||
|
ts.tcpSynPortDestCounts[conn.dstPort][conn.dstIP]++
|
||||||
|
ts.tcpSynPortTotalCounts[conn.dstPort]++
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ss *SecurityScanner) detectPortScans(name, ip string, stats *trafficStats) {
|
func (ss *SecurityScanner) detectPortScans(name, ip string, stats *trafficStats) {
|
||||||
for dstIP, portCounts := range stats.destPorts {
|
for dstIP, portCounts := range stats.tcpSynDestPorts {
|
||||||
uniquePorts := len(portCounts)
|
uniquePorts := len(portCounts)
|
||||||
switch {
|
switch {
|
||||||
case uniquePorts >= 20:
|
case uniquePorts >= 25:
|
||||||
ss.addAlert(name, "port_scan", "high", ip, dstIP, 0,
|
ss.addAlert(name, "port_scan", "high", ip, dstIP, 0,
|
||||||
fmt.Sprintf("端口扫描: 同一目标 %s 出现 %d 个不同目标端口", dstIP, uniquePorts),
|
fmt.Sprintf("端口扫描: 同一目标 %s 出现 %d 个不同 TCP 半开目标端口", dstIP, uniquePorts),
|
||||||
"")
|
"")
|
||||||
case uniquePorts >= 8:
|
case uniquePorts >= 12:
|
||||||
ss.addAlert(name, "port_scan", "medium", ip, dstIP, 0,
|
ss.addAlert(name, "port_scan", "medium", ip, dstIP, 0,
|
||||||
fmt.Sprintf("可疑端口探测: 同一目标 %s 出现 %d 个不同目标端口", dstIP, uniquePorts),
|
fmt.Sprintf("可疑端口探测: 同一目标 %s 出现 %d 个不同 TCP 半开目标端口", dstIP, uniquePorts),
|
||||||
"")
|
"")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
for port, targets := range stats.portDestCounts {
|
for port, targets := range stats.tcpSynPortDestCounts {
|
||||||
uniqueTargets := len(targets)
|
uniqueTargets := len(targets)
|
||||||
if service, ok := bruteForcePorts[port]; ok {
|
if service, ok := bruteForcePorts[port]; ok {
|
||||||
if uniqueTargets >= 30 {
|
if uniqueTargets >= 30 {
|
||||||
ss.addAlert(name, "brute_force", "critical", ip, "*", port,
|
ss.addAlert(name, "brute_force", "critical", ip, "*", port,
|
||||||
fmt.Sprintf("横向爆破: 目标服务 %s(%d) 覆盖 %d 个不同 IP", service, port, uniqueTargets),
|
fmt.Sprintf("横向爆破: 目标服务 %s(%d) 出现 TCP 半开连接并覆盖 %d 个不同 IP", service, port, uniqueTargets),
|
||||||
"")
|
"")
|
||||||
} else if uniqueTargets >= 10 {
|
} else if uniqueTargets >= 12 {
|
||||||
ss.addAlert(name, "brute_force", "high", ip, "*", port,
|
ss.addAlert(name, "brute_force", "high", ip, "*", port,
|
||||||
fmt.Sprintf("疑似横向爆破: 目标服务 %s(%d) 覆盖 %d 个不同 IP", service, port, uniqueTargets),
|
fmt.Sprintf("疑似横向爆破: 目标服务 %s(%d) 出现 TCP 半开连接并覆盖 %d 个不同 IP", service, port, uniqueTargets),
|
||||||
"")
|
"")
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if uniqueTargets >= 40 {
|
if uniqueTargets >= 50 {
|
||||||
ss.addAlert(name, "horizontal_scan", "high", ip, "*", port,
|
ss.addAlert(name, "horizontal_scan", "high", ip, "*", port,
|
||||||
fmt.Sprintf("横向扫描: 同一端口 %d 覆盖 %d 个不同目标", port, uniqueTargets),
|
fmt.Sprintf("横向扫描: 同一 TCP 端口 %d 出现半开连接并覆盖 %d 个不同目标", port, uniqueTargets),
|
||||||
"")
|
"")
|
||||||
} else if uniqueTargets >= 15 {
|
} else if uniqueTargets >= 20 {
|
||||||
ss.addAlert(name, "horizontal_scan", "medium", ip, "*", port,
|
ss.addAlert(name, "horizontal_scan", "medium", ip, "*", port,
|
||||||
fmt.Sprintf("可疑横向探测: 同一端口 %d 覆盖 %d 个不同目标", port, uniqueTargets),
|
fmt.Sprintf("可疑横向探测: 同一 TCP 端口 %d 出现半开连接并覆盖 %d 个不同目标", port, uniqueTargets),
|
||||||
"")
|
"")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -323,13 +343,25 @@ func (ss *SecurityScanner) detectBruteForce(name, ip string, stats *trafficStats
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if count >= 20 {
|
synCount := 0
|
||||||
|
if ports := stats.tcpSynDestPorts[dstIP]; ports != nil {
|
||||||
|
synCount = ports[port]
|
||||||
|
}
|
||||||
|
if synCount >= 25 {
|
||||||
ss.addAlert(name, "brute_force", "critical", ip, dstIP, port,
|
ss.addAlert(name, "brute_force", "critical", ip, dstIP, port,
|
||||||
fmt.Sprintf("暴力破解: %s(%d) 当前连接数 %d", service, port, count),
|
fmt.Sprintf("暴力破解: %s(%d) 当前 TCP 半开连接 %d 条", service, port, synCount),
|
||||||
"")
|
"")
|
||||||
} else if count >= 10 {
|
} else if synCount >= 12 {
|
||||||
ss.addAlert(name, "brute_force", "high", ip, dstIP, port,
|
ss.addAlert(name, "brute_force", "high", ip, dstIP, port,
|
||||||
fmt.Sprintf("疑似暴力破解: %s(%d) 当前连接数 %d", service, port, count),
|
fmt.Sprintf("疑似暴力破解: %s(%d) 当前 TCP 半开连接 %d 条", service, port, synCount),
|
||||||
|
"")
|
||||||
|
} else if count >= 60 {
|
||||||
|
ss.addAlert(name, "brute_force", "critical", ip, dstIP, port,
|
||||||
|
fmt.Sprintf("暴力破解: %s(%d) 当前连接数 %d 条", service, port, count),
|
||||||
|
"")
|
||||||
|
} else if count >= 30 {
|
||||||
|
ss.addAlert(name, "brute_force", "high", ip, dstIP, port,
|
||||||
|
fmt.Sprintf("疑似暴力破解: %s(%d) 当前连接数 %d 条", service, port, count),
|
||||||
"")
|
"")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -356,30 +388,41 @@ func (ss *SecurityScanner) detectSpam(name, ip string, stats *trafficStats) {
|
|||||||
func (ss *SecurityScanner) detectMassAbuse(name, ip string, stats *trafficStats) {
|
func (ss *SecurityScanner) detectMassAbuse(name, ip string, stats *trafficStats) {
|
||||||
targets := len(stats.destCounts)
|
targets := len(stats.destCounts)
|
||||||
switch {
|
switch {
|
||||||
case targets >= 100:
|
case targets >= 120 && stats.total >= 600:
|
||||||
ss.addAlert(name, "ddos", "critical", ip, "*", 0,
|
ss.addAlert(name, "ddos", "critical", ip, "*", 0,
|
||||||
fmt.Sprintf("大规模对外连接: 当前覆盖 %d 个不同目标", targets),
|
fmt.Sprintf("大规模对外连接: 当前 conntrack 出站记录 %d 条,覆盖 %d 个不同目标", stats.total, targets),
|
||||||
"")
|
"")
|
||||||
case targets >= 35:
|
case targets >= 60 && stats.total >= 300:
|
||||||
ss.addAlert(name, "ddos", "high", ip, "*", 0,
|
ss.addAlert(name, "ddos", "high", ip, "*", 0,
|
||||||
fmt.Sprintf("大量对外连接: 当前覆盖 %d 个不同目标", targets),
|
fmt.Sprintf("大量对外连接: 当前 conntrack 出站记录 %d 条,覆盖 %d 个不同目标", stats.total, targets),
|
||||||
"")
|
"")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
synTargets := len(stats.synSentByDst)
|
||||||
switch {
|
switch {
|
||||||
case stats.total >= 500:
|
case stats.totalSynSent >= 250 || (synTargets >= 80 && stats.totalSynSent >= 160):
|
||||||
ss.addAlert(name, "ddos", "critical", ip, "*", 0,
|
ss.addAlert(name, "ddos", "critical", ip, "*", 0,
|
||||||
fmt.Sprintf("异常大量连接: 当前 conntrack 出站记录 %d 条", stats.total),
|
fmt.Sprintf("大量半开连接: 当前 TCP SYN_SENT %d 条,覆盖 %d 个不同目标", stats.totalSynSent, synTargets),
|
||||||
"")
|
"")
|
||||||
case stats.total >= 200:
|
case stats.totalSynSent >= 100 || (synTargets >= 35 && stats.totalSynSent >= 70):
|
||||||
ss.addAlert(name, "ddos", "high", ip, "*", 0,
|
ss.addAlert(name, "ddos", "high", ip, "*", 0,
|
||||||
fmt.Sprintf("高连接数: 当前 conntrack 出站记录 %d 条", stats.total),
|
fmt.Sprintf("可疑大量半开连接: 当前 TCP SYN_SENT %d 条,覆盖 %d 个不同目标", stats.totalSynSent, synTargets),
|
||||||
"")
|
"")
|
||||||
}
|
}
|
||||||
|
|
||||||
if stats.totalSynSent >= 100 {
|
udpTargets := len(stats.udpDestTotalCounts)
|
||||||
|
udpTotal := 0
|
||||||
|
for _, count := range stats.udpTotalCounts {
|
||||||
|
udpTotal += count
|
||||||
|
}
|
||||||
|
switch {
|
||||||
|
case udpTargets >= 120 && udpTotal >= 300:
|
||||||
ss.addAlert(name, "ddos", "critical", ip, "*", 0,
|
ss.addAlert(name, "ddos", "critical", ip, "*", 0,
|
||||||
fmt.Sprintf("大量半开连接: 当前 SYN_SENT %d 条", stats.totalSynSent),
|
fmt.Sprintf("UDP 大规模外发: 当前 UDP 连接 %d 条,覆盖 %d 个不同目标", udpTotal, udpTargets),
|
||||||
|
"")
|
||||||
|
case udpTargets >= 50 && udpTotal >= 120:
|
||||||
|
ss.addAlert(name, "ddos", "high", ip, "*", 0,
|
||||||
|
fmt.Sprintf("可疑 UDP 大规模外发: 当前 UDP 连接 %d 条,覆盖 %d 个不同目标", udpTotal, udpTargets),
|
||||||
"")
|
"")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -404,11 +447,18 @@ func (ss *SecurityScanner) detectReflectionAbuse(name, ip string, stats *traffic
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if targets >= 30 || total >= 100 {
|
criticalTargets, criticalTotal := 40, 120
|
||||||
|
highTargets, highTotal := 15, 45
|
||||||
|
if port == 53 {
|
||||||
|
criticalTargets, criticalTotal = 75, 300
|
||||||
|
highTargets, highTotal = 25, 100
|
||||||
|
}
|
||||||
|
|
||||||
|
if targets >= criticalTargets && total >= criticalTotal {
|
||||||
ss.addAlert(name, "reflection", "critical", ip, "*", port,
|
ss.addAlert(name, "reflection", "critical", ip, "*", port,
|
||||||
fmt.Sprintf("UDP 反射放大: %s(%d) 当前 UDP 连接 %d 条,覆盖 %d 个目标", service, port, total, targets),
|
fmt.Sprintf("UDP 反射放大: %s(%d) 当前 UDP 连接 %d 条,覆盖 %d 个目标", service, port, total, targets),
|
||||||
"")
|
"")
|
||||||
} else if targets >= 10 || total >= 30 {
|
} else if targets >= highTargets && total >= highTotal {
|
||||||
ss.addAlert(name, "reflection", "high", ip, "*", port,
|
ss.addAlert(name, "reflection", "high", ip, "*", port,
|
||||||
fmt.Sprintf("疑似 UDP 反射放大: %s(%d) 当前 UDP 连接 %d 条,覆盖 %d 个目标", service, port, total, targets),
|
fmt.Sprintf("疑似 UDP 反射放大: %s(%d) 当前 UDP 连接 %d 条,覆盖 %d 个目标", service, port, total, targets),
|
||||||
"")
|
"")
|
||||||
@@ -645,6 +695,9 @@ func severityRank(severity string) int {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func autoShutdownAlertContainer(containerName, alertType, severity string) {
|
func autoShutdownAlertContainer(containerName, alertType, severity string) {
|
||||||
|
if !config.AppConfig.SecurityAutoShutdown {
|
||||||
|
return
|
||||||
|
}
|
||||||
c := config.FindContainerByName(containerName)
|
c := config.FindContainerByName(containerName)
|
||||||
if c == nil || c.Status != "running" {
|
if c == nil || c.Status != "running" {
|
||||||
return
|
return
|
||||||
@@ -660,6 +713,24 @@ func autoShutdownAlertContainer(containerName, alertType, severity string) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func clearSecurityPolicyBlocks() int {
|
||||||
|
cleared := 0
|
||||||
|
for i := range config.AppConfig.Containers {
|
||||||
|
c := &config.AppConfig.Containers[i]
|
||||||
|
if !c.PolicyBlocked || !isSecurityPolicyBlockReason(c.PolicyBlockedReason) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
config.SetContainerPolicyBlock(c.ID, false, "")
|
||||||
|
config.AddAuditLog("security_policy_unblock", c.Name, "关闭安全告警自动关机后解除策略临时封禁", "system")
|
||||||
|
cleared++
|
||||||
|
}
|
||||||
|
return cleared
|
||||||
|
}
|
||||||
|
|
||||||
|
func isSecurityPolicyBlockReason(reason string) bool {
|
||||||
|
return strings.Contains(reason, "告警触发策略临时封禁")
|
||||||
|
}
|
||||||
|
|
||||||
// HandleSecurityAlerts returns all security alerts.
|
// HandleSecurityAlerts returns all security alerts.
|
||||||
func HandleSecurityAlerts(w http.ResponseWriter, r *http.Request) {
|
func HandleSecurityAlerts(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.Method != http.MethodGet {
|
if r.Method != http.MethodGet {
|
||||||
@@ -699,9 +770,17 @@ func HandleSecuritySettings(w http.ResponseWriter, r *http.Request) {
|
|||||||
jsonResponse(w, http.StatusInternalServerError, APIResponse{Success: false, Message: err.Error()})
|
jsonResponse(w, http.StatusInternalServerError, APIResponse{Success: false, Message: err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
cancelledTasks := 0
|
||||||
|
clearedBlocks := 0
|
||||||
|
if !req.AutoShutdown {
|
||||||
|
cancelledTasks = globalQueue.CancelPendingSecurityStops()
|
||||||
|
clearedBlocks = clearSecurityPolicyBlocks()
|
||||||
|
}
|
||||||
auditRequest(r, "security.settings", "auto_shutdown", fmt.Sprintf("auto_shutdown=%v", req.AutoShutdown), true, "")
|
auditRequest(r, "security.settings", "auto_shutdown", fmt.Sprintf("auto_shutdown=%v", req.AutoShutdown), true, "")
|
||||||
jsonResponse(w, http.StatusOK, APIResponse{Success: true, Data: map[string]bool{
|
jsonResponse(w, http.StatusOK, APIResponse{Success: true, Data: map[string]interface{}{
|
||||||
"auto_shutdown": config.AppConfig.SecurityAutoShutdown,
|
"auto_shutdown": config.AppConfig.SecurityAutoShutdown,
|
||||||
|
"cancelled_tasks": cancelledTasks,
|
||||||
|
"cleared_blocks": clearedBlocks,
|
||||||
}})
|
}})
|
||||||
default:
|
default:
|
||||||
jsonResponse(w, http.StatusMethodNotAllowed, APIResponse{Success: false, Message: "Method not allowed"})
|
jsonResponse(w, http.StatusMethodNotAllowed, APIResponse{Success: false, Message: "Method not allowed"})
|
||||||
|
|||||||
@@ -0,0 +1,148 @@
|
|||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"clicd/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDetectReflectionAbuseIgnoresSingleDNSResolver(t *testing.T) {
|
||||||
|
resetSecurityTestConfig()
|
||||||
|
|
||||||
|
stats := newTrafficStats()
|
||||||
|
for i := 0; i < 180; i++ {
|
||||||
|
stats.add(connEntry{
|
||||||
|
dstIP: "1.1.1.1",
|
||||||
|
dstPort: 53,
|
||||||
|
proto: "udp",
|
||||||
|
state: "UNREPLIED",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
ss := newSecurityScanner()
|
||||||
|
ss.detectReflectionAbuse("ct-dns", "10.0.0.2", stats)
|
||||||
|
|
||||||
|
if len(ss.alerts) != 0 {
|
||||||
|
t.Fatalf("normal DNS queries to one resolver should not trigger reflection alert: %+v", ss.alerts)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDetectReflectionAbuseFlagsWideDNSFanout(t *testing.T) {
|
||||||
|
resetSecurityTestConfig()
|
||||||
|
|
||||||
|
stats := newTrafficStats()
|
||||||
|
for i := 0; i < 120; i++ {
|
||||||
|
stats.add(connEntry{
|
||||||
|
dstIP: fmt.Sprintf("203.0.113.%d", i),
|
||||||
|
dstPort: 53,
|
||||||
|
proto: "udp",
|
||||||
|
state: "UNREPLIED",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
ss := newSecurityScanner()
|
||||||
|
ss.detectReflectionAbuse("ct-dns", "10.0.0.2", stats)
|
||||||
|
|
||||||
|
if len(ss.alerts) != 1 {
|
||||||
|
t.Fatalf("expected one reflection alert, got %+v", ss.alerts)
|
||||||
|
}
|
||||||
|
if got := ss.alerts[0].Type; got != "reflection" {
|
||||||
|
t.Fatalf("expected reflection alert, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDetectPortScansUsesHalfOpenConnections(t *testing.T) {
|
||||||
|
resetSecurityTestConfig()
|
||||||
|
|
||||||
|
established := newTrafficStats()
|
||||||
|
for port := 8000; port < 8020; port++ {
|
||||||
|
established.add(connEntry{
|
||||||
|
dstIP: "198.51.100.10",
|
||||||
|
dstPort: port,
|
||||||
|
proto: "tcp",
|
||||||
|
state: "ESTABLISHED",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
ss := newSecurityScanner()
|
||||||
|
ss.detectPortScans("ct-web", "10.0.0.3", established)
|
||||||
|
if len(ss.alerts) != 0 {
|
||||||
|
t.Fatalf("established multi-port connections should not trigger port scan alert: %+v", ss.alerts)
|
||||||
|
}
|
||||||
|
|
||||||
|
halfOpen := newTrafficStats()
|
||||||
|
for port := 8000; port < 8012; port++ {
|
||||||
|
halfOpen.add(connEntry{
|
||||||
|
dstIP: "198.51.100.10",
|
||||||
|
dstPort: port,
|
||||||
|
proto: "tcp",
|
||||||
|
state: "SYN_SENT",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
ss.detectPortScans("ct-web", "10.0.0.3", halfOpen)
|
||||||
|
if len(ss.alerts) != 1 {
|
||||||
|
t.Fatalf("expected one port scan alert, got %+v", ss.alerts)
|
||||||
|
}
|
||||||
|
if got := ss.alerts[0].Type; got != "port_scan" {
|
||||||
|
t.Fatalf("expected port_scan alert, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCancelPendingSecurityStops(t *testing.T) {
|
||||||
|
resetSecurityTestConfig()
|
||||||
|
|
||||||
|
q := &TaskQueue{
|
||||||
|
tasks: map[string]*Task{},
|
||||||
|
}
|
||||||
|
securityTask := &Task{
|
||||||
|
ID: "task-1",
|
||||||
|
Type: TaskStop,
|
||||||
|
ContainerID: 1,
|
||||||
|
Status: "pending",
|
||||||
|
User: "system:security",
|
||||||
|
}
|
||||||
|
userTask := &Task{
|
||||||
|
ID: "task-2",
|
||||||
|
Type: TaskStop,
|
||||||
|
ContainerID: 2,
|
||||||
|
Status: "pending",
|
||||||
|
User: "admin",
|
||||||
|
}
|
||||||
|
runningSecurityTask := &Task{
|
||||||
|
ID: "task-3",
|
||||||
|
Type: TaskStop,
|
||||||
|
ContainerID: 3,
|
||||||
|
Status: "running",
|
||||||
|
User: "system:security",
|
||||||
|
}
|
||||||
|
q.tasks[securityTask.ID] = securityTask
|
||||||
|
q.tasks[userTask.ID] = userTask
|
||||||
|
q.tasks[runningSecurityTask.ID] = runningSecurityTask
|
||||||
|
q.opQueue = []*Task{securityTask, userTask, runningSecurityTask}
|
||||||
|
|
||||||
|
if got := q.CancelPendingSecurityStops(); got != 1 {
|
||||||
|
t.Fatalf("expected one pending security stop to be cancelled, got %d", got)
|
||||||
|
}
|
||||||
|
if _, ok := q.tasks[securityTask.ID]; ok {
|
||||||
|
t.Fatal("pending security stop task was not removed")
|
||||||
|
}
|
||||||
|
if _, ok := q.tasks[userTask.ID]; !ok {
|
||||||
|
t.Fatal("user stop task should not be removed")
|
||||||
|
}
|
||||||
|
if _, ok := q.tasks[runningSecurityTask.ID]; !ok {
|
||||||
|
t.Fatal("running security stop task should be left for worker-side skip")
|
||||||
|
}
|
||||||
|
if len(q.opQueue) != 2 {
|
||||||
|
t.Fatalf("expected op queue to keep two tasks, got %d", len(q.opQueue))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func resetSecurityTestConfig() {
|
||||||
|
config.AppConfig = &config.ClicdConfig{
|
||||||
|
Containers: []config.Container{},
|
||||||
|
AuditLogs: []config.AuditLog{},
|
||||||
|
Tasks: []config.SavedTask{},
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -213,6 +213,10 @@ func (q *TaskQueue) enqueueSingleWithAudit(containerID int, containerName string
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (q *TaskQueue) EnqueueSecurityStop(containerID int, containerName string) (string, bool) {
|
func (q *TaskQueue) EnqueueSecurityStop(containerID int, containerName string) (string, bool) {
|
||||||
|
if !config.AppConfig.SecurityAutoShutdown {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
q.mu.Lock()
|
q.mu.Lock()
|
||||||
defer q.mu.Unlock()
|
defer q.mu.Unlock()
|
||||||
|
|
||||||
@@ -230,6 +234,34 @@ func (q *TaskQueue) EnqueueSecurityStop(containerID int, containerName string) (
|
|||||||
return taskID, true
|
return taskID, true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (q *TaskQueue) CancelPendingSecurityStops() int {
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
|
||||||
|
cancelled := 0
|
||||||
|
newOpQueue := make([]*Task, 0, len(q.opQueue))
|
||||||
|
for _, task := range q.opQueue {
|
||||||
|
if isSecurityStopTask(task) && task.Status == "pending" {
|
||||||
|
delete(q.tasks, task.ID)
|
||||||
|
cancelled++
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
newOpQueue = append(newOpQueue, task)
|
||||||
|
}
|
||||||
|
q.opQueue = newOpQueue
|
||||||
|
|
||||||
|
for id, task := range q.tasks {
|
||||||
|
if isSecurityStopTask(task) && task.Status == "pending" {
|
||||||
|
delete(q.tasks, id)
|
||||||
|
cancelled++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if cancelled > 0 {
|
||||||
|
q.persistTasks()
|
||||||
|
}
|
||||||
|
return cancelled
|
||||||
|
}
|
||||||
|
|
||||||
// createWorker handles TaskCreate: lxc-create, resource setup, start, and SSH init.
|
// createWorker handles TaskCreate: lxc-create, resource setup, start, and SSH init.
|
||||||
// If a restored task already has a same-name container in config, it resumes
|
// If a restored task already has a same-name container in config, it resumes
|
||||||
// initialization instead of creating another ct-{id}.
|
// initialization instead of creating another ct-{id}.
|
||||||
@@ -324,6 +356,7 @@ func (q *TaskQueue) opWorker() {
|
|||||||
q.mu.Unlock()
|
q.mu.Unlock()
|
||||||
|
|
||||||
var err error
|
var err error
|
||||||
|
skipped := false
|
||||||
err = resolveTaskContainer(task)
|
err = resolveTaskContainer(task)
|
||||||
// Block operations on expired or traffic-exceeded containers (except stop/delete)
|
// Block operations on expired or traffic-exceeded containers (except stop/delete)
|
||||||
if err == nil && (task.Type == TaskStart || task.Type == TaskRestart || task.Type == TaskReinstall) {
|
if err == nil && (task.Type == TaskStart || task.Type == TaskRestart || task.Type == TaskReinstall) {
|
||||||
@@ -336,7 +369,11 @@ func (q *TaskQueue) opWorker() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if err == nil && isSecurityStopTask(task) && !config.AppConfig.SecurityAutoShutdown {
|
||||||
|
skipped = true
|
||||||
|
}
|
||||||
if err == nil {
|
if err == nil {
|
||||||
|
if !skipped {
|
||||||
switch task.Type {
|
switch task.Type {
|
||||||
case TaskStart:
|
case TaskStart:
|
||||||
err = startByRuntime(task.ContainerID)
|
err = startByRuntime(task.ContainerID)
|
||||||
@@ -360,6 +397,7 @@ func (q *TaskQueue) opWorker() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
q.mu.Lock()
|
q.mu.Lock()
|
||||||
auditUser := task.User
|
auditUser := task.User
|
||||||
@@ -370,6 +408,9 @@ func (q *TaskQueue) opWorker() {
|
|||||||
task.Status = "failed"
|
task.Status = "failed"
|
||||||
task.Error = err.Error()
|
task.Error = err.Error()
|
||||||
config.AddAuditLogFull(string(task.Type), task.ContainerName, "失败: "+err.Error(), auditUser, task.IP, task.UserAgent, false, err.Error())
|
config.AddAuditLogFull(string(task.Type), task.ContainerName, "失败: "+err.Error(), auditUser, task.IP, task.UserAgent, false, err.Error())
|
||||||
|
} else if skipped {
|
||||||
|
task.Status = "done"
|
||||||
|
config.AddAuditLogFull(string(task.Type), task.ContainerName, "跳过: 安全告警自动关机已关闭", auditUser, task.IP, task.UserAgent, true, "")
|
||||||
} else {
|
} else {
|
||||||
task.Status = "done"
|
task.Status = "done"
|
||||||
config.AddAuditLogFull(string(task.Type), task.ContainerName, "成功", auditUser, task.IP, task.UserAgent, true, "")
|
config.AddAuditLogFull(string(task.Type), task.ContainerName, "成功", auditUser, task.IP, task.UserAgent, true, "")
|
||||||
@@ -391,6 +432,10 @@ func (q *TaskQueue) opWorker() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func isSecurityStopTask(task *Task) bool {
|
||||||
|
return task != nil && task.Type == TaskStop && task.User == "system:security"
|
||||||
|
}
|
||||||
|
|
||||||
func clearPolicyBlockAfterAdminRecovery(task *Task) {
|
func clearPolicyBlockAfterAdminRecovery(task *Task) {
|
||||||
if task == nil || strings.HasPrefix(task.User, "user:") || task.User == "system:security" {
|
if task == nil || strings.HasPrefix(task.User, "user:") || task.User == "system:security" {
|
||||||
return
|
return
|
||||||
@@ -817,6 +862,9 @@ func HandleTasks(w http.ResponseWriter, r *http.Request) {
|
|||||||
// RestoreTasks restores task queue from config
|
// RestoreTasks restores task queue from config
|
||||||
func RestoreTasks() {
|
func RestoreTasks() {
|
||||||
for _, st := range config.AppConfig.Tasks {
|
for _, st := range config.AppConfig.Tasks {
|
||||||
|
if st.Type == string(TaskStop) && st.User == "system:security" && !config.AppConfig.SecurityAutoShutdown {
|
||||||
|
continue
|
||||||
|
}
|
||||||
var cfg lxc.ContainerConfig
|
var cfg lxc.ContainerConfig
|
||||||
if st.Config != "" {
|
if st.Config != "" {
|
||||||
json.Unmarshal([]byte(st.Config), &cfg)
|
json.Unmarshal([]byte(st.Config), &cfg)
|
||||||
|
|||||||
Reference in New Issue
Block a user