diff --git a/go.mod b/go.mod index dfe7619..7e32442 100644 --- a/go.mod +++ b/go.mod @@ -2,4 +2,7 @@ module github.com/theta42/theta-agent go 1.22.2 -require gopkg.in/yaml.v3 v3.0.1 // indirect +require ( + github.com/gorilla/websocket v1.5.3 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/go.sum b/go.sum index 4bc0337..27b5003 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,5 @@ +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/main.go b/main.go index 8bbc5b8..080a301 100644 --- a/main.go +++ b/main.go @@ -30,8 +30,8 @@ func main() { cfg.Capabilities.ArbitraryBash, ) - // TODO: Initialize WebSocket connection to SSO Manager - // TODO: Start telemetry background loop if capabilities.Telemetry == true + // WebSocket connection to SSO Manager + go connectWebSocket(cfg) // Block until signal is received sigs := make(chan os.Signal, 1) diff --git a/theta-agent b/theta-agent index f6936a7..38292b2 100755 Binary files a/theta-agent and b/theta-agent differ diff --git a/websocket.go b/websocket.go new file mode 100644 index 0000000..67f86d0 --- /dev/null +++ b/websocket.go @@ -0,0 +1,107 @@ +package main + +import ( + "encoding/json" + "log" + "net/url" + "strings" + "time" + + "github.com/gorilla/websocket" +) + +type WSMessage struct { + Type string `json:"type"` + Payload map[string]interface{} `json:"payload"` +} + +func connectWebSocket(cfg *Config) { + for { + // Ensure protocol is ws/wss + serverURL := strings.Replace(cfg.ServerURL, "http://", "ws://", 1) + serverURL = strings.Replace(serverURL, "https://", "wss://", 1) + + u, err := url.Parse(serverURL) + if err != nil { + log.Fatalf("Invalid ServerURL: %v", err) + } + u.Path = "/api/agent/ws" + u.RawQuery = "token=" + cfg.AuthToken + + log.Printf("Connecting to %s", u.String()) + + c, _, err := websocket.DefaultDialer.Dial(u.String(), nil) + if err != nil { + log.Printf("Dial error: %v. Retrying in 5 seconds...", err) + time.Sleep(5 * time.Second) + continue + } + + log.Println("Successfully connected to SSO Manager.") + + // Start telemetry if enabled + var telemetryTicker *time.Ticker + var telemetryDone chan bool + if cfg.Capabilities.Telemetry { + telemetryTicker, telemetryDone = startTelemetry(c, cfg) + } + + // Read loop + for { + _, message, err := c.ReadMessage() + if err != nil { + log.Println("WebSocket read error:", err) + break // break read loop, reconnect + } + + var msg WSMessage + if err := json.Unmarshal(message, &msg); err != nil { + log.Printf("Error unmarshaling message: %v", err) + continue + } + + handleCommand(cfg, msg, c) + } + + // Cleanup on disconnect + c.Close() + if telemetryTicker != nil { + telemetryTicker.Stop() + telemetryDone <- true + } + + log.Println("WebSocket disconnected. Reconnecting in 5 seconds...") + time.Sleep(5 * time.Second) + } +} + +func handleCommand(cfg *Config, msg WSMessage, c *websocket.Conn) { + log.Printf("Received command: %s", msg.Type) + + switch msg.Type { + case "config": + log.Printf("Received config payload: %v", msg.Payload) + case "reboot": + if !cfg.Capabilities.Reboot { + log.Println("Reboot rejected: capability disabled in agent.yml") + return + } + log.Println("Reboot capability enabled. (Simulation: rebooting system...)") + case "service_restart": + serviceName, ok := msg.Payload["service"].(string) + if !ok || !cfg.Capabilities.CanManageService(serviceName) { + log.Printf("Service restart rejected for '%s': not in allowed service list", serviceName) + return + } + log.Printf("Restarting service %s...", serviceName) + default: + log.Printf("Unknown command type: %s", msg.Type) + } +} + +func startTelemetry(c *websocket.Conn, cfg *Config) (*time.Ticker, chan bool) { + ticker := time.Ticker{C: nil} // placeholder logic + done := make(chan bool) + log.Println("Telemetry loop started (placeholder).") + return &ticker, done +}