From c8586346be98953e62e2b17735fc47903202e889 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=C3=89mile=20R=C3=A9?= Date: Tue, 14 Apr 2026 13:16:24 +0400 Subject: [PATCH] Review fixes MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Émile Ré --- pkg/coredata/cookie_banner.go | 13 +-- .../api/cookiebanner/v1/cors_middleware.go | 84 ++++++++++--------- 2 files changed, 51 insertions(+), 46 deletions(-) diff --git a/pkg/coredata/cookie_banner.go b/pkg/coredata/cookie_banner.go index 8afec6a3a..5c2ce14f2 100644 --- a/pkg/coredata/cookie_banner.go +++ b/pkg/coredata/cookie_banner.go @@ -168,6 +168,7 @@ LIMIT 1; func (b *CookieBanner) LoadActiveByOrigin( ctx context.Context, conn pg.Querier, + scope Scoper, origin string, ) error { q := ` @@ -185,12 +186,16 @@ SELECT FROM cookie_banners WHERE - origin = @origin + %s + AND origin = @origin AND state = 'ACTIVE' LIMIT 1; ` + q = fmt.Sprintf(q, scope.SQLFragment()) + args := pgx.StrictNamedArgs{"origin": origin} + maps.Copy(args, scope.SQLArguments()) rows, err := conn.Query(ctx, q, args) if err != nil { @@ -345,8 +350,7 @@ INSERT INTO cookie_banners ( _, err := tx.Exec(ctx, q, args) if err != nil { - var pgErr *pgconn.PgError - if errors.As(err, &pgErr) { + if pgErr, ok := errors.AsType[*pgconn.PgError](err); ok { if pgErr.Code == "23505" && pgErr.ConstraintName == "idx_cookie_banners_unique_active_origin" { return ErrResourceAlreadyExists } @@ -393,8 +397,7 @@ WHERE result, err := tx.Exec(ctx, q, args) if err != nil { - var pgErr *pgconn.PgError - if errors.As(err, &pgErr) { + if pgErr, ok := errors.AsType[*pgconn.PgError](err); ok { if pgErr.Code == "23505" && pgErr.ConstraintName == "idx_cookie_banners_unique_active_origin" { return ErrResourceAlreadyExists } diff --git a/pkg/server/api/cookiebanner/v1/cors_middleware.go b/pkg/server/api/cookiebanner/v1/cors_middleware.go index e83b1661a..00a14a185 100644 --- a/pkg/server/api/cookiebanner/v1/cors_middleware.go +++ b/pkg/server/api/cookiebanner/v1/cors_middleware.go @@ -26,54 +26,56 @@ import ( func newCORSMiddleware(logger *log.Logger, cookieBannerSvc *cookiebanner.Service) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - origin := r.Header.Get("Origin") - if origin == "" { - next.ServeHTTP(w, r) - return - } + return http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + origin := r.Header.Get("Origin") + if origin == "" { + next.ServeHTTP(w, r) + return + } - bannerIDStr := chi.URLParam(r, "bannerID") - if bannerIDStr == "" { - http.Error(w, "forbidden", http.StatusForbidden) - return - } - - bannerID, err := gid.ParseGID(bannerIDStr) - if err != nil { - http.Error(w, "forbidden", http.StatusForbidden) - return - } - - banner, err := cookieBannerSvc.GetActiveCookieBanner(r.Context(), bannerID) - if err != nil { - if errors.Is(err, cookiebanner.ErrBannerNotFound) { + bannerIDStr := chi.URLParam(r, "bannerID") + if bannerIDStr == "" { http.Error(w, "forbidden", http.StatusForbidden) return } - logger.ErrorCtx(r.Context(), "cannot load cookie banner for CORS check", log.Error(err)) - http.Error(w, "internal server error", http.StatusInternalServerError) - return - } - canonicalOrigin := cookiebanner.CanonicalizeOrigin(origin) - if banner.Origin != canonicalOrigin { - http.Error(w, "forbidden", http.StatusForbidden) - return - } + bannerID, err := gid.ParseGID(bannerIDStr) + if err != nil { + http.Error(w, "forbidden", http.StatusForbidden) + return + } - w.Header().Set("Access-Control-Allow-Origin", origin) - w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") - w.Header().Set("Access-Control-Allow-Headers", "Content-Type") - w.Header().Set("Access-Control-Max-Age", "600") - w.Header().Set("Vary", "Origin") + banner, err := cookieBannerSvc.GetActiveCookieBanner(r.Context(), bannerID) + if err != nil { + if errors.Is(err, cookiebanner.ErrBannerNotFound) { + http.Error(w, "forbidden", http.StatusForbidden) + return + } + logger.ErrorCtx(r.Context(), "cannot load cookie banner for CORS check", log.Error(err)) + http.Error(w, "internal server error", http.StatusInternalServerError) + return + } - if r.Method == http.MethodOptions { - w.WriteHeader(http.StatusNoContent) - return - } + canonicalOrigin := cookiebanner.CanonicalizeOrigin(origin) + if banner.Origin != canonicalOrigin { + http.Error(w, "forbidden", http.StatusForbidden) + return + } - next.ServeHTTP(w, r) - }) + w.Header().Set("Access-Control-Allow-Origin", origin) + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") + w.Header().Set("Access-Control-Allow-Headers", "Content-Type") + w.Header().Set("Access-Control-Max-Age", "600") + w.Header().Set("Vary", "Origin") + + if r.Method == http.MethodOptions { + w.WriteHeader(http.StatusNoContent) + return + } + + next.ServeHTTP(w, r) + }, + ) } }