package main import ( "database/sql" "encoding/json" "fmt" "log" "net" "os" "path/filepath" "sort" "strings" ) /* WAF WHITELIST MANAGEMENT */ func updateWAFWhitelist(db *sql.DB) error { log.Println("[WAF Whitelist] Starting update process") if err := syncGitRepo(); err != nil { return fmt.Errorf("git sync failed: %w", err) } wafNetworks, err := loadWAFNetworks(db) if err != nil { return fmt.Errorf("failed to load WAF networks: %w", err) } log.Printf("[WAF Whitelist] Loaded %d WAF networks from DB", len(wafNetworks)) allOrigins, ipToSIDs, sidToDomain, err := loadAllOrigins(db) if err != nil { return fmt.Errorf("failed to load origins: %w", err) } log.Printf("[WAF Whitelist] Loaded %d auto origin IPs", len(allOrigins)) manualOrigins, manualIPToSIDs, manualSIDToDomain, err := loadManualOrigins(db) if err != nil { return fmt.Errorf("failed to load manual origins: %w", err) } log.Printf("[WAF Whitelist] Loaded %d manual origin IPs", len(manualOrigins)) // Объединяем auto и manual for ip, sids := range manualIPToSIDs { ipToSIDs[ip] = append(ipToSIDs[ip], sids...) } for sid, domain := range manualSIDToDomain { sidToDomain[sid] = domain } uniqueIPSet := make(map[string]bool) for _, ip := range allOrigins { uniqueIPSet[ip] = true } for _, ip := range manualOrigins { uniqueIPSet[ip] = true } combined := make([]string, 0, len(uniqueIPSet)) for ip := range uniqueIPSet { combined = append(combined, ip) } sort.Strings(combined) log.Printf("[WAF Whitelist] Combined total: %d unique origin IPs", len(combined)) filteredOrigins := filterNonWAFOrigins(combined, wafNetworks) log.Printf("[WAF Whitelist] Filtered to %d non-WAF origin IPs", len(filteredOrigins)) whitelistPath := filepath.Join(gitRepoPath, whitelistFile) currentWhitelist, err := readWhitelistFile(whitelistPath) if err != nil { return fmt.Errorf("failed to read whitelist file: %w", err) } log.Printf("[WAF Whitelist] Current whitelist contains %d IPs", len(currentWhitelist)) toAdd, toRemove := compareWhitelists(currentWhitelist, filteredOrigins) if len(toAdd) == 0 && len(toRemove) == 0 { log.Println("[WAF Whitelist] No changes needed") return nil } log.Println("[WAF Whitelist] ==================== CHANGES ====================") if len(toAdd) > 0 { log.Printf("[WAF Whitelist] IPs TO ADD (%d):", len(toAdd)) for _, ip := range toAdd { log.Printf("[WAF Whitelist] + %s", ip) } } if len(toRemove) > 0 { log.Printf("[WAF Whitelist] IPs TO REMOVE (%d):", len(toRemove)) for _, ip := range toRemove { log.Printf("[WAF Whitelist] - %s", ip) } } log.Println("[WAF Whitelist] ===================================================") if err := writeWhitelistFile(whitelistPath, filteredOrigins); err != nil { return fmt.Errorf("failed to write whitelist file: %w", err) } log.Println("[WAF Whitelist] File updated successfully") commitMsg := formatCommitMessage(toAdd, toRemove) if err := gitCommitAndPush(commitMsg); err != nil { return fmt.Errorf("git commit/push failed: %w", err) } log.Println("[WAF Whitelist] Changes committed and pushed") telegramMsg := formatTelegramMessage(toAdd, toRemove, ipToSIDs, sidToDomain) if err := sendAlert(telegramMsg); err != nil { log.Printf("[WAF Whitelist] Failed to send Telegram alert: %v", err) } else { log.Println("[WAF Whitelist] Telegram alert sent") } return nil } func loadWAFNetworks(db *sql.DB) ([]*net.IPNet, error) { rows, err := db.Query(`SELECT unnest(waf_networks) FROM ips`) if err != nil { return nil, err } defer rows.Close() var networks []*net.IPNet for rows.Next() { var networkStr string if err := rows.Scan(&networkStr); err != nil { return nil, err } _, ipNet, err := net.ParseCIDR(networkStr) if err != nil { log.Printf("[WAF Networks] Warning: invalid CIDR %s: %v", networkStr, err) continue } networks = append(networks, ipNet) log.Printf("[WAF Networks] Loaded network: %s", networkStr) } return networks, nil } func loadAllOrigins(db *sql.DB) ([]string, map[string][]int64, map[int64]string, error) { rows, err := db.Query(`SELECT sid, domain_name, origins FROM sp_info`) if err != nil { return nil, nil, nil, err } defer rows.Close() uniqueIPs := make(map[string]bool) ipToSIDs := make(map[string][]int64) sidToDomain := make(map[int64]string) for rows.Next() { var sid int64 var domain string var originsJSON []byte if err := rows.Scan(&sid, &domain, &originsJSON); err != nil { return nil, nil, nil, err } sidToDomain[sid] = domain var origins []originItem if err := json.Unmarshal(originsJSON, &origins); err != nil { log.Printf("[Origins] Warning: failed to parse origins JSON for SID %d: %v", sid, err) continue } for _, origin := range origins { uniqueIPs[origin.IP] = true ipToSIDs[origin.IP] = append(ipToSIDs[origin.IP], sid) } } result := make([]string, 0, len(uniqueIPs)) for ip := range uniqueIPs { result = append(result, ip) } sort.Strings(result) return result, ipToSIDs, sidToDomain, nil } func loadManualOrigins(db *sql.DB) ([]string, map[string][]int64, map[int64]string, error) { rows, err := db.Query(`SELECT sid, domain_name, origins FROM manual_info`) if err != nil { return nil, nil, nil, err } defer rows.Close() uniqueIPs := make(map[string]bool) ipToSIDs := make(map[string][]int64) sidToDomain := make(map[int64]string) for rows.Next() { var sid int64 var domain string var originsJSON []byte if err := rows.Scan(&sid, &domain, &originsJSON); err != nil { return nil, nil, nil, err } sidToDomain[sid] = domain var origins []originItem if err := json.Unmarshal(originsJSON, &origins); err != nil { log.Printf("[Manual Origins] Warning: failed to parse origins JSON for SID %d: %v", sid, err) continue } for _, origin := range origins { uniqueIPs[origin.IP] = true ipToSIDs[origin.IP] = append(ipToSIDs[origin.IP], sid) } } result := make([]string, 0, len(uniqueIPs)) for ip := range uniqueIPs { result = append(result, ip) } sort.Strings(result) return result, ipToSIDs, sidToDomain, nil } func filterNonWAFOrigins(origins []string, wafNetworks []*net.IPNet) []string { var filtered []string for _, ipStr := range origins { ip := net.ParseIP(ipStr) if ip == nil { log.Printf("[Filter] Warning: invalid IP address: %s", ipStr) continue } isWAF := false for _, network := range wafNetworks { if network.Contains(ip) { isWAF = true log.Printf("[Filter] Excluding WAF IP: %s (matches network %s)", ipStr, network.String()) break } } if !isWAF { filtered = append(filtered, ipStr) } } return filtered } func readWhitelistFile(path string) ([]string, error) { data, err := os.ReadFile(path) if err != nil { if os.IsNotExist(err) { log.Printf("[Whitelist] File does not exist, will create new one") return []string{}, nil } return nil, err } lines := strings.Split(string(data), "\n") var ips []string for _, line := range lines { line = strings.TrimSpace(line) if line == "" || strings.HasPrefix(line, "#") { continue } ips = append(ips, line) } return ips, nil } func compareWhitelists(current, desired []string) (toAdd, toRemove []string) { currentMap := make(map[string]bool) desiredMap := make(map[string]bool) for _, ip := range current { currentMap[ip] = true } for _, ip := range desired { desiredMap[ip] = true } for _, ip := range desired { if !currentMap[ip] { toAdd = append(toAdd, ip) } } for _, ip := range current { if !desiredMap[ip] { toRemove = append(toRemove, ip) } } sort.Strings(toAdd) sort.Strings(toRemove) return toAdd, toRemove } func writeWhitelistFile(path string, ips []string) error { var content strings.Builder for _, ip := range ips { content.WriteString(ip + "\n") } return os.WriteFile(path, []byte(content.String()), 0644) }