package centrifugo import ( "context" "encoding/base64" "errors" "fmt" "net/http" "net/http/httptest" "strings" "sync" "testing" "time" "git.golem15.com/golem15/summercms/modules/backpack" "git.golem15.com/golem15/summercms/modules/bouncer" "git.golem15.com/golem15/summercms/modules/lighthouse" "github.com/golang-jwt/jwt/v5" ) const tokenTestSecret = "test-only-centrifugo-token-secret-0123456789" var fixedNow = time.Date(2026, 9, 30, 12, 0, 0, 0, time.UTC) // jwtParts returns the decoded header and claims segments of token after // checking its HS256 signature with secret. func jwtParts(t *testing.T, token, secret string) (string, string) { t.Helper() parsed, err := jwt.Parse(token, func(tok *jwt.Token) (any, error) { return []byte(secret), nil }, jwt.WithValidMethods([]string{"HS256"}), jwt.WithoutClaimsValidation()) if err != nil || !parsed.Valid { t.Fatalf("token does not verify: %v", err) } seg := strings.Split(token, ".") if len(seg) != 3 { t.Fatalf("token has %d segments", len(seg)) } dec := func(s string) string { b, err := base64.RawURLEncoding.DecodeString(s) if err != nil { t.Fatal(err) } return string(b) } return dec(seg[0]), dec(seg[1]) } // TestTokenClaims covers RT-01 and T-11-11: every generator signs HS256 with // the exact WinterCMS claim set and order, the user token carries only the // name, and an empty secret signs nothing. func TestTokenClaims(t *testing.T) { iss := NewTokenIssuer(tokenTestSecret, time.Hour) iss.Now = func() time.Time { return fixedNow } exp := fixedNow.Add(time.Hour).Unix() name := "Ann" cases := []struct { name string sign func() (string, error) want string }{ {"for_user", func() (string, error) { return iss.ForUser(lighthouse.User{ID: 7, Name: &name}) }, fmt.Sprintf(`{"sub":"7","exp":%d,"info":{"name":"Ann"}}`, exp)}, {"for_user_null_name", func() (string, error) { return iss.ForUser(lighthouse.User{ID: 8}) }, fmt.Sprintf(`{"sub":"8","exp":%d,"info":{"name":null}}`, exp)}, {"subscription", func() (string, error) { return iss.Subscription(lighthouse.User{ID: 7}, "collection:5") }, fmt.Sprintf(`{"sub":"7","channel":"collection:5","exp":%d}`, exp)}, {"anonymous", iss.Anonymous, fmt.Sprintf(`{"sub":"","exp":%d}`, fixedNow.Add(5*time.Minute).Unix())}, {"for_identifier_empty_info", func() (string, error) { return iss.ForIdentifier("kiosk-1", nil) }, fmt.Sprintf(`{"sub":"kiosk-1","exp":%d,"info":[]}`, exp)}, {"for_identifier_info", func() (string, error) { return iss.ForIdentifier("kiosk-1", map[string]any{"room": "a/b"}) }, fmt.Sprintf(`{"sub":"kiosk-1","exp":%d,"info":{"room":"a/b"}}`, exp)}, {"subscription_for_identifier", func() (string, error) { return iss.SubscriptionForIdentifier("kiosk-1", "room:1") }, fmt.Sprintf(`{"sub":"kiosk-1","channel":"room:1","exp":%d}`, exp)}, } for _, c := range cases { t.Run(c.name, func(t *testing.T) { tok, err := c.sign() if err != nil { t.Fatal(err) } header, claims := jwtParts(t, tok, tokenTestSecret) if header != `{"alg":"HS256","typ":"JWT"}` { t.Fatalf("header = %s", header) } if claims != c.want { t.Fatalf("claims = %s, want %s", claims, c.want) } }) } empty := NewTokenIssuer("", 0) if empty.Configured() || empty.ttl != DefaultTokenTTL { t.Fatalf("empty issuer configured=%v ttl=%s", empty.Configured(), empty.ttl) } for name, sign := range map[string]func() (string, error){ "ForUser": func() (string, error) { return empty.ForUser(lighthouse.User{ID: 1}) }, "Subscription": func() (string, error) { return empty.Subscription(lighthouse.User{ID: 1}, "a:1") }, "Anonymous": empty.Anonymous, "ForIdentifier": func() (string, error) { return empty.ForIdentifier("x", nil) }, "SubscriptionForIdentifier": func() (string, error) { return empty.SubscriptionForIdentifier("x", "a:1") }, } { if tok, err := sign(); !errors.Is(err, ErrNotConfigured) || tok != "" { t.Errorf("%s with an empty secret = %q, %v; want ErrNotConfigured", name, tok, err) } } var nilIss *TokenIssuer if nilIss.Configured() { t.Fatal("nil issuer is configured") } if _, err := iss.ForIdentifier("x", map[string]any{"bad": make(chan int)}); err == nil { t.Fatal("unencodable info accepted") } // The default clock is used when Now is nil. live := NewTokenIssuer(tokenTestSecret, time.Minute) tok, err := live.Anonymous() if err != nil { t.Fatal(err) } if _, claims := jwtParts(t, tok, tokenTestSecret); !strings.Contains(claims, `"exp":`) { t.Fatalf("claims = %s", claims) } } func tokenService(t *testing.T, lookup lighthouse.UserLookup) *lighthouse.Service { t.Helper() svc, err := lighthouse.From(backpack.New(nil)) if err != nil { t.Fatal(err) } svc.SetUserLookup(lookup) return svc } // TestTokenHandler covers RT-01: 401 without a principal or user, 503 with // an empty secret only after the user check, and a 200 {"token"} body with // the Laravel JSON headers and no trailing newline, safe under concurrency. func TestTokenHandler(t *testing.T) { name := "Ann" svc := tokenService(t, func(_ context.Context, id uint) (lighthouse.User, bool, error) { switch id { case 7: return lighthouse.User{ID: 7, Name: &name}, true, nil case 9: return lighthouse.User{}, false, errors.New("database down") } return lighthouse.User{}, false, nil }) iss := NewTokenIssuer(tokenTestSecret, time.Hour) h := TokenHandler(svc, iss) noSecret := TokenHandler(svc, NewTokenIssuer("", time.Hour)) call := func(h http.HandlerFunc, p *bouncer.Principal) *httptest.ResponseRecorder { req := httptest.NewRequest(http.MethodGet, "/api/realtime/token", nil) if p != nil { req = req.WithContext(bouncer.WithUser(req.Context(), p)) } rec := httptest.NewRecorder() h(rec, req) return rec } cases := []struct { name string h http.HandlerFunc p *bouncer.Principal status int body string }{ {"no_principal", h, nil, 401, `{"error":"Unauthorized"}`}, {"zero_id", h, &bouncer.Principal{ID: 0}, 401, `{"error":"Unauthorized"}`}, {"unknown_user", h, &bouncer.Principal{ID: 5}, 401, `{"error":"Unauthorized"}`}, {"lookup_error", h, &bouncer.Principal{ID: 9}, 401, `{"error":"Unauthorized"}`}, {"unknown_user_before_secret", noSecret, &bouncer.Principal{ID: 5}, 401, `{"error":"Unauthorized"}`}, {"empty_secret", noSecret, &bouncer.Principal{ID: 7}, 503, `{"error":"WebSocket not configured"}`}, } for _, c := range cases { t.Run(c.name, func(t *testing.T) { rec := call(c.h, c.p) if rec.Code != c.status || rec.Body.String() != c.body { t.Fatalf("got %d %s, want %d %s", rec.Code, rec.Body.String(), c.status, c.body) } if rec.Header().Get("Content-Type") != "application/json" || rec.Header().Get("Cache-Control") != "no-cache, private" { t.Fatalf("headers = %v", rec.Header()) } }) } t.Run("ok", func(t *testing.T) { rec := call(h, &bouncer.Principal{ID: 7}) body := rec.Body.String() if rec.Code != 200 || !strings.HasPrefix(body, `{"token":"`) || !strings.HasSuffix(body, `"}`) || strings.HasSuffix(body, "\n") { t.Fatalf("got %d %q", rec.Code, body) } tok := strings.TrimSuffix(strings.TrimPrefix(body, `{"token":"`), `"}`) if _, claims := jwtParts(t, tok, tokenTestSecret); !strings.HasPrefix(claims, `{"sub":"7","exp":`) || !strings.HasSuffix(claims, `,"info":{"name":"Ann"}}`) { t.Fatalf("claims = %s", claims) } }) t.Run("concurrent", func(t *testing.T) { var wg sync.WaitGroup for range 16 { wg.Add(1) go func() { defer wg.Done() if rec := call(h, &bouncer.Principal{ID: 7}); rec.Code != 200 { t.Errorf("concurrent status %d", rec.Code) } }() } wg.Wait() }) }