package main import ( "context" "encoding/json" "net/http" "net/http/httptest" "strings" "sync" "testing" "time" "forge-mesh/internal/api/types" "forge-mesh/internal/mining" "forge-mesh/internal/policy" "github.com/gorilla/websocket" ) func TestAgentWSHeartbeatAndCommands(t *testing.T) { const secret = "ws-test-secret" var mu sync.Mutex var heartbeats int var commands []string connected := make(chan struct{}, 1) upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }} srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/api/v1/ws/fleet" { http.NotFound(w, r) return } token := r.URL.Query().Get("token") auth := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") if token != secret && auth != secret { http.Error(w, "unauthorized", http.StatusUnauthorized) return } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } defer conn.Close() connected <- struct{}{} go func() { for { _ = conn.SetReadDeadline(time.Now().Add(10 * time.Second)) _, data, err := conn.ReadMessage() if err != nil { return } var msg map[string]any if json.Unmarshal(data, &msg) == nil && msg["type"] == "heartbeat" { mu.Lock() heartbeats++ mu.Unlock() } } }() for _, action := range []string{"pause", "screenshot", "reboot"} { cmd := types.WsMessage{ Type: "command", Command: &types.FleetCommand{ ID: "test-" + action, Action: action, }, } payload, _ := json.Marshal(cmd) if err := conn.WriteMessage(websocket.TextMessage, payload); err != nil { t.Errorf("write command %s: %v", action, err) return } mu.Lock() commands = append(commands, action) mu.Unlock() time.Sleep(30 * time.Millisecond) } time.Sleep(300 * time.Millisecond) })) defer srv.Close() chain := mining.NewChain(policy.DefaultMiningProfile("ws-wallet"), mining.DefaultChainConfig()) ctx, cancel := context.WithCancel(context.Background()) defer cancel() go func() { _ = chain.Run(ctx) }() agent := &Agent{ serverURL: srv.URL, secret: secret, hostID: "host-ws-test", hostname: "test-host", chain: chain, } agentCtx, agentCancel := context.WithTimeout(context.Background(), 10*time.Second) defer agentCancel() errCh := make(chan error, 1) go func() { errCh <- agent.connect(agentCtx) }() select { case <-connected: case err := <-errCh: t.Fatalf("connect failed: %v", err) case <-time.After(5 * time.Second): t.Fatal("mock server did not receive agent connection") } waitForAgentConn(t, agent, 2*time.Second) if err := agent.sendHeartbeat(); err != nil { t.Fatalf("send heartbeat: %v", err) } select { case err := <-errCh: if err != nil && agentCtx.Err() == nil { t.Fatalf("connect ended: %v", err) } case <-time.After(4 * time.Second): agentCancel() <-errCh } mu.Lock() defer mu.Unlock() if heartbeats < 1 { t.Fatalf("expected at least one heartbeat, got %d", heartbeats) } if len(commands) < 3 { t.Fatalf("expected pause/screenshot/reboot commands, got %v", commands) } } func TestAgentBeaconFallbackCommands(t *testing.T) { const secret = "beacon-test-secret" var received []string srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/api/v1/fleet/beacon" { http.NotFound(w, r) return } if strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") != secret { http.Error(w, "unauthorized", http.StatusUnauthorized) return } received = append(received, r.Method) _ = json.NewEncoder(w).Encode(types.BeaconResponse{ OK: true, Commands: []types.FleetCommand{ {ID: "b1", Action: "pause"}, {ID: "b2", Action: "screenshot"}, }, }) })) defer srv.Close() chain := mining.NewChain(policy.DefaultMiningProfile("beacon-wallet"), mining.DefaultChainConfig()) agent := &Agent{ serverURL: srv.URL, secret: secret, chain: chain, } agent.beaconFallback() if len(received) != 1 || received[0] != http.MethodPost { t.Fatalf("expected one beacon POST, got %v", received) } } func TestAgentSendHeartbeatPayload(t *testing.T) { upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }} got := make(chan []byte, 1) srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } defer conn.Close() _, data, err := conn.ReadMessage() if err == nil { got <- data } })) defer srv.Close() wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) if err != nil { t.Fatal(err) } defer conn.Close() chain := mining.NewChain(policy.DefaultMiningProfile("hb-wallet"), mining.DefaultChainConfig()) agent := &Agent{chain: chain, conn: conn} if err := agent.sendHeartbeat(); err != nil { t.Fatal(err) } select { case payload := <-got: var msg map[string]any if err := json.Unmarshal(payload, &msg); err != nil { t.Fatal(err) } if msg["type"] != "heartbeat" { t.Fatalf("type: got %v", msg["type"]) } if msg["hostname"] == nil { t.Fatal("expected hostname in heartbeat") } case <-time.After(2 * time.Second): t.Fatal("heartbeat not received by mock server") } } func waitForAgentConn(t *testing.T, agent *Agent, timeout time.Duration) { t.Helper() deadline := time.Now().Add(timeout) for time.Now().Before(deadline) { agent.mu.Lock() ok := agent.conn != nil agent.mu.Unlock() if ok { return } time.Sleep(20 * time.Millisecond) } t.Fatal("agent websocket not connected") } func TestHandleCommandPauseStopsChain(t *testing.T) { chain := mining.NewChain(policy.DefaultMiningProfile("pause-wallet"), mining.DefaultChainConfig()) agent := &Agent{chain: chain} agent.handleCommand(types.FleetCommand{ID: "1", Action: "pause"}) chain.Stop() } func TestHandleCommandMiningProfile(t *testing.T) { chain := mining.NewChain(policy.DefaultMiningProfile("old-wallet"), mining.DefaultChainConfig()) agent := &Agent{chain: chain} profile := types.MiningProfile{ WalletAddress: "new-wallet", Tiers: []types.MiningTierSpec{ {Type: "stratum", Duration: 1}, }, } raw, _ := json.Marshal(profile) var profileArg map[string]any _ = json.Unmarshal(raw, &profileArg) agent.handleCommand(types.FleetCommand{ ID: "profile-1", Action: "mining_profile", Args: map[string]any{"profile": profileArg}, }) updated := chain.Profile() if updated.WalletAddress != "new-wallet" { t.Fatalf("wallet: got %q want new-wallet", updated.WalletAddress) } if len(updated.Tiers) != 1 || updated.Tiers[0].Type != "stratum" { t.Fatalf("tiers not updated: %+v", updated.Tiers) } }