diff --git a/samlsp/new.go b/samlsp/new.go index 3145eec5..33397435 100644 --- a/samlsp/new.go +++ b/samlsp/new.go @@ -24,6 +24,7 @@ type Options struct { SignRequest bool ForceAuthn bool // TODO(ross): this should be *bool CookieSameSite http.SameSite + RelayStateFunc func(w http.ResponseWriter, r *http.Request) string } // DefaultSessionCodec returns the default SessionCodec for the provided options, @@ -72,6 +73,7 @@ func DefaultRequestTracker(opts Options, serviceProvider *saml.ServiceProvider) NamePrefix: "saml_", Codec: DefaultTrackedRequestCodec(opts), MaxAge: saml.MaxIssueDelay, + RelayStateFunc: opts.RelayStateFunc, SameSite: opts.CookieSameSite, } } diff --git a/samlsp/request_tracker_cookie.go b/samlsp/request_tracker_cookie.go index 0347253f..d9189f63 100644 --- a/samlsp/request_tracker_cookie.go +++ b/samlsp/request_tracker_cookie.go @@ -19,6 +19,7 @@ type CookieRequestTracker struct { NamePrefix string Codec TrackedRequestCodec MaxAge time.Duration + RelayStateFunc func(w http.ResponseWriter, r *http.Request) string SameSite http.SameSite } @@ -30,6 +31,14 @@ func (t CookieRequestTracker) TrackRequest(w http.ResponseWriter, r *http.Reques SAMLRequestID: samlRequestID, URI: r.URL.String(), } + + if t.RelayStateFunc != nil { + relayState := t.RelayStateFunc(w, r) + if relayState != "" { + trackedRequest.Index = relayState + } + } + signedTrackedRequest, err := t.Codec.Encode(trackedRequest) if err != nil { return "", err