api_sync/whitelist.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)
}