314 lines
7.8 KiB
Go
314 lines
7.8 KiB
Go
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)
|
|
}
|