package wsclient import ( "context" "encoding/json" "net/http" "net/http/httptest" "strings" "testing" "time" "github.com/coder/websocket" "gitea.dcglab.co.uk/steve/restic-manager/internal/api" ) func TestConnectOnceCleanDisconnectDoesNotPanic(t *testing.T) { serverErr := make(chan error, 1) srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := websocket.Accept(w, r, nil) if err != nil { serverErr <- err return } defer conn.CloseNow() //nolint:errcheck // Wait for the agent hello so Dial and the first client write have both // completed before ending the connection normally. if _, _, err := conn.Read(r.Context()); err != nil { serverErr <- err return } serverErr <- conn.Close(websocket.StatusNormalClosure, "test complete") })) defer srv.Close() ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() err := connectOnce(ctx, Config{ ServerURL: srv.URL, AgentToken: "test-token", HeartbeatPeriod: time.Hour, }, nil) if err == nil { t.Fatal("connectOnce returned nil after server disconnected") } if err := <-serverErr; err != nil { t.Fatalf("server websocket: %v", err) } } func TestConnectOnceAcceptsMessageLargerThanDefaultReadLimit(t *testing.T) { received := make(chan struct{}, 1) serverErr := make(chan error, 1) srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := websocket.Accept(w, r, nil) if err != nil { serverErr <- err return } defer conn.CloseNow() //nolint:errcheck if _, _, err := conn.Read(r.Context()); err != nil { serverErr <- err return } env := api.Envelope{ Type: api.MsgConfigUpdate, Payload: json.RawMessage(`{"padding":"` + strings.Repeat("x", 40*1024) + `"}`), } raw, _ := json.Marshal(env) serverErr <- conn.Write(r.Context(), websocket.MessageText, raw) <-r.Context().Done() })) defer srv.Close() ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) go func() { done <- connectOnce(ctx, Config{ ServerURL: srv.URL, AgentToken: "test-token", HeartbeatPeriod: time.Hour, }, func(_ context.Context, env api.Envelope, _ Sender) error { if env.Type == api.MsgConfigUpdate { received <- struct{}{} cancel() } return nil }) }() select { case <-received: case <-time.After(5 * time.Second): t.Fatal("agent did not receive oversized server message") } if err := <-serverErr; err != nil { t.Fatalf("server websocket: %v", err) } if err := <-done; err == nil { t.Fatal("connectOnce returned nil") } }