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