daemon: golang WS hub, REST, gitlab, docker spawner, sqlite (18/18 tests)
This commit is contained in:
@@ -0,0 +1,332 @@
|
||||
package main
|
||||
|
||||
// hub_test.go — agent handshake, replay, prompt routing, live fan-out.
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
const testToken string = "test-token"
|
||||
|
||||
// newTestServer spins a Server over a temp store and returns both.
|
||||
func newTestServer(t *testing.T) (*httptest.Server, *Store) {
|
||||
t.Helper()
|
||||
store := openTestStore(t)
|
||||
daemonToken = testToken // tests run sequentially; package var set per-server
|
||||
hub := NewHub(store)
|
||||
srv := &Server{store: store, hub: hub, gitlab: NewGitLab(store, "https://gitlab.example")}
|
||||
ts := httptest.NewServer(srv.Routes(""))
|
||||
t.Cleanup(ts.Close)
|
||||
return ts, store
|
||||
}
|
||||
|
||||
func dialAgent(t *testing.T, ts *httptest.Server) *websocket.Conn {
|
||||
t.Helper()
|
||||
wsURL := "ws" + strings.TrimPrefix(ts.URL, "http") + "/agent/ws"
|
||||
hdr := http.Header{"Authorization": []string{"Bearer " + testToken}}
|
||||
ws, _, err := websocket.DefaultDialer.Dial(wsURL, hdr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial agent ws: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = ws.Close() })
|
||||
return ws
|
||||
}
|
||||
|
||||
func readFrame(t *testing.T, ws *websocket.Conn) map[string]any {
|
||||
t.Helper()
|
||||
var m map[string]any
|
||||
if err := ws.ReadJSON(&m); err != nil {
|
||||
t.Fatalf("read frame: %v", err)
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
func helloFrame(sessionID string) map[string]any {
|
||||
return map[string]any{
|
||||
"v": 1, "type": evHello, "sessionId": sessionID, "seq": 0, "ts": 1,
|
||||
"session": map[string]any{
|
||||
"id": sessionID, "name": nil, "cwd": "/w", "model": "glm-5.3",
|
||||
"provider": "zai-renaud", "agent": true, "repo": nil, "startedAt": 99,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestHubHandshakeWelcomeLastSeq(t *testing.T) {
|
||||
ts, store := newTestServer(t)
|
||||
ws := dialAgent(t, ts)
|
||||
if err := ws.WriteJSON(helloFrame("s1")); err != nil {
|
||||
t.Fatalf("send hello: %v", err)
|
||||
}
|
||||
if welcome := readFrame(t, ws); welcome["type"] != evWelcome || welcome["lastSeq"].(float64) != 0 {
|
||||
t.Fatalf("welcome = %v", welcome)
|
||||
}
|
||||
|
||||
for _, typ := range []string{evMessageStart, evMessageEnd, evAgentSettled} {
|
||||
if err := ws.WriteJSON(map[string]any{"v": 1, "type": typ, "sessionId": "s1", "seq": 1, "ts": 100}); err != nil {
|
||||
t.Fatalf("send %s: %v", typ, err)
|
||||
}
|
||||
}
|
||||
waitFor(t, 2*time.Second, func() bool {
|
||||
seq, _ := store.LastSeq("s1")
|
||||
return seq == 1
|
||||
})
|
||||
_ = ws.Close()
|
||||
|
||||
ws2 := dialAgent(t, ts)
|
||||
_ = ws2.WriteJSON(helloFrame("s1"))
|
||||
welcome2 := readFrame(t, ws2)
|
||||
if welcome2["lastSeq"].(float64) != 1 {
|
||||
t.Fatalf("welcome lastSeq after reconnect = %v, want 1", welcome2["lastSeq"])
|
||||
}
|
||||
}
|
||||
|
||||
func waitFor(t *testing.T, timeout time.Duration, cond func() bool) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
if cond() {
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
t.Fatal("condition not met before timeout")
|
||||
}
|
||||
|
||||
func TestHubReconnectReplacesOldConn(t *testing.T) {
|
||||
ts, _ := newTestServer(t)
|
||||
ws1 := dialAgent(t, ts)
|
||||
_ = ws1.WriteJSON(helloFrame("s1"))
|
||||
_ = readFrame(t, ws1)
|
||||
|
||||
ws2 := dialAgent(t, ts)
|
||||
_ = ws2.WriteJSON(helloFrame("s1"))
|
||||
_ = readFrame(t, ws2)
|
||||
|
||||
_ = ws1.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
if _, _, err := ws1.ReadMessage(); err == nil {
|
||||
t.Fatal("old agent conn should be closed after replacement")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHubPromptRoutingAndOffline409(t *testing.T) {
|
||||
ts, _ := newTestServer(t)
|
||||
|
||||
resp, err := http.Post(ts.URL+"/api/sessions/s1/prompt",
|
||||
"application/json", strings.NewReader(`{"message":"hi"}`))
|
||||
if err != nil {
|
||||
t.Fatalf("post prompt: %v", err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusUnauthorized {
|
||||
t.Fatalf("unauthenticated prompt = %d, want 401", resp.StatusCode)
|
||||
}
|
||||
|
||||
postJSON := func(path, body string) int {
|
||||
req, _ := http.NewRequest(http.MethodPost, ts.URL+path, strings.NewReader(body))
|
||||
req.Header.Set("Authorization", "Bearer "+testToken)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("post %s: %v", path, err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
return resp.StatusCode
|
||||
}
|
||||
|
||||
if code := postJSON("/api/sessions/s1/prompt", `{"message":"hi"}`); code != http.StatusConflict {
|
||||
t.Fatalf("offline prompt = %d, want 409", code)
|
||||
}
|
||||
|
||||
ws := dialAgent(t, ts)
|
||||
_ = ws.WriteJSON(helloFrame("s1"))
|
||||
_ = readFrame(t, ws)
|
||||
|
||||
if code := postJSON("/api/sessions/s1/prompt", `{"message":"do it"}`); code != http.StatusOK {
|
||||
t.Fatalf("online prompt = %d, want 200", code)
|
||||
}
|
||||
got := readFrame(t, ws)
|
||||
if got["type"] != evPrompt || got["message"] != "do it" {
|
||||
t.Fatalf("prompt frame = %v", got)
|
||||
}
|
||||
if _, ok := got["promptId"].(string); !ok {
|
||||
t.Fatalf("prompt frame missing promptId: %v", got)
|
||||
}
|
||||
|
||||
if code := postJSON("/api/sessions/s1/abort", ""); code != http.StatusOK {
|
||||
t.Fatalf("abort = %d, want 200", code)
|
||||
}
|
||||
abort := readFrame(t, ws)
|
||||
if abort["type"] != evAbort {
|
||||
t.Fatalf("abort frame = %v", abort)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHubMessageUpdateNotPersistedOthersAre(t *testing.T) {
|
||||
ts, store := newTestServer(t)
|
||||
ws := dialAgent(t, ts)
|
||||
_ = ws.WriteJSON(helloFrame("s1"))
|
||||
_ = readFrame(t, ws)
|
||||
|
||||
frames := []map[string]any{
|
||||
{"v": 1, "type": evMessageUpdate, "sessionId": "s1", "seq": 1, "ts": 10, "delta": "chunk"},
|
||||
{"v": 1, "type": evMessageEnd, "sessionId": "s1", "seq": 2, "ts": 11, "message": map[string]any{"role": "assistant", "id": "m2", "text": "chunk"}},
|
||||
}
|
||||
for _, f := range frames {
|
||||
if err := ws.WriteJSON(f); err != nil {
|
||||
t.Fatalf("send: %v", err)
|
||||
}
|
||||
}
|
||||
waitFor(t, 2*time.Second, func() bool {
|
||||
events, _ := store.EventsAfter("s1", 0, 100)
|
||||
return len(events) == 1
|
||||
})
|
||||
events, _ := store.EventsAfter("s1", 0, 100)
|
||||
if len(events) != 1 || events[0].Type != evMessageEnd || events[0].Seq != 2 {
|
||||
t.Fatalf("persisted events = %+v, want only message_end seq 2", events)
|
||||
}
|
||||
if last, _ := store.LastSeq("s1"); last != 2 {
|
||||
t.Fatalf("lastSeq = %d, want 2 (deltas unpersisted but seq advanced)", last)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHubMalformedFramesDoNotKillConn(t *testing.T) {
|
||||
ts, _ := newTestServer(t)
|
||||
ws := dialAgent(t, ts)
|
||||
if err := ws.WriteMessage(websocket.TextMessage, []byte("not json")); err != nil {
|
||||
t.Fatalf("send garbage: %v", err)
|
||||
}
|
||||
if err := ws.WriteJSON(map[string]any{"v": 1, "type": 42, "sessionId": "s1"}); err != nil {
|
||||
t.Fatalf("send wrong-type: %v", err)
|
||||
}
|
||||
_ = ws.WriteJSON(helloFrame("s1"))
|
||||
if welcome := readFrame(t, ws); welcome["type"] != evWelcome {
|
||||
t.Fatalf("conn died after malformed frames: %v", welcome)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHubWebWSSubscribeReceivesLiveEvents(t *testing.T) {
|
||||
ts, _ := newTestServer(t)
|
||||
wsURL := "ws" + strings.TrimPrefix(ts.URL, "http") + "/ws?token=" + testToken
|
||||
web, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("dial web ws: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = web.Close() })
|
||||
|
||||
if first := readFrame(t, web); first["type"] != frameSessionList {
|
||||
t.Fatalf("first web frame = %v", first)
|
||||
}
|
||||
|
||||
agent := dialAgent(t, ts)
|
||||
_ = agent.WriteJSON(helloFrame("s1"))
|
||||
_ = readFrame(t, agent)
|
||||
|
||||
// session_list broadcast on connect arrives before the event stream.
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
_ = web.SetReadDeadline(deadline)
|
||||
for {
|
||||
if m := readFrame(t, web); m["type"] == frameSessionList {
|
||||
break
|
||||
}
|
||||
}
|
||||
_ = web.WriteJSON(map[string]any{"type": frameSubscribe, "sessionId": "s1"})
|
||||
|
||||
// The subscription registers server-side asynchronously; an event sent too
|
||||
// early would be dropped (per protocol the UI refetches from REST). Resend
|
||||
// until a batch arrives. A reader goroutine avoids deadline-based reads
|
||||
// (gorilla conns fail permanently after a read timeout).
|
||||
frames := make(chan map[string]any, 16)
|
||||
go func() {
|
||||
defer close(frames)
|
||||
for {
|
||||
var m map[string]any
|
||||
if err := web.ReadJSON(&m); err != nil {
|
||||
return
|
||||
}
|
||||
frames <- m
|
||||
}
|
||||
}()
|
||||
start := time.Now()
|
||||
seq := int64(5)
|
||||
var got map[string]any
|
||||
for got == nil {
|
||||
seq++
|
||||
_ = agent.WriteJSON(map[string]any{"v": 1, "type": evMessageUpdate, "sessionId": "s1", "seq": seq, "ts": 50, "delta": "hi"})
|
||||
select {
|
||||
case m, ok := <-frames:
|
||||
if !ok {
|
||||
t.Fatal("web conn closed before events frame")
|
||||
}
|
||||
if m["type"] == frameEvents {
|
||||
got = m
|
||||
}
|
||||
case <-time.After(200 * time.Millisecond):
|
||||
}
|
||||
if time.Since(start) > 3*time.Second {
|
||||
t.Fatal("no events frame delivered after subscribe")
|
||||
}
|
||||
}
|
||||
evList := got["events"].([]any)
|
||||
if len(evList) == 0 {
|
||||
t.Fatalf("events batch empty: %v", got)
|
||||
}
|
||||
first := evList[0].(map[string]any)
|
||||
if got["sessionId"] != "s1" || got["after"].(float64) != first["seq"].(float64)-1 {
|
||||
t.Fatalf("events frame = %v", got)
|
||||
}
|
||||
for _, e := range evList {
|
||||
if e.(map[string]any)["type"] != evMessageUpdate {
|
||||
t.Fatalf("unexpected event in batch: %v", e)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestHubWebWSBadTokenRejected(t *testing.T) {
|
||||
ts, _ := newTestServer(t)
|
||||
wsURL := "ws" + strings.TrimPrefix(ts.URL, "http") + "/ws?token=wrong"
|
||||
if web, _, err := websocket.DefaultDialer.Dial(wsURL, nil); err == nil {
|
||||
_ = web.Close()
|
||||
t.Fatal("web ws with bad token should be rejected before upgrade")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHubSessionsViewOnlineFlag(t *testing.T) {
|
||||
ts, _ := newTestServer(t)
|
||||
ws := dialAgent(t, ts)
|
||||
_ = ws.WriteJSON(helloFrame("s1"))
|
||||
_ = readFrame(t, ws)
|
||||
|
||||
fetchSessions := func() []map[string]any {
|
||||
req, _ := http.NewRequest(http.MethodGet, ts.URL+"/api/sessions", nil)
|
||||
req.Header.Set("Authorization", "Bearer "+testToken)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("get sessions: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
var sessions []map[string]any
|
||||
if err := json.NewDecoder(resp.Body).Decode(&sessions); err != nil {
|
||||
t.Fatalf("decode: %v", err)
|
||||
}
|
||||
return sessions
|
||||
}
|
||||
|
||||
sessions := fetchSessions()
|
||||
if len(sessions) != 1 || sessions[0]["id"] != "s1" {
|
||||
t.Fatalf("sessions = %v", sessions)
|
||||
}
|
||||
if sessions[0]["online"] != true {
|
||||
t.Fatalf("online flag = %v, want true", sessions[0]["online"])
|
||||
}
|
||||
|
||||
_ = ws.Close()
|
||||
waitFor(t, 2*time.Second, func() bool {
|
||||
return fetchSessions()[0]["online"] == false
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user