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