You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
 

122 lines
4.8 KiB

package controlplane
import (
"crypto/tls"
"net"
"os"
"path/filepath"
"strconv"
"testing"
"time"
"github.com/gorilla/websocket"
)
func newRestartTestController(t *testing.T) (*ControlPlane, string) {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
port := listener.Addr().(*net.TCPAddr).Port
_ = listener.Close()
configPath := filepath.Join(t.TempDir(), "controller-config.json")
if err := os.WriteFile(configPath, []byte(`{"agent_host":"127.0.0.1","agent_port":`+strconv.Itoa(port)+`}`), 0o600); err != nil {
t.Fatal(err)
}
t.Setenv("MULTICLAW_CONTROLLER_DATA_DIR", t.TempDir())
controller, err := Open(configPath)
if err != nil {
t.Fatal(err)
}
if err := controller.StartAgentServer(); err != nil {
t.Fatal(err)
}
return controller, "wss://127.0.0.1:" + strconv.Itoa(port)
}
func pairRestartTestAgent(t *testing.T, controller *ControlPlane, url string) (*websocket.Conn, string, string) {
t.Helper()
client, _, err := (&websocket.Dialer{TLSClientConfig: &tls.Config{InsecureSkipVerify: true}}).Dial(url, nil) //nolint:gosec -- test-only TLS.
if err != nil {
t.Fatal(err)
}
code, err := controller.CreatePairingCode()
if err != nil {
t.Fatal(err)
}
if err := client.WriteJSON(map[string]any{"type": "pair_request", "code": code, "agent_name": "restart-agent"}); err != nil {
t.Fatal(err)
}
var paired map[string]any
if err := client.ReadJSON(&paired); err != nil {
t.Fatal(err)
}
return client, paired["agent_id"].(string), paired["credential"].(string)
}
func TestRestartOnlineAgentIsIdempotentAndReconnects(t *testing.T) {
controller, url := newRestartTestController(t)
defer controller.Close()
client, agentID, credential := pairRestartTestAgent(t, controller, url)
defer client.Close()
first, err := controller.RestartAgent(agentID)
if err != nil || first.Status != "restarting" {
t.Fatalf("restart request failed: %#v / %v", first, err)
}
var command map[string]any
if err := client.ReadJSON(&command); err != nil || command["type"] != "agent_restart" || command["request_id"] != first.RequestID {
t.Fatalf("expected restart command, got %#v / %v", command, err)
}
second, err := controller.RestartAgent(agentID)
if err != nil || second.RequestID != first.RequestID {
t.Fatalf("duplicate must reuse operation: %#v / %v", second, err)
}
if err := client.WriteJSON(map[string]any{"type": "agent_restart_ack", "request_id": first.RequestID, "status": "accepted"}); err != nil {
t.Fatal(err)
}
time.Sleep(25 * time.Millisecond)
_ = client.Close()
replacement, _, err := (&websocket.Dialer{TLSClientConfig: &tls.Config{InsecureSkipVerify: true}}).Dial(url, nil) //nolint:gosec -- test-only TLS.
if err != nil {
t.Fatal(err)
}
defer replacement.Close()
if err := replacement.WriteJSON(map[string]any{"type": "authenticate", "agent_id": agentID, "credential": credential, "agent_name": "restart-agent"}); err != nil {
t.Fatal(err)
}
var accepted map[string]any
if err := replacement.ReadJSON(&accepted); err != nil || accepted["type"] != "auth_accepted" {
t.Fatalf("replacement did not reconnect: %#v / %v", accepted, err)
}
var status string
if err := controller.db.QueryRow("SELECT status FROM agent_restart_operations WHERE agent_id=?", agentID).Scan(&status); err != nil || status != "reconnected" {
t.Fatalf("restart state = %q / %v", status, err)
}
}
func TestRestartOfflineAgentIsRejectedAndLateCompletionCannotWin(t *testing.T) {
controller, _ := newRestartTestController(t)
defer controller.Close()
now := time.Now().UTC().Format(time.RFC3339Nano)
if _, err := controller.db.Exec("INSERT INTO managed_agents(agent_id, name, credential_hash, status, last_seen_at) VALUES ('offline', 'offline', 'hash', 'offline', ?)", now); err != nil {
t.Fatal(err)
}
if _, err := controller.RestartAgent("offline"); err == nil {
t.Fatal("offline Agent restart must be rejected")
}
if _, err := controller.db.Exec("INSERT INTO task_runs(task_id, agent_id, agent_name, instruction, status, created_at, timeout_seconds) VALUES ('cancelled-late', 'offline', 'offline', 'test', 'cancelling', ?, 30)", now); err != nil {
t.Fatal(err)
}
controller.applyTaskEvent("offline", map[string]any{"task_id": "cancelled-late", "status": "completed", "output": "late"})
var taskStatus string
if err := controller.db.QueryRow("SELECT status FROM task_runs WHERE task_id='cancelled-late'").Scan(&taskStatus); err != nil || taskStatus != "cancelling" {
t.Fatalf("late completion overwrote cancellation: %q / %v", taskStatus, err)
}
controller.applyTaskEvent("offline", map[string]any{"task_id": "cancelled-late", "status": "cancelled", "message": "已取消"})
if err := controller.db.QueryRow("SELECT status FROM task_runs WHERE task_id='cancelled-late'").Scan(&taskStatus); err != nil || taskStatus != "cancelled" {
t.Fatalf("cancel acknowledgement missing: %q / %v", taskStatus, err)
}
}