177 lines
7.6 KiB
Go
177 lines
7.6 KiB
Go
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)
|
|
}
|
|
}
|