From c4b06ddc7d564de3e45285396eff214040cd2278 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=C3=89mile=20R=C3=A9?= Date: Tue, 3 Mar 2026 16:39:13 +0400 Subject: [PATCH] Add compliange page base URL in context and use it to validate redirect URLs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Émile Ré --- pkg/probod/probod.go | 11 ++++++++++- pkg/server/api/compliancepage/context.go | 10 ++++++++-- pkg/server/api/compliancepage/id_middleware.go | 11 ++++++++++- pkg/server/api/compliancepage/sni_middleware.go | 14 ++++++++++++++ pkg/server/api/trust/v1/graphql_handler.go | 2 -- pkg/server/api/trust/v1/resolver.go | 2 -- pkg/server/api/trust/v1/v1_resolver.go | 13 +++++++++++++ pkg/server/server.go | 6 +++--- 8 files changed, 58 insertions(+), 11 deletions(-) diff --git a/pkg/probod/probod.go b/pkg/probod/probod.go index 0eb563407..d4a947ef5 100644 --- a/pkg/probod/probod.go +++ b/pkg/probod/probod.go @@ -551,7 +551,16 @@ func (impl *Implm) Run( defer stopTrustCenterServer() wg.Go( func() { - if err := impl.runTrustCenterServer(trustCenterServerCtx, l, r, tp, pgClient, serverHandler.TrustCenterHandler(), acmeService, proboService); err != nil { + if err := impl.runTrustCenterServer( + trustCenterServerCtx, + l, + r, + tp, + pgClient, + serverHandler.TrustCenterHandler(), + acmeService, + proboService, + ); err != nil { cancel(fmt.Errorf("trust center server crashed: %w", err)) } }, diff --git a/pkg/server/api/compliancepage/context.go b/pkg/server/api/compliancepage/context.go index 58736896a..1dc971239 100644 --- a/pkg/server/api/compliancepage/context.go +++ b/pkg/server/api/compliancepage/context.go @@ -23,8 +23,9 @@ import ( type ctxKey struct{ name string } var ( - compliancePageKey = &ctxKey{name: "compliance_page"} - complianceMembershipKey = &ctxKey{name: "compliance_membership"} + compliancePageKey = &ctxKey{name: "compliance_page"} + complianceMembershipKey = &ctxKey{name: "compliance_membership"} + compliancePageBaseURLKey = &ctxKey{name: "compliance_page_base_url"} ) func CompliancePageFromContext(ctx context.Context) *coredata.TrustCenter { @@ -36,3 +37,8 @@ func ComplianceMembershipFromContext(ctx context.Context) *coredata.TrustCenterA membership, _ := ctx.Value(complianceMembershipKey).(*coredata.TrustCenterAccess) return membership } + +func CompliancePageBaseURLFromContext(ctx context.Context) *string { + page, _ := ctx.Value(compliancePageBaseURLKey).(*string) + return page +} diff --git a/pkg/server/api/compliancepage/id_middleware.go b/pkg/server/api/compliancepage/id_middleware.go index 68980cf02..f7482907a 100644 --- a/pkg/server/api/compliancepage/id_middleware.go +++ b/pkg/server/api/compliancepage/id_middleware.go @@ -23,12 +23,13 @@ import ( "github.com/go-chi/chi/v5" "github.com/vektah/gqlparser/v2/gqlerror" "go.gearno.de/kit/httpserver" + "go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/server/gqlutils" "go.probo.inc/probo/pkg/trust" ) -func NewIDMiddleware(trustSvc *trust.Service) func(next http.Handler) http.Handler { +func NewIDMiddleware(trustSvc *trust.Service, baseURL string) func(next http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc( func(w http.ResponseWriter, r *http.Request) { @@ -56,6 +57,10 @@ func NewIDMiddleware(trustSvc *trust.Service) func(next http.Handler) http.Handl return } + baseURL := baseurl.MustParse(baseURL).AppendPath("/trust/" + id.String()).MustString() + ctx = context.WithValue(ctx, compliancePageBaseURLKey, &baseURL) + r = r.WithContext(ctx) + if !compliancePage.Active { next.ServeHTTP(w, r) return @@ -85,6 +90,10 @@ func NewIDMiddleware(trustSvc *trust.Service) func(next http.Handler) http.Handl return } + baseURL := baseurl.MustParse(baseURL).AppendPath("/trust/" + value).MustString() + ctx = context.WithValue(ctx, compliancePageBaseURLKey, &baseURL) + r = r.WithContext(ctx) + if compliancePage.Active { ctx = context.WithValue(ctx, compliancePageKey, compliancePage) next.ServeHTTP(w, r.WithContext(ctx)) diff --git a/pkg/server/api/compliancepage/sni_middleware.go b/pkg/server/api/compliancepage/sni_middleware.go index 737264d73..0e6bc24bd 100644 --- a/pkg/server/api/compliancepage/sni_middleware.go +++ b/pkg/server/api/compliancepage/sni_middleware.go @@ -18,6 +18,7 @@ import ( "context" "errors" "net/http" + "net/url" "github.com/99designs/gqlgen/graphql" "github.com/vektah/gqlparser/v2/gqlerror" @@ -55,6 +56,19 @@ func NewSNIMiddleware(trustSvc *trust.Service) func(next http.Handler) http.Hand return } + baseURL := &url.URL{ + Host: r.Host, + Path: r.URL.Path, + Scheme: "https", + } + + ctx = context.WithValue( + ctx, + compliancePageBaseURLKey, + new(baseURL.String()), + ) + r = r.WithContext(ctx) + if compliancePage.Active { ctx = context.WithValue(ctx, compliancePageKey, compliancePage) next.ServeHTTP(w, r.WithContext(ctx)) diff --git a/pkg/server/api/trust/v1/graphql_handler.go b/pkg/server/api/trust/v1/graphql_handler.go index b40ca46de..31aa599a7 100644 --- a/pkg/server/api/trust/v1/graphql_handler.go +++ b/pkg/server/api/trust/v1/graphql_handler.go @@ -21,7 +21,6 @@ import ( "go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/esign" "go.probo.inc/probo/pkg/iam" - "go.probo.inc/probo/pkg/saferedirect" "go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/server/api/authn" "go.probo.inc/probo/pkg/server/api/trust/v1/schema" @@ -39,7 +38,6 @@ func NewGraphQLHandler(iamSvc *iam.Service, trustSvc *trust.Service, esignSvc *e logger: logger, baseURL: baseURL, sessionCookie: authn.NewCookie(&cookieConfig), - safeRedirect: &saferedirect.SafeRedirect{AllowedHost: baseURL.Host()}, }, Directives: schema.DirectiveRoot{ Nda: newNDADirectiveFunc(logger, trustSvc, esignSvc), diff --git a/pkg/server/api/trust/v1/resolver.go b/pkg/server/api/trust/v1/resolver.go index dcdd058b6..bb0b384cf 100644 --- a/pkg/server/api/trust/v1/resolver.go +++ b/pkg/server/api/trust/v1/resolver.go @@ -26,7 +26,6 @@ import ( "go.probo.inc/probo/pkg/esign" "go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/iam" - "go.probo.inc/probo/pkg/saferedirect" "go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/server/api/authn" "go.probo.inc/probo/pkg/server/api/compliancepage" @@ -53,7 +52,6 @@ type ( iam *iam.Service sessionCookie *authn.Cookie baseURL *baseurl.BaseURL - safeRedirect *saferedirect.SafeRedirect } ) diff --git a/pkg/server/api/trust/v1/v1_resolver.go b/pkg/server/api/trust/v1/v1_resolver.go index 863b59cd6..aa74c570a 100644 --- a/pkg/server/api/trust/v1/v1_resolver.go +++ b/pkg/server/api/trust/v1/v1_resolver.go @@ -14,11 +14,13 @@ import ( "time" "go.gearno.de/kit/log" + "go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/esign" "go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/iam" "go.probo.inc/probo/pkg/page" + "go.probo.inc/probo/pkg/saferedirect" "go.probo.inc/probo/pkg/server/api/authn" "go.probo.inc/probo/pkg/server/api/compliancepage" "go.probo.inc/probo/pkg/server/api/trust/v1/schema" @@ -157,6 +159,17 @@ func (r *frameworkResolver) DarkLogoURL(ctx context.Context, obj *types.Framewor func (r *mutationResolver) SendMagicLink(ctx context.Context, input types.SendMagicLinkInput) (*types.SendMagicLinkPayload, error) { trustCenter := compliancepage.CompliancePageFromContext(ctx) + baseURL := compliancepage.CompliancePageBaseURLFromContext(ctx) + + safeRedirect := &saferedirect.SafeRedirect{AllowedHost: baseurl.MustParse(*baseURL).Host()} + + if input.Continue != nil { + _, ok := safeRedirect.Validate(*input.Continue) + if !ok { + return nil, gqlutils.Invalidf(ctx, "invalid continue URL") + } + } + req := &iam.SendMagicLinkRequest{ Email: input.Email, CompliancePageID: &trustCenter.ID, diff --git a/pkg/server/server.go b/pkg/server/server.go index 4f63580c4..2de2d8017 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -109,16 +109,16 @@ func NewServer(cfg Config) (*Server, error) { logger: cfg.Logger, } - server.setupRoutes() + server.setupRoutes(cfg.BaseURL.String()) return server, nil } -func (s *Server) setupRoutes() { +func (s *Server) setupRoutes(baseURL string) { s.router.Mount("/api", http.StripPrefix("/api", s.apiServer)) s.router.Route("/trust/{slugOrId}", func(r chi.Router) { - r.Use(compliancepage.NewIDMiddleware(s.trustService)) + r.Use(compliancepage.NewIDMiddleware(s.trustService, baseURL)) r.Use(s.stripTrustPrefix) r.Mount("/", s.trustCenterRouter()) })