Handle cookie banner origin validation + unicity
Signed-off-by: Émile Ré <emile@getprobo.com>
This commit is contained in:
@@ -25,4 +25,5 @@ var (
|
|||||||
ErrVersionNotPublished = errors.New("cookie banner version is not published")
|
ErrVersionNotPublished = errors.New("cookie banner version is not published")
|
||||||
ErrNoDraftVersion = errors.New("no draft cookie banner version to publish")
|
ErrNoDraftVersion = errors.New("no draft cookie banner version to publish")
|
||||||
ErrCannotDeleteRequiredCategory = errors.New("cannot delete required cookie category")
|
ErrCannotDeleteRequiredCategory = errors.New("cannot delete required cookie category")
|
||||||
|
ErrOriginAlreadyInUse = errors.New("origin is already used by another active cookie banner")
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -100,7 +100,7 @@ func (r *CreateCookieBannerRequest) Validate() error {
|
|||||||
|
|
||||||
v.Check(r.OrganizationID, "organization_id", validator.Required(), validator.GID(coredata.OrganizationEntityType))
|
v.Check(r.OrganizationID, "organization_id", validator.Required(), validator.GID(coredata.OrganizationEntityType))
|
||||||
v.Check(r.Name, "name", validator.Required(), validator.SafeTextNoNewLine(255))
|
v.Check(r.Name, "name", validator.Required(), validator.SafeTextNoNewLine(255))
|
||||||
v.Check(r.Origin, "origin", validator.Required(), validator.NotEmpty())
|
v.Check(r.Origin, "origin", validator.Required(), validator.Origin())
|
||||||
v.Check(r.PrivacyPolicyURL, "privacy_policy_url", validator.Required(), validator.URL())
|
v.Check(r.PrivacyPolicyURL, "privacy_policy_url", validator.Required(), validator.URL())
|
||||||
v.Check(r.ConsentExpiryDays, "consent_expiry_days", validator.Required(), validator.Min(1))
|
v.Check(r.ConsentExpiryDays, "consent_expiry_days", validator.Required(), validator.Min(1))
|
||||||
v.Check(r.ConsentMode, "consent_mode", validator.Required(), validator.OneOfSlice(coredata.CookieConsentModes()))
|
v.Check(r.ConsentMode, "consent_mode", validator.Required(), validator.OneOfSlice(coredata.CookieConsentModes()))
|
||||||
@@ -113,7 +113,7 @@ func (r *UpdateCookieBannerRequest) Validate() error {
|
|||||||
|
|
||||||
v.Check(r.CookieBannerID, "cookie_banner_id", validator.Required(), validator.GID(coredata.CookieBannerEntityType))
|
v.Check(r.CookieBannerID, "cookie_banner_id", validator.Required(), validator.GID(coredata.CookieBannerEntityType))
|
||||||
v.Check(r.Name, "name", validator.SafeTextNoNewLine(255))
|
v.Check(r.Name, "name", validator.SafeTextNoNewLine(255))
|
||||||
v.Check(r.Origin, "origin", validator.NotEmpty())
|
v.Check(r.Origin, "origin", validator.Origin())
|
||||||
v.Check(r.PrivacyPolicyURL, "privacy_policy_url", validator.URL())
|
v.Check(r.PrivacyPolicyURL, "privacy_policy_url", validator.URL())
|
||||||
v.Check(r.ConsentExpiryDays, "consent_expiry_days", validator.Min(1))
|
v.Check(r.ConsentExpiryDays, "consent_expiry_days", validator.Min(1))
|
||||||
v.Check(r.ConsentMode, "consent_mode", validator.OneOfSlice(coredata.CookieConsentModes()))
|
v.Check(r.ConsentMode, "consent_mode", validator.OneOfSlice(coredata.CookieConsentModes()))
|
||||||
@@ -260,6 +260,9 @@ func (s *Service) CreateCookieBanner(
|
|||||||
}
|
}
|
||||||
|
|
||||||
if err := banner.Insert(ctx, tx, scope); err != nil {
|
if err := banner.Insert(ctx, tx, scope); err != nil {
|
||||||
|
if errors.Is(err, coredata.ErrResourceAlreadyExists) {
|
||||||
|
return ErrOriginAlreadyInUse
|
||||||
|
}
|
||||||
return fmt.Errorf("cannot insert cookie banner: %w", err)
|
return fmt.Errorf("cannot insert cookie banner: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -425,6 +428,9 @@ func (s *Service) UpdateCookieBanner(
|
|||||||
banner.UpdatedAt = time.Now()
|
banner.UpdatedAt = time.Now()
|
||||||
|
|
||||||
if err := banner.Update(ctx, tx, scope); err != nil {
|
if err := banner.Update(ctx, tx, scope); err != nil {
|
||||||
|
if errors.Is(err, coredata.ErrResourceAlreadyExists) {
|
||||||
|
return ErrOriginAlreadyInUse
|
||||||
|
}
|
||||||
return fmt.Errorf("cannot update cookie banner: %w", err)
|
return fmt.Errorf("cannot update cookie banner: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -512,6 +518,9 @@ func (s *Service) ActivateCookieBanner(
|
|||||||
banner.UpdatedAt = time.Now()
|
banner.UpdatedAt = time.Now()
|
||||||
|
|
||||||
if err := banner.Update(ctx, tx, scope); err != nil {
|
if err := banner.Update(ctx, tx, scope); err != nil {
|
||||||
|
if errors.Is(err, coredata.ErrResourceAlreadyExists) {
|
||||||
|
return ErrOriginAlreadyInUse
|
||||||
|
}
|
||||||
return fmt.Errorf("cannot update cookie banner: %w", err)
|
return fmt.Errorf("cannot update cookie banner: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/jackc/pgx/v5"
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
"go.gearno.de/kit/pg"
|
"go.gearno.de/kit/pg"
|
||||||
"go.probo.inc/probo/pkg/gid"
|
"go.probo.inc/probo/pkg/gid"
|
||||||
"go.probo.inc/probo/pkg/page"
|
"go.probo.inc/probo/pkg/page"
|
||||||
@@ -118,10 +119,10 @@ LIMIT 1;
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *CookieBanner) LoadActiveByID(
|
func (b *CookieBanner) LoadActiveByOrigin(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
conn pg.Querier,
|
conn pg.Querier,
|
||||||
bannerID gid.GID,
|
origin string,
|
||||||
) error {
|
) error {
|
||||||
q := `
|
q := `
|
||||||
SELECT
|
SELECT
|
||||||
@@ -138,12 +139,12 @@ SELECT
|
|||||||
FROM
|
FROM
|
||||||
cookie_banners
|
cookie_banners
|
||||||
WHERE
|
WHERE
|
||||||
id = @banner_id
|
origin = @origin
|
||||||
AND state = 'ACTIVE'
|
AND state = 'ACTIVE'
|
||||||
LIMIT 1;
|
LIMIT 1;
|
||||||
`
|
`
|
||||||
|
|
||||||
args := pgx.StrictNamedArgs{"banner_id": bannerID}
|
args := pgx.StrictNamedArgs{"origin": origin}
|
||||||
|
|
||||||
rows, err := conn.Query(ctx, q, args)
|
rows, err := conn.Query(ctx, q, args)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -298,6 +299,12 @@ INSERT INTO cookie_banners (
|
|||||||
|
|
||||||
_, err := tx.Exec(ctx, q, args)
|
_, err := tx.Exec(ctx, q, args)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
var pgErr *pgconn.PgError
|
||||||
|
if errors.As(err, &pgErr) {
|
||||||
|
if pgErr.Code == "23505" && pgErr.ConstraintName == "idx_cookie_banners_unique_active_origin" {
|
||||||
|
return ErrResourceAlreadyExists
|
||||||
|
}
|
||||||
|
}
|
||||||
return fmt.Errorf("cannot insert cookie banner: %w", err)
|
return fmt.Errorf("cannot insert cookie banner: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -340,6 +347,12 @@ WHERE
|
|||||||
|
|
||||||
result, err := tx.Exec(ctx, q, args)
|
result, err := tx.Exec(ctx, q, args)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
var pgErr *pgconn.PgError
|
||||||
|
if errors.As(err, &pgErr) {
|
||||||
|
if pgErr.Code == "23505" && pgErr.ConstraintName == "idx_cookie_banners_unique_active_origin" {
|
||||||
|
return ErrResourceAlreadyExists
|
||||||
|
}
|
||||||
|
}
|
||||||
return fmt.Errorf("cannot update cookie banner: %w", err)
|
return fmt.Errorf("cannot update cookie banner: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
16
pkg/coredata/migrations/20260413T114822Z.sql
Normal file
16
pkg/coredata/migrations/20260413T114822Z.sql
Normal file
@@ -0,0 +1,16 @@
|
|||||||
|
-- Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
--
|
||||||
|
-- Permission to use, copy, modify, and/or distribute this software for any
|
||||||
|
-- purpose with or without fee is hereby granted, provided that the above
|
||||||
|
-- copyright notice and this permission notice appear in all copies.
|
||||||
|
--
|
||||||
|
-- THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||||
|
-- REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||||
|
-- AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||||
|
-- INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||||
|
-- LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||||
|
-- OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
|
-- PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
|
CREATE UNIQUE INDEX idx_cookie_banners_unique_active_origin
|
||||||
|
ON cookie_banners (tenant_id, origin) WHERE state = 'ACTIVE';
|
||||||
@@ -133,6 +133,49 @@ func GID(entityTypes ...uint16) ValidatorFunc {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Origin validates that a string is a valid web origin (scheme + host + optional port).
|
||||||
|
// No path, query, fragment, or userinfo is allowed.
|
||||||
|
func Origin() ValidatorFunc {
|
||||||
|
return func(value any) *ValidationError {
|
||||||
|
actualValue, isNil := dereferenceValue(value)
|
||||||
|
if isNil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
str, ok := actualValue.(string)
|
||||||
|
if !ok {
|
||||||
|
return newValidationError(ErrorCodeInvalidFormat, "value must be a string")
|
||||||
|
}
|
||||||
|
|
||||||
|
if str == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
parsedURL, err := url.Parse(str)
|
||||||
|
if err != nil {
|
||||||
|
return newValidationError(ErrorCodeInvalidFormat, "must be a valid origin (e.g. https://example.com)")
|
||||||
|
}
|
||||||
|
|
||||||
|
if parsedURL.Scheme != "http" && parsedURL.Scheme != "https" {
|
||||||
|
return newValidationError(ErrorCodeInvalidFormat, "must be a valid origin (e.g. https://example.com)")
|
||||||
|
}
|
||||||
|
|
||||||
|
if parsedURL.Host == "" {
|
||||||
|
return newValidationError(ErrorCodeInvalidFormat, "must be a valid origin (e.g. https://example.com)")
|
||||||
|
}
|
||||||
|
|
||||||
|
if parsedURL.Path != "" && parsedURL.Path != "/" {
|
||||||
|
return newValidationError(ErrorCodeInvalidFormat, "must be a valid origin (e.g. https://example.com)")
|
||||||
|
}
|
||||||
|
|
||||||
|
if parsedURL.RawQuery != "" || parsedURL.Fragment != "" || parsedURL.User != nil {
|
||||||
|
return newValidationError(ErrorCodeInvalidFormat, "must be a valid origin (e.g. https://example.com)")
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Domain validates that a string is a valid domain name.
|
// Domain validates that a string is a valid domain name.
|
||||||
func Domain() ValidatorFunc {
|
func Domain() ValidatorFunc {
|
||||||
return func(value any) *ValidationError {
|
return func(value any) *ValidationError {
|
||||||
|
|||||||
@@ -128,6 +128,42 @@ func TestHTTPSUrl(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestOrigin(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
value any
|
||||||
|
wantError bool
|
||||||
|
}{
|
||||||
|
{"valid https origin", "https://example.com", false},
|
||||||
|
{"valid http origin", "http://example.com", false},
|
||||||
|
{"valid with port", "http://localhost:3000", false},
|
||||||
|
{"valid https with port", "https://example.com:8443", false},
|
||||||
|
{"valid with trailing slash", "https://example.com/", false},
|
||||||
|
{"invalid - has path", "https://example.com/path", true},
|
||||||
|
{"invalid - has query", "https://example.com?q=1", true},
|
||||||
|
{"invalid - has fragment", "https://example.com#section", true},
|
||||||
|
{"invalid - has userinfo", "https://user:pass@example.com", true},
|
||||||
|
{"invalid - no scheme", "example.com", true},
|
||||||
|
{"invalid - ftp scheme", "ftp://example.com", true},
|
||||||
|
{"invalid - no host", "https://", true},
|
||||||
|
{"empty string", "", false},
|
||||||
|
{"nil pointer", (*string)(nil), false},
|
||||||
|
{"non-string", 123, true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
err := Origin()(tt.value)
|
||||||
|
if (err != nil) != tt.wantError {
|
||||||
|
t.Errorf("Origin() error = %v, wantError %v", err, tt.wantError)
|
||||||
|
}
|
||||||
|
if err != nil && err.Code != ErrorCodeInvalidFormat {
|
||||||
|
t.Errorf("Expected error code %s, got %s", ErrorCodeInvalidFormat, err.Code)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestDomain(t *testing.T) {
|
func TestDomain(t *testing.T) {
|
||||||
t.Run("valid domain", func(t *testing.T) {
|
t.Run("valid domain", func(t *testing.T) {
|
||||||
str := "example.com"
|
str := "example.com"
|
||||||
|
|||||||
Reference in New Issue
Block a user