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