Files
lvmh/daemon/hub_test.go
T

333 lines
9.6 KiB
Go

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