From 81083c61ac7a6bdc876a928352cd43fe9f98baa3 Mon Sep 17 00:00:00 2001 From: kirillius Date: Wed, 15 Jul 2026 15:50:14 +0300 Subject: [PATCH] initial commit --- .gitignore | 2 + config.go | 33 ++++++++ controller.go | 51 +++++++++++ gateway.go | 23 +++++ go.mod | 3 + main.go | 229 ++++++++++++++++++++++++++++++++++++++++++++++++++ routes.go | 190 +++++++++++++++++++++++++++++++++++++++++ srm.json | 12 +++ 8 files changed, 543 insertions(+) create mode 100644 .gitignore create mode 100644 config.go create mode 100644 controller.go create mode 100644 gateway.go create mode 100644 go.mod create mode 100644 main.go create mode 100644 routes.go create mode 100644 srm.json diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..72c13ae --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +ТЗ.txt +srmd diff --git a/config.go b/config.go new file mode 100644 index 0000000..6a9c0ca --- /dev/null +++ b/config.go @@ -0,0 +1,33 @@ +package main + +import ( + "encoding/json" + "os" +) + +type Gateway struct { + Name string `json:"name"` + Address string `json:"address"` +} + +type Config struct { + Controller string `json:"controller"` + Gateway []Gateway `json:"gateway"` + Metric int `json:"metric"` + Timeout int `json:"timeout"` + Interval int `json:"interval"` + Port int `json:"port"` + Hotswap bool `json:"hotswap"` +} + +func LoadConfig(path string) (*Config, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, err + } + var cfg Config + if err := json.Unmarshal(data, &cfg); err != nil { + return nil, err + } + return &cfg, nil +} diff --git a/controller.go b/controller.go new file mode 100644 index 0000000..5134771 --- /dev/null +++ b/controller.go @@ -0,0 +1,51 @@ +package main + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "time" +) + +func GetUpdate(controllerURL string, timeoutSec int) (int64, error) { + client := &http.Client{Timeout: time.Duration(timeoutSec) * time.Second} + resp, err := client.Get(controllerURL + "/update") + if err != nil { + return 0, err + } + defer resp.Body.Close() + if resp.StatusCode != 200 { + return 0, fmt.Errorf("unexpected status: %d", resp.StatusCode) + } + body, err := io.ReadAll(resp.Body) + if err != nil { + return 0, err + } + var ts int64 + if err := json.Unmarshal(body, &ts); err != nil { + return 0, err + } + return ts, nil +} + +func GetSubnets(controllerURL string, timeoutSec int) ([]string, error) { + client := &http.Client{Timeout: time.Duration(timeoutSec) * time.Second} + resp, err := client.Get(controllerURL + "/subnets") + if err != nil { + return nil, err + } + defer resp.Body.Close() + if resp.StatusCode != 200 { + return nil, fmt.Errorf("unexpected status: %d", resp.StatusCode) + } + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, err + } + var subnets []string + if err := json.Unmarshal(body, &subnets); err != nil { + return nil, err + } + return subnets, nil +} diff --git a/gateway.go b/gateway.go new file mode 100644 index 0000000..da2de9c --- /dev/null +++ b/gateway.go @@ -0,0 +1,23 @@ +package main + +import ( + "fmt" + "os/exec" + "time" +) + +func Ping(host string, timeoutSec int) bool { + cmd := exec.Command("ping", "-c", "1", "-W", fmt.Sprintf("%d", timeoutSec), host) + done := make(chan error, 1) + go func() { + done <- cmd.Run() + }() + + select { + case err := <-done: + return err == nil + case <-time.After(time.Duration(timeoutSec+2) * time.Second): + cmd.Process.Kill() + return false + } +} diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..11bbfd4 --- /dev/null +++ b/go.mod @@ -0,0 +1,3 @@ +module srmd + +go 1.26.5 diff --git a/main.go b/main.go new file mode 100644 index 0000000..6feea35 --- /dev/null +++ b/main.go @@ -0,0 +1,229 @@ +package main + +import ( + "flag" + "fmt" + "html/template" + "log" + "net/http" + "os" + "sync" + "time" +) + +type State struct { + mu sync.RWMutex + GatewayStatus map[string]bool + GatewayIPs map[string]string + SubnetCount int + LastUpdate int64 + ControllerUp bool + Subnets []string +} + +var state State + +const indexHTML = ` + + + + + SRM Daemon + + + +

SRM Daemon

+ + + {{range $name, $up := .GatewayStatus}} + + + + + {{end}} + + + + + + + + + +
ParameterValue
Gateway "{{html $name}}" ({{index $.GatewayIPs $name}}){{if $up}}UP{{else}}DOWN{{end}}
Current Routes{{.SubnetCount}}
Last Update{{if eq .LastUpdate 0}}Never{{else}}{{formatTime .LastUpdate}}{{end}}
Controller{{if .ControllerUp}}UP{{else}}DOWN{{end}}
+ +` + +func formatTime(ts int64) string { + t := time.UnixMilli(ts) + return t.Format("2006-01-02 15:04:05 MST") +} + +func startWebServer(port int) { + funcMap := template.FuncMap{ + "formatTime": formatTime, + } + tmpl, err := template.New("index").Funcs(funcMap).Parse(indexHTML) + if err != nil { + log.Fatalf("Failed to parse template: %v", err) + } + + http.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + state.mu.RLock() + defer state.mu.RUnlock() + w.Header().Set("Content-Type", "text/html; charset=utf-8") + tmpl.Execute(w, &state) + }) + + addr := fmt.Sprintf(":%d", port) + log.Printf("Web interface listening on %s", addr) + if err := http.ListenAndServe(addr, nil); err != nil { + log.Fatalf("Web server failed: %v", err) + } +} + +func main() { + configPath := flag.String("c", "", "Path to config JSON file") + flag.Parse() + + if *configPath == "" { + fmt.Fprintln(os.Stderr, "Usage: srmd -c /path/to/config.json") + os.Exit(1) + } + + cfg, err := LoadConfig(*configPath) + if err != nil { + log.Fatalf("Failed to load config: %v", err) + } + + state.GatewayStatus = make(map[string]bool) + state.GatewayIPs = make(map[string]string) + for _, gw := range cfg.Gateway { + state.GatewayStatus[gw.Name] = false + state.GatewayIPs[gw.Name] = gw.Address + } + + state.Subnets = []string{} + state.LastUpdate = 0 + + go startWebServer(cfg.Port) + + var lastUpdate int64 + var subnets []string + gatewayStates := make(map[string]bool) + for _, gw := range cfg.Gateway { + gatewayStates[gw.Name] = false + } + + for { + log.Println("--- Step 2: Checking gateways ---") + prevStates := make(map[string]bool) + for k, v := range gatewayStates { + prevStates[k] = v + } + + newStates := make(map[string]bool) + for _, gw := range cfg.Gateway { + up := Ping(gw.Address, cfg.Timeout) + newStates[gw.Name] = up + log.Printf("Gateway %q (%s): %v", gw.Name, gw.Address, up) + } + + changed := false + for _, gw := range cfg.Gateway { + if prevStates[gw.Name] != newStates[gw.Name] { + changed = true + break + } + } + gatewayStates = newStates + + state.mu.Lock() + for k, v := range gatewayStates { + state.GatewayStatus[k] = v + } + state.mu.Unlock() + + if changed { + log.Println("Gateway states changed, updating routes (step 3)") + DoRouteUpdate(cfg, gatewayStates, subnets) + } else { + log.Println("Gateway states unchanged, skipping route update") + } + + log.Println("--- Step 4: Getting update from controller ---") + ts, err := GetUpdate(cfg.Controller, cfg.Timeout) + if err != nil { + log.Printf("Controller update failed: %v", err) + state.mu.Lock() + state.ControllerUp = false + state.mu.Unlock() + time.Sleep(time.Duration(cfg.Interval) * time.Second) + continue + } + state.mu.Lock() + state.ControllerUp = true + state.mu.Unlock() + log.Printf("Controller timestamp: %d (lastUpdate: %d)", ts, lastUpdate) + + if ts == lastUpdate { + log.Println("Timestamp unchanged, waiting and retrying") + time.Sleep(time.Duration(cfg.Interval) * time.Second) + continue + } + + newSubnets, err := GetSubnets(cfg.Controller, cfg.Timeout) + if err != nil { + log.Printf("Failed to get subnets: %v", err) + state.mu.Lock() + state.ControllerUp = false + state.mu.Unlock() + time.Sleep(time.Duration(cfg.Interval) * time.Second) + continue + } + state.mu.Lock() + state.ControllerUp = true + state.mu.Unlock() + + subnetsChanged := !stringSliceEqual(subnets, newSubnets) + subnets = newSubnets + + state.mu.Lock() + state.Subnets = make([]string, len(subnets)) + copy(state.Subnets, subnets) + state.SubnetCount = len(subnets) + state.mu.Unlock() + + if subnetsChanged { + log.Printf("Subnets changed (%d subnets), updating lastUpdate", len(subnets)) + lastUpdate = ts + state.mu.Lock() + state.LastUpdate = ts + state.mu.Unlock() + DoRouteUpdate(cfg, gatewayStates, subnets) + } else { + log.Println("Subnets unchanged, waiting") + time.Sleep(time.Duration(cfg.Interval) * time.Second) + } + } +} + +func stringSliceEqual(a, b []string) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} diff --git a/routes.go b/routes.go new file mode 100644 index 0000000..ff7bccd --- /dev/null +++ b/routes.go @@ -0,0 +1,190 @@ +package main + +import ( + "fmt" + "log" + "os/exec" + "strings" +) + +type Route struct { + Subnet string + Gateway string +} + +func ParseRoutes() []Route { + out, err := exec.Command("ip", "route", "list").Output() + if err != nil { + log.Printf("ip route list failed: %v", err) + return nil + } + return parseRouteOutput(string(out)) +} + +func parseRouteOutput(output string) []Route { + var routes []Route + lines := strings.Split(strings.TrimSpace(output), "\n") + for _, line := range lines { + line = strings.TrimSpace(line) + if line == "" { + continue + } + fields := strings.Fields(line) + if len(fields) < 3 { + continue + } + subnet := fields[0] + if subnet == "default" { + subnet = "0.0.0.0/0" + } else if !strings.Contains(subnet, "/") { + subnet = subnet + "/32" + } + + gw := "" + for i := 1; i < len(fields)-1; i++ { + if fields[i] == "via" { + gw = fields[i+1] + break + } + } + if gw != "" { + routes = append(routes, Route{Subnet: subnet, Gateway: gw}) + } + } + return routes +} + +func ipToGatewayName(gateways []Gateway, ip string) (string, bool) { + for _, gw := range gateways { + if gw.Address == ip { + return gw.Name, true + } + } + return "", false +} + +func subnetsContains(subnets []string, subnet string) bool { + for _, s := range subnets { + if s == subnet { + return true + } + } + return false +} + +func routeListHas(routes []Route, subnet, gw string) bool { + for _, r := range routes { + if r.Subnet == subnet && r.Gateway == gw { + return true + } + } + return false +} + +func DoRouteUpdate(cfg *Config, gatewayStates map[string]bool, subnets []string) { + if cfg.Hotswap { + hotswapUpdate(cfg, gatewayStates, subnets) + } else { + failoverUpdate(cfg, gatewayStates, subnets) + } + + finalRoutes := ParseRoutes() + log.Printf("Routes updated, total: %d", len(finalRoutes)) +} + +func hotswapUpdate(cfg *Config, gatewayStates map[string]bool, subnets []string) { + routes := ParseRoutes() + log.Printf("Current system routes: %d", len(routes)) + + for _, r := range routes { + name, known := ipToGatewayName(cfg.Gateway, r.Gateway) + if !known { + continue + } + if !gatewayStates[name] { + log.Printf("Gateway %q is down, deleting route %s via %s", name, r.Subnet, r.Gateway) + delRoute(r.Subnet, r.Gateway) + continue + } + if !subnetsContains(subnets, r.Subnet) { + log.Printf("Subnet %s not in subnets list, deleting route via %s", r.Subnet, r.Gateway) + delRoute(r.Subnet, r.Gateway) + } + } + + for _, subnet := range subnets { + for i, gw := range cfg.Gateway { + if !gatewayStates[gw.Name] { + continue + } + routes = ParseRoutes() + if !routeListHas(routes, subnet, gw.Address) { + metric := cfg.Metric + i + log.Printf("Adding route %s via %s (gateway %q, metric %d)", subnet, gw.Address, gw.Name, metric) + addRoute(subnet, gw.Address, metric) + } + } + } +} + +func failoverUpdate(cfg *Config, gatewayStates map[string]bool, subnets []string) { + routes := ParseRoutes() + log.Printf("Current system routes: %d", len(routes)) + + var primaryGW *Gateway + for _, gw := range cfg.Gateway { + if gatewayStates[gw.Name] { + primaryGW = &gw + log.Printf("Selected primary gateway: %q (%s)", gw.Name, gw.Address) + break + } + } + + for _, r := range routes { + name, known := ipToGatewayName(cfg.Gateway, r.Gateway) + if !known { + continue + } + if primaryGW != nil && r.Gateway == primaryGW.Address { + if subnetsContains(subnets, r.Subnet) { + continue + } + log.Printf("Subnet %s not in subnets, deleting primary route via %s", r.Subnet, r.Gateway) + delRoute(r.Subnet, r.Gateway) + continue + } + log.Printf("Deleting route %s via %s (gateway %q)", r.Subnet, r.Gateway, name) + delRoute(r.Subnet, r.Gateway) + } + + if primaryGW == nil { + log.Println("No active gateway available") + return + } + + for _, subnet := range subnets { + routes = ParseRoutes() + if !routeListHas(routes, subnet, primaryGW.Address) { + log.Printf("Adding route %s via %s (metric %d)", subnet, primaryGW.Address, cfg.Metric) + addRoute(subnet, primaryGW.Address, cfg.Metric) + } + } +} + +func delRoute(subnet, gw string) { + cmd := exec.Command("ip", "route", "del", subnet, "via", gw) + if err := cmd.Run(); err != nil { + log.Printf("Failed to delete route %s via %s: %v", subnet, gw, err) + } else { + log.Printf("Deleted route %s via %s", subnet, gw) + } +} + +func addRoute(subnet, gw string, metric int) { + cmd := exec.Command("ip", "route", "add", subnet, "via", gw, "metric", fmt.Sprintf("%d", metric)) + if err := cmd.Run(); err != nil { + log.Printf("Failed to add route %s via %s: %v", subnet, gw, err) + } else { + log.Printf("Added route %s via %s", subnet, gw) + } +} diff --git a/srm.json b/srm.json new file mode 100644 index 0000000..c17cbbe --- /dev/null +++ b/srm.json @@ -0,0 +1,12 @@ +{ + "controller": "http://localhost:8181/webhook/SRM", + "gateway": [ + {"name":"main", "address":"172.16.217.2"}, + {"name":"backup", "address":"172.16.217.3"} + ], + "metric": 50, + "timeout": 3, + "interval": 10, + "port": 8089, + "hotswap": false +}