package controlplane import ( "encoding/json" "net/http" "net/http/httptest" "path/filepath" "testing" "time" ) func TestClientsRemainIsolatedWhenOneDisconnects(t *testing.T) { server, err := NewServer(ServerConfig{ DataFile: filepath.Join(t.TempDir(), "control-plane.json"), NodeTokens: map[string]string{"client-a": "secret-a", "client-b": "secret-b"}, WebUsers: map[string]string{"admin": "web-secret"}, HeartbeatTimeout: time.Minute, LeaseTTL: time.Minute, }) if err != nil { t.Fatal(err) } defer server.Close() httpServer := httptest.NewServer(server.Handler()) defer httpServer.Close() client := httpServer.Client() login := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/auth/login", "", map[string]string{"username": "admin", "password": "web-secret"}) var session struct { AccessToken string `json:"access_token"` } decodeBody(t, login, &session) webAuth := "Bearer " + session.AccessToken register := func(nodeID, token, connectionID string) { response := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/nodes/register", "Bearer "+token, NodeRegistration{ NodeID: nodeID, ConnectionID: connectionID, AgentVersion: "test", ProtocolVersion: ProtocolVersion, Capabilities: []string{"heartbeat", "poll-tasks", "send-text"}, Accounts: []AccountSummary{{AccountID: "account-a", Active: true, Verified: true}}, }) if response.Code != http.StatusOK { t.Fatalf("register %s status = %d: %s", nodeID, response.Code, response.Body.String()) } } register("client-a", "secret-a", "connection-a-1") register("client-b", "secret-b", "connection-b-1") create := func(nodeID, key string) TaskSubmissionResponse { response := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/tasks", webAuth, TaskSubmission{ NodeID: nodeID, AccountID: "account-a", Kind: "send-text", IdempotencyKey: key, Payload: json.RawMessage(`{"target_id":"target","text":"text","confirmed":true}`), }) if response.Code != http.StatusAccepted { t.Fatalf("create %s status = %d: %s", nodeID, response.Code, response.Body.String()) } var result TaskSubmissionResponse decodeBody(t, response, &result) return result } taskA := create("client-a", "a-1") taskB := create("client-b", "b-1") stale := time.Now().UTC().Add(-time.Hour) if err := server.store.Mutate(func(state *PersistedState) error { node := state.Nodes["client-a"] node.LastHeartbeatAt = &stale state.Nodes["client-a"] = node return nil }); err != nil { t.Fatal(err) } nodesResponse := doJSON(t, client, http.MethodGet, httpServer.URL+"/v1/nodes", webAuth, nil) if nodesResponse.Code != http.StatusOK { t.Fatalf("node list status = %d: %s", nodesResponse.Code, nodesResponse.Body.String()) } var nodesBody struct { Nodes []Node `json:"nodes"` } decodeBody(t, nodesResponse, &nodesBody) statuses := map[string]NodeStatus{} for _, node := range nodesBody.Nodes { statuses[node.NodeID] = node.Status } if statuses["client-a"] != NodeOffline || statuses["client-b"] != NodeOnline { t.Fatalf("unexpected client statuses: %+v", statuses) } getTask := func(taskID string) Task { response := doJSON(t, client, http.MethodGet, httpServer.URL+"/v1/tasks/"+taskID, webAuth, nil) if response.Code != http.StatusOK { t.Fatalf("get task %s status = %d: %s", taskID, response.Code, response.Body.String()) } var task Task decodeBody(t, response, &task) return task } if task := getTask(taskA.TaskID); task.Status != TaskWaitingForClient { t.Fatalf("client A task status = %s, want %s", task.Status, TaskWaitingForClient) } if task := getTask(taskB.TaskID); task.Status != TaskPending { t.Fatalf("client B task status = %s, want %s", task.Status, TaskPending) } bPoll := doJSON(t, client, http.MethodGet, httpServer.URL+"/v1/nodes/client-b/tasks?account_id=account-a", "Bearer secret-b", nil) var bBatch TaskBatch decodeBody(t, bPoll, &bBatch) if len(bBatch.Tasks) != 1 || bBatch.Tasks[0].TaskID != taskB.TaskID { t.Fatalf("client B received the wrong tasks: %+v", bBatch.Tasks) } bLease := bBatch.Tasks[0] bAck := TaskAck{TaskID: bLease.TaskID, AccountID: bLease.AccountID, LeaseGeneration: bLease.LeaseGeneration} for _, phase := range []string{"ack", "start"} { response := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/nodes/client-b/tasks/"+bLease.TaskID+"/"+phase, "Bearer secret-b", bAck) if response.Code != http.StatusOK { t.Fatalf("client B %s status = %d: %s", phase, response.Code, response.Body.String()) } } result := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/nodes/client-b/tasks/"+bLease.TaskID+"/result", "Bearer secret-b", TaskResult{ TaskID: bLease.TaskID, AccountID: bLease.AccountID, LeaseGeneration: bLease.LeaseGeneration, Status: TaskSucceeded, HasSideEffect: true, CorrelationID: "client-b-result", }) if result.Code != http.StatusOK { t.Fatalf("client B result status = %d: %s", result.Code, result.Body.String()) } bEvent := MessageEvent{NodeID: "client-b", AccountID: "account-a", ChatID: "client-b-chat", ChatType: ChatPrivate, EventSeq: 1, EventType: "message", OccurredAt: time.Now().UTC(), Content: "client-b-content", ConfigVersion: 1, AuthorizationVersion: 1, Authorized: true} bEventResponse := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/nodes/client-b/events", "Bearer secret-b", bEvent) if bEventResponse.Code != http.StatusAccepted { t.Fatalf("client B event status = %d: %s", bEventResponse.Code, bEventResponse.Body.String()) } bEvents := doJSON(t, client, http.MethodGet, httpServer.URL+"/v1/events?node_id=client-b", webAuth, nil) var bEventList struct { Events []StoredEvent `json:"events"` } decodeBody(t, bEvents, &bEventList) if len(bEventList.Events) != 1 || bEventList.Events[0].NodeID != "client-b" { t.Fatalf("client B event isolation failed: %+v", bEventList.Events) } aPoll := doJSON(t, client, http.MethodGet, httpServer.URL+"/v1/nodes/client-a/tasks?account_id=account-a", "Bearer secret-a", nil) var aBatch TaskBatch decodeBody(t, aPoll, &aBatch) if len(aBatch.Tasks) != 0 { t.Fatalf("offline client A received tasks: %+v", aBatch.Tasks) } register("client-a", "secret-a", "connection-a-2") staleHeartbeat := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/nodes/client-a/heartbeat", "Bearer secret-a", Heartbeat{ NodeID: "client-a", ConnectionID: "connection-a-1", AgentVersion: "test", ProtocolVersion: ProtocolVersion, NodeStatus: NodeOnline, CorrelationID: "stale-heartbeat", }) if staleHeartbeat.Code != http.StatusConflict { t.Fatalf("stale heartbeat status = %d: %s", staleHeartbeat.Code, staleHeartbeat.Body.String()) } freshHeartbeat := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/nodes/client-a/heartbeat", "Bearer secret-a", Heartbeat{ NodeID: "client-a", ConnectionID: "connection-a-2", AgentVersion: "test", ProtocolVersion: ProtocolVersion, NodeStatus: NodeOnline, CorrelationID: "fresh-heartbeat", }) if freshHeartbeat.Code != http.StatusOK { t.Fatalf("fresh heartbeat status = %d: %s", freshHeartbeat.Code, freshHeartbeat.Body.String()) } resume := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/tasks/"+taskA.TaskID+"/resume", webAuth, nil) var resumed Task decodeBody(t, resume, &resumed) if resume.Code != http.StatusOK || resumed.Status != TaskPending { t.Fatalf("resume status = %d task=%+v", resume.Code, resumed) } aPoll = doJSON(t, client, http.MethodGet, httpServer.URL+"/v1/nodes/client-a/tasks?account_id=account-a", "Bearer secret-a", nil) aBatch = TaskBatch{} decodeBody(t, aPoll, &aBatch) if len(aBatch.Tasks) != 1 || aBatch.Tasks[0].TaskID != taskA.TaskID { t.Fatalf("resumed client A received the wrong tasks: %+v", aBatch.Tasks) } }