Add cookie banner REST API with per-banner CORS middleware

Introduce /cookie-banner/v1/{bannerID}/config endpoint for the JS SDK.
The custom CORS middleware validates each request origin against the
specific banner being requested, preventing cross-customer leakage.

Signed-off-by: Émile Ré <emile@getprobo.com>
This commit is contained in:
Émile Ré
2026-04-13 17:37:09 +04:00
parent 41b57a61de
commit 36062310be
6 changed files with 247 additions and 3 deletions

View File

@@ -156,7 +156,7 @@ func (r *CreateCookieConsentRecordRequest) Validate() error {
return v.Error()
}
func canonicalizeOrigin(raw string) string {
func CanonicalizeOrigin(raw string) string {
u, err := url.Parse(raw)
if err != nil {
return raw
@@ -269,7 +269,7 @@ func (s *Service) CreateCookieBanner(
ID: gid.New(scope.GetTenantID(), coredata.CookieBannerEntityType),
OrganizationID: req.OrganizationID,
Name: req.Name,
Origin: canonicalizeOrigin(req.Origin),
Origin: CanonicalizeOrigin(req.Origin),
State: coredata.CookieBannerStateActive,
PrivacyPolicyURL: req.PrivacyPolicyURL,
ConsentExpiryDays: req.ConsentExpiryDays,
@@ -350,6 +350,58 @@ func (s *Service) GetCookieBanner(
return &banner, nil
}
func (s *Service) GetActiveCookieBanner(
ctx context.Context,
bannerID gid.GID,
) (*coredata.CookieBanner, error) {
var banner coredata.CookieBanner
err := s.pg.WithConn(
ctx,
func(ctx context.Context, conn pg.Querier) error {
if err := banner.LoadActiveByID(ctx, conn, bannerID); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return ErrBannerNotFound
}
return fmt.Errorf("cannot load cookie banner: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return &banner, nil
}
func (s *Service) GetActiveCookieBannerByOrigin(
ctx context.Context,
origin string,
) (*coredata.CookieBanner, error) {
var banner coredata.CookieBanner
err := s.pg.WithConn(
ctx,
func(ctx context.Context, conn pg.Querier) error {
if err := banner.LoadActiveByOrigin(ctx, conn, CanonicalizeOrigin(origin)); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return ErrBannerNotFound
}
return fmt.Errorf("cannot load cookie banner: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return &banner, nil
}
func (s *Service) ListCookieBannersForOrganization(
ctx context.Context,
scope coredata.Scoper,
@@ -432,7 +484,7 @@ func (s *Service) UpdateCookieBanner(
banner.Name = *req.Name
}
if req.Origin != nil {
banner.Origin = canonicalizeOrigin(*req.Origin)
banner.Origin = CanonicalizeOrigin(*req.Origin)
}
if req.PrivacyPolicyURL != nil {
banner.PrivacyPolicyURL = *req.PrivacyPolicyURL

View File

@@ -119,6 +119,52 @@ LIMIT 1;
return nil
}
func (b *CookieBanner) LoadActiveByID(
ctx context.Context,
conn pg.Querier,
bannerID gid.GID,
) error {
q := `
SELECT
id,
organization_id,
name,
origin,
state,
privacy_policy_url,
consent_expiry_days,
consent_mode,
created_at,
updated_at
FROM
cookie_banners
WHERE
id = @banner_id
AND state = 'ACTIVE'
LIMIT 1;
`
args := pgx.StrictNamedArgs{"banner_id": bannerID}
rows, err := conn.Query(ctx, q, args)
if err != nil {
return fmt.Errorf("cannot query cookie banners: %w", err)
}
banner, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[CookieBanner])
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return ErrResourceNotFound
}
return fmt.Errorf("cannot collect cookie banner: %w", err)
}
*b = banner
return nil
}
func (b *CookieBanner) LoadActiveByOrigin(
ctx context.Context,
conn pg.Querier,

View File

@@ -28,6 +28,7 @@ import (
"go.probo.inc/probo/pkg/accessreview"
"go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/connector"
"go.probo.inc/probo/pkg/cookiebanner"
"go.probo.inc/probo/pkg/esign"
"go.probo.inc/probo/pkg/file"
"go.probo.inc/probo/pkg/iam"
@@ -36,6 +37,7 @@ import (
"go.probo.inc/probo/pkg/securecookie"
connect_v1 "go.probo.inc/probo/pkg/server/api/connect/v1"
console_v1 "go.probo.inc/probo/pkg/server/api/console/v1"
cookiebanner_v1 "go.probo.inc/probo/pkg/server/api/cookiebanner/v1"
files_v1 "go.probo.inc/probo/pkg/server/api/files/v1"
mcp_v1 "go.probo.inc/probo/pkg/server/api/mcp/v1"
slack_v1 "go.probo.inc/probo/pkg/server/api/slack/v1"
@@ -56,6 +58,7 @@ type (
AccessReview *accessreview.Service
Slack *slack.Service
Mailman *mailman.Service
CookieBanner *cookiebanner.Service
Cookie securecookie.Config
TokenSecret string
ConnectorRegistry *connector.ConnectorRegistry
@@ -74,6 +77,7 @@ type (
csrf *http.CrossOriginProtection
compliancePageHandler http.Handler
consoleHandler http.Handler
cookieBannerHandler http.Handler
filesHandler http.Handler
mcpHandler http.Handler
slackHandler http.Handler
@@ -135,6 +139,12 @@ func NewServer(cfg Config) (*Server, error) {
// POSTs from external identity providers by design.
csrf.AddInsecureBypassPattern("POST /connect/v1/saml/2.0/consume")
// The cookie banner API is called cross-origin from customer websites
// by the JS SDK. CORS is handled by the cookie banner middleware.
csrf.AddInsecureBypassPattern("GET /cookie-banner/v1/*")
csrf.AddInsecureBypassPattern("POST /cookie-banner/v1/*")
csrf.AddInsecureBypassPattern("OPTIONS /cookie-banner/v1/*")
csrf.SetDenyHandler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
httpserver.RenderJSON(
w,
@@ -171,6 +181,10 @@ func NewServer(cfg Config) (*Server, error) {
cfg.BaseURL,
cfg.CustomDomainCname,
),
cookieBannerHandler: cookiebanner_v1.NewMux(
cfg.Logger.Named("cookiebanner.v1"),
cfg.CookieBanner,
),
filesHandler: files_v1.NewMux(
cfg.Logger.Named("files.v1"),
cfg.File,
@@ -240,6 +254,7 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
router.Mount("/console/v1", http.StripPrefix("/console/v1", s.consoleHandler))
router.Mount("/connect/v1", http.StripPrefix("/connect/v1", s.connectHandler))
router.Mount("/cookie-banner/v1", http.StripPrefix("/cookie-banner/v1", s.cookieBannerHandler))
router.Mount("/files/v1", http.StripPrefix("/files/v1", s.filesHandler))
router.Mount("/trust/v1", http.StripPrefix("/trust/v1", s.compliancePageHandler))
router.Mount("/mcp/v1", http.StripPrefix("/mcp/v1", s.mcpHandler))

View File

@@ -0,0 +1,79 @@
// 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.
package cookiebanner_v1
import (
"errors"
"net/http"
"github.com/go-chi/chi/v5"
"go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/cookiebanner"
"go.probo.inc/probo/pkg/gid"
)
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
}
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) {
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
}
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)
})
}
}

View File

@@ -0,0 +1,49 @@
// 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.
package cookiebanner_v1
import (
"net/http"
"github.com/go-chi/chi/v5"
"go.gearno.de/kit/httpserver"
"go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/cookiebanner"
)
type Handler struct {
logger *log.Logger
cookieBannerSvc *cookiebanner.Service
}
func NewMux(
logger *log.Logger,
cookieBannerSvc *cookiebanner.Service,
) *chi.Mux {
h := &Handler{
logger: logger,
cookieBannerSvc: cookieBannerSvc,
}
r := chi.NewMux()
r.Use(newCORSMiddleware(logger, cookieBannerSvc))
r.Get("/{bannerID}/config", h.handleGetConfig)
return r
}
func (h *Handler) handleGetConfig(w http.ResponseWriter, r *http.Request) {
httpserver.RenderJSON(w, http.StatusOK, map[string]string{"status": "ok"})
}

View File

@@ -27,6 +27,7 @@ import (
"go.probo.inc/probo/pkg/accessreview"
"go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/connector"
"go.probo.inc/probo/pkg/cookiebanner"
"go.probo.inc/probo/pkg/esign"
"go.probo.inc/probo/pkg/file"
"go.probo.inc/probo/pkg/iam"
@@ -54,6 +55,7 @@ type Config struct {
AccessReview *accessreview.Service
Slack *slack.Service
Mailman *mailman.Service
CookieBanner *cookiebanner.Service
Cookie securecookie.Config
TokenSecret string
ConnectorRegistry *connector.ConnectorRegistry
@@ -85,6 +87,7 @@ func NewServer(cfg Config) (*Server, error) {
AccessReview: cfg.AccessReview,
Slack: cfg.Slack,
Mailman: cfg.Mailman,
CookieBanner: cfg.CookieBanner,
Cookie: cfg.Cookie,
TokenSecret: cfg.TokenSecret,
ConnectorRegistry: cfg.ConnectorRegistry,