Files
AetherForge/server/internal/api/websocket_test.go

315 lines
9.0 KiB
Go

package api
import (
"encoding/base64"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"crypto-miner-server/internal/db"
"crypto-miner-server/internal/models"
"crypto-miner-server/internal/pool"
"github.com/gorilla/websocket"
)
const wsDefaultPingIntervalSec = 30
func resetWSAuthUsers(t *testing.T, user, pass string) {
t.Helper()
hashed, err := hashPassword(pass)
if err != nil {
t.Fatal(err)
}
usersMu.Lock()
authUsers = map[string]string{user: hashed}
usersMu.Unlock()
t.Cleanup(func() {
usersMu.Lock()
authUsers = map[string]string{}
usersMu.Unlock()
})
}
func wsDashboardToken(user, pass string) string {
return base64.StdEncoding.EncodeToString([]byte(user + ":" + pass))
}
func dialAgentWS(t *testing.T, hub *WSHub) (*websocket.Conn, string) {
t.Helper()
srv := httptest.NewServer(http.HandlerFunc(hub.HandleAgentWS))
t.Cleanup(srv.Close)
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http")
conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
if err != nil {
t.Fatalf("dial: %v", err)
}
t.Cleanup(func() { _ = conn.Close() })
return conn, wsURL
}
func authAgentConn(t *testing.T, conn *websocket.Conn, payload map[string]interface{}) Message {
t.Helper()
data, _ := json.Marshal(payload)
if err := conn.WriteJSON(Message{Type: "auth", Payload: data}); err != nil {
t.Fatal(err)
}
var resp Message
if err := conn.ReadJSON(&resp); err != nil {
t.Fatalf("read auth_response: %v", err)
}
return resp
}
func TestWSHubPingIntervalConstants(t *testing.T) {
hub := NewWSHub(nil)
if hub.pingIntervalSec != wsDefaultPingIntervalSec {
t.Fatalf("default ping interval = %d", hub.pingIntervalSec)
}
hub.SetPingInterval(5)
if hub.pingIntervalSec != wsDefaultPingIntervalSec {
t.Fatalf("below-minimum ping should clamp to %d, got %d", wsDefaultPingIntervalSec, hub.pingIntervalSec)
}
hub.SetPingInterval(15)
if hub.pingIntervalSec != 15 {
t.Fatalf("expected 15, got %d", hub.pingIntervalSec)
}
if hub.pingInterval().Seconds() != 15 {
t.Fatalf("pingInterval duration = %v", hub.pingInterval())
}
}
func TestCheckDashboardWSTokenBcryptUser(t *testing.T) {
resetWSAuthUsers(t, "dash", "secret-pass")
req := httptest.NewRequest(http.MethodGet, "/ws/dashboard?token="+wsDashboardToken("dash", "secret-pass"), nil)
if !checkDashboardWSToken(req) {
t.Fatal("valid bcrypt user token should pass")
}
req = httptest.NewRequest(http.MethodGet, "/ws/dashboard?token="+wsDashboardToken("dash", "wrong"), nil)
if checkDashboardWSToken(req) {
t.Fatal("wrong password should fail")
}
req = httptest.NewRequest(http.MethodGet, "/ws/dashboard", nil)
if checkDashboardWSToken(req) {
t.Fatal("missing token should fail")
}
}
func TestHandleDashboardWSUnauthorized(t *testing.T) {
database, err := db.New(t.TempDir())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
hub := NewWSHub(database)
srv := httptest.NewServer(http.HandlerFunc(hub.HandleDashboardWS))
t.Cleanup(srv.Close)
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http")
_, resp, err := websocket.DefaultDialer.Dial(wsURL, nil)
if err == nil {
t.Fatal("expected dial failure without token")
}
if resp == nil || resp.StatusCode != http.StatusUnauthorized {
t.Fatalf("expected 401 upgrade rejection, got err=%v status=%v", err, resp)
}
}
func TestHandleDashboardWSAuthorizedInit(t *testing.T) {
resetWSAuthUsers(t, testAuthUser, testAuthPass)
database, err := db.New(t.TempDir())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
hub := NewWSHub(database)
srv := httptest.NewServer(http.HandlerFunc(hub.HandleDashboardWS))
t.Cleanup(srv.Close)
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "?token=" + wsDashboardToken(testAuthUser, testAuthPass)
conn, resp, err := websocket.DefaultDialer.Dial(wsURL, nil)
if err != nil {
t.Fatalf("dial: %v status=%v", err, resp)
}
t.Cleanup(func() { _ = conn.Close() })
var msg Message
if err := conn.ReadJSON(&msg); err != nil {
t.Fatalf("read init: %v", err)
}
if msg.Type != "init" {
t.Fatalf("expected init message, got %q", msg.Type)
}
}
func TestHandleAgentWSBadFleetSecret(t *testing.T) {
database, err := db.New(t.TempDir())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
hub := NewWSHub(database)
hub.SetFleetSecret("required-secret")
conn, _ := dialAgentWS(t, hub)
resp := authAgentConn(t, conn, map[string]interface{}{
"agent_id": "agent-bad-secret", "fleet_secret": "wrong", "hostname": "host",
})
var body map[string]interface{}
if err := json.Unmarshal(resp.Payload, &body); err != nil {
t.Fatal(err)
}
if body["success"] != false {
t.Fatalf("expected auth failure, got %+v", body)
}
if hub.isAgentConnected("agent-bad-secret") {
t.Fatal("agent should not register with bad fleet secret")
}
}
func TestHandleAgentWSStatsAndLogTail(t *testing.T) {
database, err := db.New(t.TempDir())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
hub := NewWSHub(database)
agentID := "stats-agent"
conn := connectTestAgent(t, hub, agentID)
statsPayload, _ := json.Marshal(map[string]interface{}{
"hashrate_15s": 100.0, "hashrate_1m": 90.0, "hashrate_15m": 80.0,
"shares_submitted": 5, "shares_accepted": 4,
"cpu_usage_pct": 12.5, "memory_usage_pct": 40.0, "uptime_seconds": 60,
})
if err := conn.WriteJSON(Message{Type: "stats", Payload: statsPayload}); err != nil {
t.Fatal(err)
}
logPayload, _ := json.Marshal(map[string]interface{}{"content": "line1\nline2", "lines": 2})
if err := conn.WriteJSON(Message{Type: "log_tail", Payload: logPayload}); err != nil {
t.Fatal(err)
}
time.Sleep(50 * time.Millisecond)
if got := hub.GetAgentLog(agentID); got != "line1\nline2" {
t.Fatalf("log tail = %q", got)
}
cmdPayload, _ := json.Marshal(map[string]interface{}{"action": "exec", "success": true})
if err := conn.WriteJSON(Message{Type: "command_result", Payload: cmdPayload}); err != nil {
t.Fatal(err)
}
}
func TestHandleAgentWSMaxAgentsPolicy(t *testing.T) {
database, err := db.New(t.TempDir())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
hub := NewWSHub(database)
hub.SetServerPolicy(ServerPolicy{MaxAgents: 1})
conn1, _ := dialAgentWS(t, hub)
authAgentConn(t, conn1, map[string]interface{}{"agent_id": "first", "hostname": "h1"})
conn2, _ := dialAgentWS(t, hub)
resp := authAgentConn(t, conn2, map[string]interface{}{"agent_id": "second", "hostname": "h2"})
var body map[string]interface{}
if err := json.Unmarshal(resp.Payload, &body); err != nil {
t.Fatal(err)
}
if body["success"] != false {
t.Fatalf("second agent should be rejected at max=1: %+v", body)
}
}
func TestHandleAgentWSInvalidAuthPayload(t *testing.T) {
database, err := db.New(t.TempDir())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
hub := NewWSHub(database)
conn, _ := dialAgentWS(t, hub)
if err := conn.WriteJSON(Message{Type: "auth", Payload: json.RawMessage(`"not-an-object"`)}); err != nil {
t.Fatal(err)
}
var resp Message
if err := conn.ReadJSON(&resp); err != nil {
t.Fatal(err)
}
var body map[string]interface{}
_ = json.Unmarshal(resp.Payload, &body)
if body["success"] != false {
t.Fatalf("invalid auth payload should fail: %+v", body)
}
}
func TestWSHubSendAgentCommandNotConnected(t *testing.T) {
hub := NewWSHub(nil)
if err := hub.SendAgentCommand("missing", "restart", nil); err == nil {
t.Fatal("expected error for disconnected agent")
}
}
func TestWSHubBroadcastHelpers(t *testing.T) {
hub := NewWSHub(nil)
hub.BroadcastServerLog(" ")
hub.BroadcastFleetAlert(map[string]string{"level": "info"})
hub.BroadcastPoolStatus(map[string]string{"connected": "true"})
hub.BroadcastAIActivity(map[string]string{"agent_id": "a"})
}
func TestWSHubEnrichAgentsCapabilities(t *testing.T) {
database, err := db.New(t.TempDir())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
hub := NewWSHub(database)
agentID := "cap-agent"
connectTestAgent(t, hub, agentID)
agents := []*models.Agent{{ID: agentID, Name: "x"}}
hub.enrichAgentsCapabilities(agents)
if agents[0].Capabilities == nil {
t.Fatal("expected capabilities enrichment")
}
}
func TestWSHubAgentPoolConfigDefaults(t *testing.T) {
hub := NewWSHub(nil)
hub.defaultPool = pool.Config{Host: "primary.pool", Port: 3333, Wallet: "48wallet", Password: "pw"}
hub.agentConfigs["a1"] = AgentForgeConfig{PoolHost: "custom.pool", PoolPort: 4444}
cfg := hub.agentPoolConfig("a1")
if cfg.Host != "custom.pool" || cfg.Port != 4444 {
t.Fatalf("unexpected pool cfg: %+v", cfg)
}
cfg = hub.agentPoolConfig("missing")
if cfg.Host != "primary.pool" {
t.Fatalf("missing agent should use default pool: %+v", cfg)
}
}
func TestWSHubConnectedAgentCount(t *testing.T) {
database, err := db.New(t.TempDir())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
hub := NewWSHub(database)
if hub.connectedAgentCount() != 0 {
t.Fatal("expected zero agents initially")
}
connectTestAgent(t, hub, "count-agent")
if hub.connectedAgentCount() != 1 {
t.Fatalf("expected 1 connected agent, got %d", hub.connectedAgentCount())
}
}