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.
236 lines
9.5 KiB
236 lines
9.5 KiB
package controlplane
|
|
|
|
import (
|
|
"crypto/tls"
|
|
"net"
|
|
"os"
|
|
"path/filepath"
|
|
"strconv"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/gorilla/websocket"
|
|
)
|
|
|
|
func TestContinueTaskSessionStartsAThreadFromACompletedTask(t *testing.T) {
|
|
t.Setenv("MULTICLAW_CONTROLLER_DATA_DIR", t.TempDir())
|
|
controller, err := Open(filepath.Join(t.TempDir(), "controller-config.json"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
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 ('agent-1', '会话机', 'hash', 'offline', ?)", now); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
// 单独指派 sends a one-shot task: it must not carry any conversation.
|
|
oneShot, err := controller.CreateTaskWithImages("agent-1", "帮我下载一个文件到 D:\\demo", 300, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var sessionID string
|
|
if err := controller.db.QueryRow("SELECT COALESCE(session_id, '') FROM task_runs WHERE task_id = ?", oneShot).Scan(&sessionID); err != nil || sessionID != "" {
|
|
t.Fatalf("a one-shot task must have no conversation: %q / %v", sessionID, err)
|
|
}
|
|
if _, err := controller.ContinueTaskSession(oneShot, "把它复制到 D:\\target", 300, nil); err == nil {
|
|
t.Fatal("only a completed task can start a conversation")
|
|
}
|
|
|
|
controller.applyTaskEvent("agent-1", map[string]any{"task_id": oneShot, "status": "completed", "output": "已下载到 D:\\demo\\file.zip", "execution": "builtin"})
|
|
|
|
// The completed task becomes the first turn of a brand new conversation.
|
|
followUp, err := controller.ContinueTaskSession(oneShot, "把它复制到 D:\\target", 300, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if followUp.SessionID == "" || followUp.Turns != 2 || followUp.TaskID == oneShot {
|
|
t.Fatalf("expected a second turn in a new conversation, got %#v", followUp)
|
|
}
|
|
if err := controller.db.QueryRow("SELECT session_id FROM task_runs WHERE task_id = ?", oneShot).Scan(&sessionID); err != nil || sessionID != followUp.SessionID {
|
|
t.Fatalf("the original task must join its conversation: %q / %q / %v", sessionID, followUp.SessionID, err)
|
|
}
|
|
history, err := controller.sessionHistory(followUp.SessionID, 50)
|
|
if err != nil || len(history) != 2 {
|
|
t.Fatalf("expected the original turn as context, got %#v / %v", history, err)
|
|
}
|
|
if history[0]["content"] != "帮我下载一个文件到 D:\\demo" || history[1]["content"] != "已下载到 D:\\demo\\file.zip" {
|
|
t.Fatalf("unexpected seeded conversation: %#v", history)
|
|
}
|
|
|
|
// Continuing again appends to the same conversation.
|
|
controller.applyTaskEvent("agent-1", map[string]any{"task_id": followUp.TaskID, "status": "completed", "output": "已复制到 D:\\target", "execution": "builtin"})
|
|
third, err := controller.ContinueTaskSession(followUp.TaskID, "再确认一下文件数量", 300, nil)
|
|
if err != nil || third.SessionID != followUp.SessionID || third.Turns != 3 {
|
|
t.Fatalf("expected the third turn in the same conversation, got %#v / %v", third, err)
|
|
}
|
|
}
|
|
|
|
func TestFailedTurnDoesNotEnterHistoryAndOldTurnsAreTrimmed(t *testing.T) {
|
|
t.Setenv("MULTICLAW_CONTROLLER_DATA_DIR", t.TempDir())
|
|
controller, err := Open(filepath.Join(t.TempDir(), "controller-config.json"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
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 ('agent-1', '会话机', 'hash', 'offline', ?)", now); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
sessionID, err := controller.createTaskSession("agent-1")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
// Shrink the conversation so trimming is observable.
|
|
if _, err := controller.db.Exec("UPDATE task_sessions SET turn_limit = 2 WHERE session_id = ?", sessionID); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
failed, err := controller.createTaskWithImages("agent-1", "失败的一轮", 300, 1, "", nil, sessionID)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
controller.applyTaskEvent("agent-1", map[string]any{"task_id": failed, "status": "failed", "error": "下载失败"})
|
|
if history, err := controller.sessionHistory(sessionID, 2); err != nil || len(history) != 0 {
|
|
t.Fatalf("a failed turn must not enter the conversation: %#v / %v", history, err)
|
|
}
|
|
|
|
for _, instruction := range []string{"第一轮成功", "第二轮成功", "第三轮成功"} {
|
|
taskID, err := controller.createTaskWithImages("agent-1", instruction, 300, 1, "", nil, sessionID)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
controller.applyTaskEvent("agent-1", map[string]any{"task_id": taskID, "status": "completed", "output": "完成:" + instruction})
|
|
}
|
|
history, err := controller.sessionHistory(sessionID, 2)
|
|
if err != nil || len(history) != 4 {
|
|
t.Fatalf("expected only the newest two turns to remain: %#v / %v", history, err)
|
|
}
|
|
if history[0]["content"] != "第二轮成功" || history[3]["content"] != "完成:第三轮成功" {
|
|
t.Fatalf("unexpected trimmed conversation: %#v", history)
|
|
}
|
|
}
|
|
|
|
func TestFollowUpDispatchCarriesTheConversationHistory(t *testing.T) {
|
|
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)
|
|
}
|
|
defer controller.Close()
|
|
if err := controller.StartAgentServer(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
dialer := websocket.Dialer{TLSClientConfig: &tls.Config{InsecureSkipVerify: true}}
|
|
client, _, err := dialer.Dial("wss://127.0.0.1:"+strconv.Itoa(port), nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer client.Close()
|
|
code, err := controller.CreatePairingCode()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := client.WriteJSON(map[string]any{"type": "pair_request", "protocol": 1, "code": code, "agent_name": "会话机"}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var paired map[string]any
|
|
if err := client.ReadJSON(&paired); err != nil || paired["type"] != "pair_accepted" {
|
|
t.Fatalf("expected pairing: %#v / %v", paired, err)
|
|
}
|
|
agentID, _ := paired["agent_id"].(string)
|
|
|
|
// A one-shot task carries no conversation.
|
|
oneShot, err := controller.CreateTaskWithImages(agentID, "帮我下载一个文件到 D:\\demo", 300, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var firstDispatch map[string]any
|
|
if err := client.ReadJSON(&firstDispatch); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if firstDispatch["session_id"] != nil || firstDispatch["history"] != nil {
|
|
t.Fatalf("a one-shot task must not carry context: %#v", firstDispatch)
|
|
}
|
|
if err := client.WriteJSON(map[string]any{"type": "task_event", "task_id": oneShot, "status": "completed", "output": "已下载到 D:\\demo\\file.zip", "execution": "builtin"}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
waitForTaskStatus(t, controller, oneShot, "completed")
|
|
|
|
if _, err := controller.ContinueTaskSession(oneShot, "把它复制到 D:\\target", 300, nil); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var followUp map[string]any
|
|
if err := client.ReadJSON(&followUp); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if followUp["session_id"] == nil {
|
|
t.Fatalf("a follow-up must carry its conversation: %#v", followUp)
|
|
}
|
|
history, _ := followUp["history"].([]any)
|
|
if len(history) != 2 {
|
|
t.Fatalf("expected the original turn as context: %#v", followUp)
|
|
}
|
|
previous, _ := history[1].(map[string]any)
|
|
if previous["role"] != "assistant" || previous["content"] != "已下载到 D:\\demo\\file.zip" {
|
|
t.Fatalf("unexpected replayed history: %#v", history)
|
|
}
|
|
}
|
|
|
|
func TestTaskSessionTurnsExposeTheStoredContext(t *testing.T) {
|
|
t.Setenv("MULTICLAW_CONTROLLER_DATA_DIR", t.TempDir())
|
|
controller, err := Open(filepath.Join(t.TempDir(), "controller-config.json"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer controller.Close()
|
|
sessionID, err := controller.createTaskSession("agent-1")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
controller.appendSessionTurn(sessionID, "帮我在桌面建一个 mc-chat 文件夹", "已在桌面创建文件夹:mc-chat")
|
|
controller.appendSessionTurn(sessionID, "把它复制到 D:\\target", "已复制到 D:\\target")
|
|
|
|
turns, err := controller.TaskSessionTurns(sessionID)
|
|
if err != nil || len(turns) != 2 {
|
|
t.Fatalf("expected two stored turns, got %#v / %v", turns, err)
|
|
}
|
|
if turns[0].Instruction != "帮我在桌面建一个 mc-chat 文件夹" || turns[0].Answer != "已在桌面创建文件夹:mc-chat" {
|
|
t.Fatalf("unexpected first turn: %#v", turns[0])
|
|
}
|
|
if turns[1].Instruction != "把它复制到 D:\\target" || turns[1].Answer != "已复制到 D:\\target" {
|
|
t.Fatalf("unexpected second turn: %#v", turns[1])
|
|
}
|
|
// A task without a conversation (or an unknown id) must render as empty.
|
|
if turns, err := controller.TaskSessionTurns(""); err != nil || len(turns) != 0 {
|
|
t.Fatalf("expected an empty context: %#v / %v", turns, err)
|
|
}
|
|
if turns, err := controller.TaskSessionTurns("missing-session"); err != nil || len(turns) != 0 {
|
|
t.Fatalf("expected an empty context for an unknown session: %#v / %v", turns, err)
|
|
}
|
|
}
|
|
|
|
// waitForTaskStatus waits until the controller has processed an Agent event.
|
|
func waitForTaskStatus(t *testing.T, controller *ControlPlane, taskID, status string) {
|
|
t.Helper()
|
|
deadline := time.Now().Add(2 * time.Second)
|
|
for {
|
|
var current string
|
|
if controller.db.QueryRow("SELECT status FROM task_runs WHERE task_id = ?", taskID).Scan(¤t) == nil && current == status {
|
|
return
|
|
}
|
|
if time.Now().After(deadline) {
|
|
t.Fatalf("task did not reach %q in time", status)
|
|
}
|
|
time.Sleep(10 * time.Millisecond)
|
|
}
|
|
}
|
|
|