package discourse import ( "bytes" "crypto/hmac" "crypto/sha256" "encoding/base64" "encoding/hex" "encoding/json" "errors" "net/http" "net/http/httptest" "net/url" "strings" "testing" ) const testWebhookSecret = "whsec-test-0123456789" const testSSOSecret = "ssosec-test-0123456789" func hmacHex(secret string, data []byte) string { m := hmac.New(sha256.New, []byte(secret)) m.Write(data) return hex.EncodeToString(m.Sum(nil)) } // --- webhook payload verification --- func TestVerifyWebhookGoodSignature(t *testing.T) { body := []byte(`{"post":{"id":55,"topic_id":12}}`) sig := hmacHex(testWebhookSecret, body) if err := VerifyWebhook([]byte(testWebhookSecret), body, sig); err != nil { t.Fatalf("good sig rejected: %v", err) } if err := VerifyWebhook([]byte(testWebhookSecret), body, "sha256="+sig); err != nil { t.Fatalf("sha256=-prefixed sig rejected: %v", err) } } func TestVerifyWebhookBadSignature(t *testing.T) { body := []byte(`{"post":{"id":55}}`) err := VerifyWebhook([]byte(testWebhookSecret), body, "deadbeef") if !errors.Is(err, ErrBadSignature) { t.Fatalf("want ErrBadSignature, got %v", err) } err = VerifyWebhook([]byte("other-secret"), body, hmacHex(testWebhookSecret, body)) if !errors.Is(err, ErrBadSignature) { t.Fatalf("wrong secret must fail, got %v", err) } if err := VerifyWebhook([]byte(testWebhookSecret), nil, ""); !errors.Is(err, ErrBadSignature) { t.Fatalf("empty signature must fail, got %v", err) } } func TestParseWebhookVerifiesAndExtracts(t *testing.T) { body := []byte(`{"post":{"id":55,"post_number":2}}`) req := httptest.NewRequest(http.MethodPost, "/hooks/discourse", bytes.NewReader(body)) req.Header.Set("X-Discourse-Event-Id", "0d8417a0-1c2b-4d7b") req.Header.Set("X-Discourse-Event-Type", "post") req.Header.Set("X-Discourse-Event", "post_created") req.Header.Set("X-Discourse-Event-Signature", "sha256="+hmacHex(testWebhookSecret, body)) evt, err := ParseWebhook([]byte(testWebhookSecret), req) if err != nil { t.Fatalf("ParseWebhook: %v", err) } if evt.ID != "0d8417a0-1c2b-4d7b" || evt.Type != "post" || evt.Event != "post_created" { t.Fatalf("headers not extracted: %+v", evt) } var payload struct { Post struct { ID int `json:"id"` PostNumber int `json:"post_number"` } `json:"post"` } if err := json.Unmarshal(evt.Body, &payload); err != nil { t.Fatalf("body not preserved: %v", err) } if payload.Post.ID != 55 || payload.Post.PostNumber != 2 { t.Fatalf("body mangled: %s", evt.Body) } } func TestWebhookHandlerRejectsAndAccepts(t *testing.T) { body := []byte(`{"topic":{"id":12}}`) got := 0 h := WebhookHandler([]byte(testWebhookSecret), func(evt *WebhookEvent) { got++ }) r := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body)) r.Header.Set("X-Discourse-Event-Signature", hmacHex("wrong", body)) w := httptest.NewRecorder() h(w, r) if w.Code != http.StatusForbidden || got != 0 { t.Fatalf("bad sig: want 403 + no dispatch, got %d dispatch=%d", w.Code, got) } r = httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body)) r.Header.Set("X-Discourse-Event-Signature", "sha256="+hmacHex(testWebhookSecret, body)) r.Header.Set("X-Discourse-Event", "topic_created") w = httptest.NewRecorder() h(w, r) if w.Code != http.StatusOK || got != 1 { t.Fatalf("good sig: want 200 + dispatch, got %d dispatch=%d", w.Code, got) } } // --- SSO (Discourse Connect) --- func ssoPayload(t *testing.T, fields string) (string, string) { t.Helper() b64 := base64.StdEncoding.EncodeToString([]byte(fields)) return b64, hmacHex(testSSOSecret, []byte(b64)) } func TestVerifySSO(t *testing.T) { b64, sig := ssoPayload(t, "nonce=cb68251eefb35f7e&return_sso_url=https%3A%2F%2Fapp.example.test%2Fsso%2Fcallback") vals, err := VerifySSO([]byte(testSSOSecret), b64, sig) if err != nil { t.Fatalf("VerifySSO: %v", err) } if vals.Get("nonce") != "cb68251eefb35f7e" { t.Fatalf("nonce not decoded: %q", vals.Get("nonce")) } if vals.Get("return_sso_url") != "https://app.example.test/sso/callback" { t.Fatalf("return url not decoded: %q", vals.Get("return_sso_url")) } } func TestVerifySSORejectsTampering(t *testing.T) { b64, sig := ssoPayload(t, "nonce=abc") if _, err := VerifySSO([]byte(testSSOSecret), b64+"x", sig); !errors.Is(err, ErrBadSignature) { t.Fatalf("tampered payload must fail, got %v", err) } if _, err := VerifySSO([]byte("other"), b64, sig); !errors.Is(err, ErrBadSignature) { t.Fatalf("wrong secret must fail, got %v", err) } if _, err := VerifySSO([]byte(testSSOSecret), "!!not-base64!!", sig); err == nil { t.Fatal("undecodable payload must fail") } if _, err := VerifySSO([]byte(testSSOSecret), "", ""); !errors.Is(err, ErrBadSignature) { t.Fatalf("empty sso/sig must fail, got %v", err) } } func TestBuildSSOResponseRoundTrip(t *testing.T) { b64, sig, err := BuildSSOResponse([]byte(testSSOSecret), "cb68251eefb35f7e", map[string]string{ "email": "charles@turnsys.com", "external_id": "7", "username": "reachableceo", "admin": "true", }) if err != nil { t.Fatalf("BuildSSOResponse: %v", err) } vals, err := VerifySSO([]byte(testSSOSecret), b64, sig) if err != nil { t.Fatalf("round-trip verify: %v", err) } for k, want := range map[string]string{ "nonce": "cb68251eefb35f7e", "email": "charles@turnsys.com", "external_id": "7", "username": "reachableceo", "admin": "true", } { if got := vals.Get(k); got != want { t.Fatalf("field %s: want %q got %q", k, want, got) } } if _, err := base64.StdEncoding.DecodeString(b64); err != nil { t.Fatalf("response not std base64: %v", err) } } func TestBuildSSOResponseRequiresNonce(t *testing.T) { if _, _, err := BuildSSOResponse([]byte(testSSOSecret), "", map[string]string{"email": "x@y.test"}); !errors.Is(err, ErrInvalidRequest) { t.Fatalf("empty nonce: want ErrInvalidRequest, got %v", err) } } func TestSSORedirectURL(t *testing.T) { u, err := SSORedirectURL("https://app.example.test/sso/callback", "c29tZXBheWxvYWQ=", "aabbcc") if err != nil { t.Fatalf("SSORedirectURL: %v", err) } parsed, err := url.Parse(u) if err != nil { t.Fatalf("parse: %v", err) } if parsed.Scheme != "https" || parsed.Host != "app.example.test" { t.Fatalf("host mangled: %s", u) } q := parsed.Query() if q.Get("sso") != "c29tZXBheWxvYWQ=" || q.Get("sig") != "aabbcc" { t.Fatalf("query mangled: %s", u) } } // --- secret loading from env --- func TestSecretsFromEnv(t *testing.T) { t.Setenv("DISCOURSE_WEBHOOK_SECRET", " hook-secret ") t.Setenv("DISCOURSE_SSO_SECRET", " sso-secret ") wh, err := WebhookSecretFromEnv() if err != nil || string(wh) != "hook-secret" { t.Fatalf("webhook secret: %q %v", wh, err) } sso, err := SSOSecretFromEnv() if err != nil || string(sso) != "sso-secret" { t.Fatalf("sso secret: %q %v", sso, err) } } func TestSecretsFromEnvMissing(t *testing.T) { t.Setenv("DISCOURSE_WEBHOOK_SECRET", "") t.Setenv("DISCOURSE_SSO_SECRET", "") if _, err := WebhookSecretFromEnv(); !errors.Is(err, ErrInvalidRequest) { t.Fatalf("missing webhook secret: want ErrInvalidRequest, got %v", err) } if _, err := SSOSecretFromEnv(); !errors.Is(err, ErrInvalidRequest) { t.Fatalf("missing sso secret: want ErrInvalidRequest, got %v", err) } } func TestWebhookSecretNeverLeaksInError(t *testing.T) { body := []byte(`{}`) err := VerifyWebhook([]byte(testWebhookSecret), body, "bad") if strings.Contains(err.Error(), testWebhookSecret) { t.Fatalf("secret leaked in error: %v", err) } }