278 lines
6.7 KiB
Go
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)
|
|
}
|
|
}
|