Files
LINUX-AETHERFORGE/cmd/agent/ws_test.go
drjones 3678b199d0
Some checks failed
Test / test (push) Has been cancelled
Initial commit: AetherForge Linux (forge-mesh) v0.1.0-dev
2026-07-04 09:31:23 +00:00

278 lines
6.7 KiB
Go

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)
}
}