package service import ( "bytes" "context" "crypto/aes" "crypto/cipher" "crypto/ecdh" "crypto/rand" "crypto/sha256" "encoding/base64" "errors" "fmt" "io" "net/http" "net/http/httptest" "testing" "github.com/gochat/gochat/internal/model" "github.com/gochat/gochat/internal/repository" "golang.org/x/crypto/hkdf" "gorm.io/driver/sqlite" "gorm.io/gorm" ) func TestDeliverWebPushProducesDecryptableRFC8291Record(t *testing.T) { receiverKey, err := ecdh.P256().GenerateKey(rand.Reader) if err != nil { t.Fatal(err) } vapidKey, err := ecdh.P256().GenerateKey(rand.Reader) if err != nil { t.Fatal(err) } authSecret := bytes.Repeat([]byte{0x42}, 16) payload := []byte(`{"title":"hello"}`) var encrypted []byte server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if got := r.Header.Get("Content-Encoding"); got != "aes128gcm" { t.Errorf("Content-Encoding = %q", got) } encrypted, err = io.ReadAll(r.Body) if err != nil { t.Errorf("read request body: %v", err) } w.WriteHeader(http.StatusCreated) })) defer server.Close() service := NewPushDeliveryService(nil, base64.RawURLEncoding.EncodeToString(vapidKey.PublicKey().Bytes()), base64.RawURLEncoding.EncodeToString(vapidKey.Bytes()), "mailto:test@example.com", ) service.httpClient = server.Client() token := model.PushToken{ Token: server.URL, P256DHKey: base64.RawURLEncoding.EncodeToString(receiverKey.PublicKey().Bytes()), AuthKey: base64.RawURLEncoding.EncodeToString(authSecret), } if err := service.deliverWebPush(context.Background(), token, payload); err != nil { t.Fatalf("deliver web push: %v", err) } decrypted, err := decryptRFC8291Record(encrypted, receiverKey, authSecret) if err != nil { t.Fatalf("decrypt RFC 8291 record: %v", err) } if !bytes.Equal(decrypted, payload) { t.Fatalf("decrypted payload = %q, want %q", decrypted, payload) } } func decryptRFC8291Record(record []byte, receiverKey *ecdh.PrivateKey, authSecret []byte) ([]byte, error) { if len(record) < 21 { return nil, fmt.Errorf("record too short: %d", len(record)) } salt := record[:16] keyIDLen := int(record[20]) if keyIDLen == 0 || len(record) < 21+keyIDLen { return nil, fmt.Errorf("invalid key id length: %d", keyIDLen) } senderPublic, err := ecdh.P256().NewPublicKey(record[21 : 21+keyIDLen]) if err != nil { return nil, err } sharedSecret, err := receiverKey.ECDH(senderPublic) if err != nil { return nil, err } info := append([]byte("WebPush: info\x00"), receiverKey.PublicKey().Bytes()...) info = append(info, senderPublic.Bytes()...) ikm, err := readHKDF(sharedSecret, authSecret, info, 32) if err != nil { return nil, err } cek, err := readHKDF(ikm, salt, []byte("Content-Encoding: aes128gcm\x00"), 16) if err != nil { return nil, err } nonce, err := readHKDF(ikm, salt, []byte("Content-Encoding: nonce\x00"), 12) if err != nil { return nil, err } block, err := aes.NewCipher(cek) if err != nil { return nil, err } gcm, err := cipher.NewGCM(block) if err != nil { return nil, err } plaintext, err := gcm.Open(nil, nonce, record[21+keyIDLen:], nil) if err != nil { return nil, err } delimiter := bytes.LastIndexByte(plaintext, 0x02) if delimiter < 0 || len(bytes.Trim(plaintext[delimiter+1:], "\x00")) != 0 { return nil, errors.New("invalid aes128gcm padding") } return plaintext[:delimiter], nil } func readHKDF(secret, salt, info []byte, size int) ([]byte, error) { value := make([]byte, size) _, err := io.ReadFull(hkdf.New(sha256.New, secret, salt, info), value) return value, err } func TestWebhookDeliveryPersistsAfterSubscriptionUpdateFailure(t *testing.T) { db := newWebhookRegressionDB(t) subscriptionErr := errors.New("subscription update failed") if err := db.Callback().Update().Before("gorm:update").Register("fail_subscription_update", func(tx *gorm.DB) { if _, ok := tx.Statement.Dest.(*model.WebhookSubscription); ok { _ = tx.AddError(subscriptionErr) } }); err != nil { t.Fatal(err) } server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusNoContent) })) defer server.Close() sub := model.WebhookSubscription{AccountID: 1, URL: server.URL, Events: []byte(`[]`), Secret: "secret", Active: true} if err := db.Create(&sub).Error; err != nil { t.Fatal(err) } service := NewWebhookDeliveryService(repository.NewWebhookSubscriptionRepo(db)) service.httpClient = server.Client() err := service.deliverToSubscription(context.Background(), sub, "message_created", []byte(`{"id":1}`)) if !errors.Is(err, subscriptionErr) { t.Fatalf("error = %v, want subscription update failure", err) } var delivery model.WebhookDelivery if err := db.First(&delivery).Error; err != nil { t.Fatal(err) } if delivery.Status != model.WebhookDeliveryStatusSuccess || delivery.Attempts != 1 || delivery.ResponseCode != http.StatusNoContent { t.Fatalf("delivery not finalized: %#v", delivery) } } func TestWebhookDeliveryJoinsResponseReadAndFinalUpdateFailures(t *testing.T) { db := newWebhookRegressionDB(t) readErr := errors.New("response read failed") updateErr := errors.New("delivery update failed") updateAttempted := false if err := db.Callback().Update().Before("gorm:update").Register("fail_delivery_update", func(tx *gorm.DB) { if _, ok := tx.Statement.Dest.(*model.WebhookDelivery); ok { updateAttempted = true _ = tx.AddError(updateErr) } }); err != nil { t.Fatal(err) } sub := model.WebhookSubscription{AccountID: 1, URL: "https://example.com/webhook", Events: []byte(`[]`), Secret: "secret", Active: true} if err := db.Create(&sub).Error; err != nil { t.Fatal(err) } service := NewWebhookDeliveryService(repository.NewWebhookSubscriptionRepo(db)) service.httpClient = &http.Client{Transport: roundTripperFunc(func(*http.Request) (*http.Response, error) { return &http.Response{StatusCode: http.StatusBadGateway, Body: errorReadCloser{err: readErr}}, nil })} err := service.deliverToSubscription(context.Background(), sub, "message_created", []byte(`{"id":1}`)) if !errors.Is(err, readErr) || !errors.Is(err, updateErr) || !updateAttempted { t.Fatalf("error = %v, update attempted = %v", err, updateAttempted) } } func TestDeliverEventReturnsDeliveryErrors(t *testing.T) { db := newWebhookRegressionDB(t) sub := model.WebhookSubscription{AccountID: 1, URL: "://invalid", Events: []byte(`["message_created"]`), Secret: "secret", Active: true} if err := db.Create(&sub).Error; err != nil { t.Fatal(err) } service := NewWebhookDeliveryService(repository.NewWebhookSubscriptionRepo(db)) if err := service.DeliverEvent(context.Background(), 1, "message_created", map[string]interface{}{"id": 1}); err == nil { t.Fatal("expected delivery error") } } func newWebhookRegressionDB(t *testing.T) *gorm.DB { t.Helper() db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{}) if err != nil { t.Fatal(err) } if err := db.AutoMigrate(&model.Account{}, &model.Inbox{}, &model.WebhookSubscription{}, &model.WebhookDelivery{}); err != nil { t.Fatal(err) } return db } type roundTripperFunc func(*http.Request) (*http.Response, error) func (fn roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) { return fn(req) } type errorReadCloser struct{ err error } func (r errorReadCloser) Read([]byte) (int, error) { return 0, r.err } func (errorReadCloser) Close() error { return nil }