diff --git a/.env.example b/.env.example index baf02c49..e05ff2d4 100644 --- a/.env.example +++ b/.env.example @@ -167,6 +167,8 @@ TINYAUTH_OAUTH_PROVIDERS_name_AUTHURL= TINYAUTH_OAUTH_PROVIDERS_name_TOKENURL= # OAuth userinfo URL. TINYAUTH_OAUTH_PROVIDERS_name_USERINFOURL= +# OpenID Connect RP-Initiated Logout end_session_endpoint URL. +TINYAUTH_OAUTH_PROVIDERS_name_LOGOUTURL= # Allow insecure OAuth connections. TINYAUTH_OAUTH_PROVIDERS_name_INSECURE=false # Provider name in UI. @@ -194,6 +196,8 @@ TINYAUTH_OIDC_CLIENTS_name_CLIENTSECRET= TINYAUTH_OIDC_CLIENTS_name_CLIENTSECRETFILE= # List of trusted redirect URIs. TINYAUTH_OIDC_CLIENTS_name_TRUSTEDREDIRECTURIS= +# List of trusted post-logout redirect URIs. +TINYAUTH_OIDC_CLIENTS_name_TRUSTEDPOSTLOGOUTREDIRECTURIS= # Client name in UI. TINYAUTH_OIDC_CLIENTS_name_NAME= diff --git a/docker-compose.dev.yml b/docker-compose.dev.yml index eb4f7ce8..082894cd 100644 --- a/docker-compose.dev.yml +++ b/docker-compose.dev.yml @@ -13,6 +13,8 @@ services: labels: traefik.enable: true traefik.http.routers.whoami.rule: Host(`whoami.127.0.0.1.sslip.io`) + traefik.http.routers.whoami.entrypoints: websecure + traefik.http.routers.whoami.tls: true traefik.http.routers.whoami.middlewares: tinyauth tinyauth-frontend: diff --git a/frontend/src/components/quick-actions/quick-actions.tsx b/frontend/src/components/quick-actions/quick-actions.tsx index cc33b2f4..5890dee8 100644 --- a/frontend/src/components/quick-actions/quick-actions.tsx +++ b/frontend/src/components/quick-actions/quick-actions.tsx @@ -77,6 +77,10 @@ export const QuickActions = () => { } return ""; })(); + const logoutParams = + screenParams.redirect_uri && screenParams.login_for !== "oidc" + ? { login_for: "app", redirect_uri: screenParams.redirect_uri } + : undefined; const [isOpen, setIsOpen] = useState(false); @@ -122,15 +126,25 @@ export const QuickActions = () => { })(); const logoutMutation = useMutation({ - mutationFn: () => axios.post("/api/user/logout"), + // redirect_uri is Tinyauth's existing application-navigation parameter. + // It is not the OIDC RP-Initiated Logout post_logout_redirect_uri. + mutationFn: () => + axios.post("/api/user/logout", undefined, { + params: logoutParams, + }), mutationKey: ["logout"], - onSuccess: () => { + onSuccess: (response) => { toast.success(t("logoutSuccessTitle"), { description: t("logoutSuccessSubtitle"), }); + const redirectUrl = response.data?.redirectUrl; redirectTimer.current = window.setTimeout(() => { - window.location.replace(`/login${compiledParams}`); + if (typeof redirectUrl === "string" && redirectUrl.length > 0) { + window.location.replace(redirectUrl); + } else { + window.location.replace(`/login${compiledParams}`); + } }, 500); }, onError: () => { diff --git a/frontend/src/pages/logout-page.tsx b/frontend/src/pages/logout-page.tsx index 78ef0555..ffb04bb9 100644 --- a/frontend/src/pages/logout-page.tsx +++ b/frontend/src/pages/logout-page.tsx @@ -36,17 +36,31 @@ export const LogoutPage = () => { } return ""; })(); + const logoutParams = + screenParams.redirect_uri && screenParams.login_for !== "oidc" + ? { login_for: "app", redirect_uri: screenParams.redirect_uri } + : undefined; const logoutMutation = useMutation({ - mutationFn: () => axios.post("/api/user/logout"), + // redirect_uri is Tinyauth's existing application-navigation parameter. + // It is not the OIDC RP-Initiated Logout post_logout_redirect_uri. + mutationFn: () => + axios.post("/api/user/logout", undefined, { + params: logoutParams, + }), mutationKey: ["logout"], - onSuccess: () => { + onSuccess: (response) => { toast.success(t("logoutSuccessTitle"), { description: t("logoutSuccessSubtitle"), }); + const redirectUrl = response.data?.redirectUrl; redirectTimer.current = window.setTimeout(() => { - window.location.replace(`/login${compiledParams}`); + if (typeof redirectUrl === "string" && redirectUrl.length > 0) { + window.location.replace(redirectUrl); + } else { + window.location.replace(`/login${compiledParams}`); + } }, 500); }, onError: () => { diff --git a/internal/assets/migrations/postgres/000004_oauth_id_token.down.sql b/internal/assets/migrations/postgres/000004_oauth_id_token.down.sql new file mode 100644 index 00000000..5b72180e --- /dev/null +++ b/internal/assets/migrations/postgres/000004_oauth_id_token.down.sql @@ -0,0 +1 @@ +ALTER TABLE "sessions" DROP COLUMN "oauth_id_token"; diff --git a/internal/assets/migrations/postgres/000004_oauth_id_token.up.sql b/internal/assets/migrations/postgres/000004_oauth_id_token.up.sql new file mode 100644 index 00000000..6faec95d --- /dev/null +++ b/internal/assets/migrations/postgres/000004_oauth_id_token.up.sql @@ -0,0 +1 @@ +ALTER TABLE "sessions" ADD COLUMN "oauth_id_token" TEXT NOT NULL DEFAULT ''; diff --git a/internal/assets/migrations/sqlite/000012_oauth_id_token.down.sql b/internal/assets/migrations/sqlite/000012_oauth_id_token.down.sql new file mode 100644 index 00000000..5b72180e --- /dev/null +++ b/internal/assets/migrations/sqlite/000012_oauth_id_token.down.sql @@ -0,0 +1 @@ +ALTER TABLE "sessions" DROP COLUMN "oauth_id_token"; diff --git a/internal/assets/migrations/sqlite/000012_oauth_id_token.up.sql b/internal/assets/migrations/sqlite/000012_oauth_id_token.up.sql new file mode 100644 index 00000000..6faec95d --- /dev/null +++ b/internal/assets/migrations/sqlite/000012_oauth_id_token.up.sql @@ -0,0 +1 @@ +ALTER TABLE "sessions" ADD COLUMN "oauth_id_token" TEXT NOT NULL DEFAULT ''; diff --git a/internal/controller/oauth_controller.go b/internal/controller/oauth_controller.go index fd6c2658..99c5bfef 100644 --- a/internal/controller/oauth_controller.go +++ b/internal/controller/oauth_controller.go @@ -160,7 +160,7 @@ func (controller *OAuthController) oauthCallbackHandler(c *gin.Context) { } code := c.Query("code") - _, err = controller.auth.GetOAuthToken(sessionIdCookie, code) + token, err := controller.auth.GetOAuthToken(sessionIdCookie, code) if err != nil { controller.log.App.Error().Err(err).Msg("Failed to exchange code for token") @@ -235,6 +235,9 @@ func (controller *OAuthController) oauthCallbackHandler(c *gin.Context) { OAuthName: svc.Name(), OAuthSub: user.Sub, } + if idToken, ok := token.Extra("id_token").(string); ok { + sessionCookie.OAuthIDToken = idToken + } controller.log.App.Debug().Msg("Creating session cookie for user") diff --git a/internal/controller/oidc_controller.go b/internal/controller/oidc_controller.go index 9064135f..7e51e48f 100644 --- a/internal/controller/oidc_controller.go +++ b/internal/controller/oidc_controller.go @@ -34,6 +34,7 @@ type authorizeErrorParams struct { type OIDCController struct { log *logger.Logger oidc *service.OIDCService + auth *service.AuthService runtime *model.RuntimeConfig } @@ -88,6 +89,7 @@ type OIDCControllerInput struct { Log *logger.Logger OIDCService *service.OIDCService + AuthService *service.AuthService RuntimeConfig *model.RuntimeConfig RouterGroup *gin.RouterGroup `name:"apiRouterGroup"` MainRouter *gin.RouterGroup `name:"mainRouterGroup"` @@ -97,6 +99,7 @@ func NewOIDCController(i OIDCControllerInput) *OIDCController { controller := &OIDCController{ log: i.Log, oidc: i.OIDCService, + auth: i.AuthService, runtime: i.RuntimeConfig, } @@ -105,6 +108,8 @@ func NewOIDCController(i OIDCControllerInput) *OIDCController { oidcGroup := i.RouterGroup.Group("/oidc") oidcGroup.POST("/authorize-complete", controller.authorizeComplete) + oidcGroup.GET("/end-session", controller.endSession) + oidcGroup.POST("/end-session", controller.endSession) oidcGroup.POST("/token", controller.Token) oidcGroup.GET("/userinfo", controller.Userinfo) oidcGroup.POST("/userinfo", controller.Userinfo) @@ -112,6 +117,78 @@ func NewOIDCController(i OIDCControllerInput) *OIDCController { return controller } +func (controller *OIDCController) endSession(c *gin.Context) { + if controller.oidc == nil { + c.JSON(http.StatusNotFound, gin.H{ + "status": http.StatusNotFound, + "message": "OIDC service not configured", + }) + return + } + + req := service.EndSessionRequest{} + err := c.ShouldBind(&req) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{ + "status": http.StatusBadRequest, + "message": "Bad Request", + }) + return + } + + userContext, err := new(model.UserContext).NewFromGin(c) + if err != nil { + userContext = nil + } + + redirectURI, err := controller.oidc.ValidateEndSessionRequest(c, req, userContext) + if err != nil { + if errors.Is(err, service.ErrEndSessionConfirmationNeeded) { + c.Redirect(http.StatusFound, controller.runtime.AppURL+"/logout") + return + } + controller.log.App.Warn().Err(err).Msg("Rejected OIDC end-session request") + c.JSON(http.StatusBadRequest, gin.H{ + "status": http.StatusBadRequest, + "message": "Invalid end-session request", + }) + return + } + + sessionID, err := c.Cookie(controller.runtime.SessionCookieName) + if err != nil && !errors.Is(err, http.ErrNoCookie) { + controller.log.App.Error().Err(err).Msg("Error retrieving session cookie on OIDC logout") + c.JSON(http.StatusInternalServerError, gin.H{ + "status": http.StatusInternalServerError, + "message": "Internal Server Error", + }) + return + } + + callbackTicket := controller.auth.CreateLogoutCallbackTicket(redirectURI) + result, err := controller.auth.Logout(c, service.LogoutRequest{ + SessionID: sessionID, + UserContext: userContext, + ClientIP: c.ClientIP(), + RedirectURI: redirectURI, + ProviderCallbackURL: controller.runtime.AppURL + "/api/user/logout/callback", + ProviderState: callbackTicket, + }) + if err != nil { + controller.log.App.Error().Err(err).Msg("Error deleting session on OIDC logout") + c.JSON(http.StatusInternalServerError, gin.H{ + "status": http.StatusInternalServerError, + "message": "Internal Server Error", + }) + return + } + if result.Cookie != nil { + http.SetCookie(c.Writer, result.Cookie) + } + + c.Redirect(http.StatusFound, result.RedirectURL) +} + // This endpoint does **not** return a code, it handles param validation, ticket creation // and then redirects to the frontend to handle the consent screen. It performs no destructive // actions (like logging out an existing session) diff --git a/internal/controller/oidc_controller_logout_test.go b/internal/controller/oidc_controller_logout_test.go new file mode 100644 index 00000000..38e22a3a --- /dev/null +++ b/internal/controller/oidc_controller_logout_test.go @@ -0,0 +1,254 @@ +package controller + +import ( + "context" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/steveiliop56/ding" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/tinyauthapp/tinyauth/internal/model" + "github.com/tinyauthapp/tinyauth/internal/repository" + "github.com/tinyauthapp/tinyauth/internal/repository/memory" + "github.com/tinyauthapp/tinyauth/internal/service" + testutil "github.com/tinyauthapp/tinyauth/internal/test" + "github.com/tinyauthapp/tinyauth/internal/utils/logger" +) + +func TestOIDCEndSession(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + description string + method string + configureRequest func(values url.Values, idToken string) + provider *model.OAuthServiceConfig + expectedStatus int + expectedRedirect string + expectedCallback string + expectSessionGone bool + validateRedirect func(t *testing.T, location string) + }{ + { + description: "GET cascades a valid end-session request to the upstream provider", + method: http.MethodGet, + configureRequest: func(values url.Values, idToken string) { + values.Set("id_token_hint", idToken) + values.Set("client_id", "some-client-id") + values.Set("post_logout_redirect_uri", "https://rp.example.net/logged-out") + values.Set("state", "state-123") + }, + provider: &model.OAuthServiceConfig{ + ClientID: "upstream-client", + LogoutURL: "https://id.example.com/api/oidc/end-session", + }, + expectedStatus: http.StatusFound, + expectedCallback: "https://rp.example.net/logged-out?state=state-123", + expectSessionGone: true, + validateRedirect: func(t *testing.T, location string) { + parsed, err := url.Parse(location) + require.NoError(t, err) + assert.Equal(t, "id.example.com", parsed.Host) + assert.Equal(t, "upstream-id-token", parsed.Query().Get("id_token_hint")) + assert.Equal(t, "upstream-client", parsed.Query().Get("client_id")) + assert.NotEmpty(t, parsed.Query().Get("state")) + assert.NotContains(t, parsed.Query().Get("state"), "rp.example.net") + }, + }, + { + description: "POST redirects directly to a registered post-logout URI", + method: http.MethodPost, + configureRequest: func(values url.Values, idToken string) { + values.Set("id_token_hint", idToken) + values.Set("post_logout_redirect_uri", "https://rp.example.net/logged-out") + values.Set("state", "state-123") + }, + expectedStatus: http.StatusFound, + expectedRedirect: "https://rp.example.net/logged-out?state=state-123", + expectSessionGone: true, + }, + { + description: "Rejects an unregistered post-logout redirect URI", + method: http.MethodGet, + configureRequest: func(values url.Values, idToken string) { + values.Set("id_token_hint", idToken) + values.Set("post_logout_redirect_uri", "https://evil.example.com/logged-out") + }, + expectedStatus: http.StatusBadRequest, + }, + { + description: "Requires confirmation when the ID token hint is missing", + method: http.MethodGet, + configureRequest: func(values url.Values, _ string) { + values.Set("client_id", "some-client-id") + values.Set("post_logout_redirect_uri", "https://rp.example.net/logged-out") + }, + expectedStatus: http.StatusFound, + expectedRedirect: "https://tinyauth.example.com/logout", + }, + { + description: "Requires confirmation when client ID does not match the token audience", + method: http.MethodGet, + configureRequest: func(values url.Values, idToken string) { + values.Set("id_token_hint", idToken) + values.Set("client_id", "another-client") + values.Set("post_logout_redirect_uri", "https://rp.example.net/logged-out") + }, + expectedStatus: http.StatusFound, + expectedRedirect: "https://tinyauth.example.com/logout", + }, + } + + for _, test := range tests { + t.Run(test.description, func(t *testing.T) { + cfg, runtime := testutil.CreateTestConfigs(t) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + log := logger.NewLogger().WithTestConfig() + log.Init() + store := memory.New() + dg := ding.New(ctx) + + oidcService, err := service.NewOIDCService(service.OIDCServiceInput{ + Log: log, + Config: &cfg, + Runtime: &runtime, + Queries: store, + Ding: dg, + }) + require.NoError(t, err) + + policyEngine, err := service.NewPolicyEngine(service.PolicyEngineInput{ + Log: log, + Config: &cfg, + }) + require.NoError(t, err) + broker := service.NewOAuthBrokerService(service.OAuthBrokerServiceInput{ + Log: log, + Runtime: &runtime, + Ctx: ctx, + }) + authService, err := service.NewAuthService(service.AuthServiceInput{ + Log: log, + Config: &cfg, + Runtime: &runtime, + Ctx: ctx, + Ding: dg, + Queries: store, + OAuthBroker: broker, + PolicyEngine: policyEngine, + }) + require.NoError(t, err) + + userContext := newOAuthUserContext("pocketid", "upstream-id-token") + client, ok := oidcService.GetClient("some-client-id") + require.True(t, ok) + sub := oidcService.CreateSub(*userContext, client.ClientID) + tokenResponse, err := oidcService.GenerateAccessToken(ctx, client, service.AuthorizeCodeEntry{ + Scope: "openid", + Userinfo: service.UserinfoResponse{ + Sub: sub, + Email: userContext.GetEmail(), + PreferredUsername: userContext.GetUsername(), + }, + }, time.Now().Unix()) + require.NoError(t, err) + + _, err = store.CreateSession(ctx, repository.CreateSessionParams{ + UUID: "browser-session", + Username: userContext.GetUsername(), + Email: userContext.GetEmail(), + Name: userContext.GetName(), + Provider: "pocketid", + Expiry: time.Now().Add(time.Hour).Unix(), + CreatedAt: time.Now().Unix(), + OAuthName: "Pocket ID", + OAuthSub: "upstream-sub", + OAuthIDToken: "upstream-id-token", + }) + require.NoError(t, err) + + if test.provider != nil { + runtime.OAuthProviders = map[string]model.OAuthServiceConfig{ + "pocketid": *test.provider, + } + } + + router := gin.New() + router.Use(func(c *gin.Context) { + c.Set("context", userContext) + c.Next() + }) + NewOIDCController(OIDCControllerInput{ + Log: log, + OIDCService: oidcService, + AuthService: authService, + RuntimeConfig: &runtime, + RouterGroup: router.Group("/api"), + MainRouter: &router.RouterGroup, + }) + NewUserController(UserControllerInput{ + Log: log, + StaticConfig: &cfg, + RuntimeConfig: &runtime, + RouterGroup: router.Group("/api"), + AuthService: authService, + }) + + values := url.Values{} + test.configureRequest(values, tokenResponse.IDToken) + requestURL := "/api/oidc/end-session" + var request *http.Request + if test.method == http.MethodPost { + request = httptest.NewRequest(test.method, requestURL, strings.NewReader(values.Encode())) + request.Header.Set("Content-Type", "application/x-www-form-urlencoded") + } else { + request = httptest.NewRequest(test.method, requestURL+"?"+values.Encode(), nil) + } + request.AddCookie(&http.Cookie{Name: runtime.SessionCookieName, Value: "browser-session"}) + + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, request) + + assert.Equal(t, test.expectedStatus, recorder.Code) + if test.expectedRedirect != "" { + assert.Equal(t, test.expectedRedirect, recorder.Header().Get("Location")) + } + if test.validateRedirect != nil { + test.validateRedirect(t, recorder.Header().Get("Location")) + } + if test.expectedCallback != "" { + providerRedirect, err := url.Parse(recorder.Header().Get("Location")) + require.NoError(t, err) + callbackRecorder := httptest.NewRecorder() + callbackRequest := httptest.NewRequest( + http.MethodGet, + "/api/user/logout/callback?state="+url.QueryEscape(providerRedirect.Query().Get("state")), + nil, + ) + router.ServeHTTP(callbackRecorder, callbackRequest) + assert.Equal(t, http.StatusFound, callbackRecorder.Code) + assert.Equal(t, test.expectedCallback, callbackRecorder.Header().Get("Location")) + + callbackRecorder = httptest.NewRecorder() + router.ServeHTTP(callbackRecorder, callbackRequest) + assert.Equal(t, http.StatusFound, callbackRecorder.Code) + assert.Equal(t, runtime.AppURL, callbackRecorder.Header().Get("Location")) + } + + _, err = store.GetSession(ctx, "browser-session") + if test.expectSessionGone { + require.ErrorIs(t, err, repository.ErrNotFound) + } else { + require.NoError(t, err) + } + }) + } +} diff --git a/internal/controller/user_controller.go b/internal/controller/user_controller.go index 65b12de0..82678cd9 100644 --- a/internal/controller/user_controller.go +++ b/internal/controller/user_controller.go @@ -4,6 +4,8 @@ import ( "errors" "fmt" "net/http" + "net/url" + "strings" "time" "github.com/tinyauthapp/tinyauth/internal/model" @@ -11,6 +13,7 @@ import ( "github.com/tinyauthapp/tinyauth/internal/service" "github.com/tinyauthapp/tinyauth/internal/utils" "github.com/tinyauthapp/tinyauth/internal/utils/logger" + "github.com/tinyauthapp/tinyauth/pkg/validators" "go.uber.org/dig" "github.com/gin-gonic/gin" @@ -28,6 +31,7 @@ type TotpRequest struct { type UserController struct { log *logger.Logger + config *model.Config runtime *model.RuntimeConfig auth *service.AuthService } @@ -36,6 +40,7 @@ type UserControllerInput struct { dig.In Log *logger.Logger + StaticConfig *model.Config RuntimeConfig *model.RuntimeConfig RouterGroup *gin.RouterGroup `name:"apiRouterGroup"` AuthService *service.AuthService @@ -44,6 +49,7 @@ type UserControllerInput struct { func NewUserController(i UserControllerInput) *UserController { controller := &UserController{ log: i.Log, + config: i.StaticConfig, runtime: i.RuntimeConfig, auth: i.AuthService, } @@ -51,6 +57,7 @@ func NewUserController(i UserControllerInput) *UserController { userGroup := i.RouterGroup.Group("/user") userGroup.POST("/login", controller.loginHandler) userGroup.POST("/logout", controller.logoutHandler) + userGroup.GET("/logout/callback", controller.ssoLogoutCallbackHandler) userGroup.POST("/totp", controller.totpHandler) userGroup.POST("/tailscale", controller.tailscaleHandler) @@ -227,51 +234,119 @@ func (controller *UserController) loginHandler(c *gin.Context) { func (controller *UserController) logoutHandler(c *gin.Context) { controller.log.App.Debug().Msg("Logout attempt") - uuid, err := c.Cookie(controller.runtime.SessionCookieName) + // redirect_uri is a Tinyauth UI/navigation parameter. It is not an + // OpenID Connect RP-Initiated Logout parameter. The standardized OP-facing + // parameters are added later when compiling the provider logout request. + requestedRedirectURI := "" + if c.Query("login_for") == "app" { + requestedRedirectURI = c.Query("redirect_uri") + } + redirectURI := controller.safeLogoutRedirect(requestedRedirectURI) + userContext, err := new(model.UserContext).NewFromGin(c) if err != nil { - if errors.Is(err, http.ErrNoCookie) { - controller.log.App.Warn().Msg("Logout attempt without session cookie, treating as successful logout") - c.JSON(200, gin.H{ - "status": 200, - "message": "Logout successful", - }) - return - } + userContext = nil + } + + sessionID, err := c.Cookie(controller.runtime.SessionCookieName) + if err != nil && !errors.Is(err, http.ErrNoCookie) { controller.log.App.Error().Err(err).Msg("Error retrieving session cookie on logout") - c.JSON(500, gin.H{ - "status": 500, + c.JSON(http.StatusInternalServerError, gin.H{ + "status": http.StatusInternalServerError, "message": "Internal Server Error", }) return } - cookie, err := controller.auth.DeleteSession(c, uuid) - + result, err := controller.auth.Logout(c, service.LogoutRequest{ + SessionID: sessionID, + UserContext: userContext, + ClientIP: c.ClientIP(), + RedirectURI: redirectURI, + ProviderCallbackURL: controller.runtime.AppURL + "/api/user/logout/callback", + ProviderState: redirectURI, + }) if err != nil { controller.log.App.Error().Err(err).Msg("Error deleting session on logout") - c.JSON(500, gin.H{ - "status": 500, + c.JSON(http.StatusInternalServerError, gin.H{ + "status": http.StatusInternalServerError, "message": "Internal Server Error", }) return } + if result.Cookie != nil { + http.SetCookie(c.Writer, result.Cookie) + } - context, err := new(model.UserContext).NewFromGin(c) + response := gin.H{ + "status": http.StatusOK, + "message": "Logout successful", + } + if result.ProviderLogout || requestedRedirectURI != "" { + response["redirectUrl"] = result.RedirectURL + } - if err == nil { - controller.log.AuditLogout(context.GetUsername(), context.GetProviderID(), c.ClientIP()) - } else { - controller.log.App.Warn().Err(err).Msg("Failed to get user context during logout, logging audit with unknown user") - controller.log.AuditLogout("unknown", "unknown", c.ClientIP()) + c.JSON(http.StatusOK, response) +} + +func (controller *UserController) ssoLogoutCallbackHandler(c *gin.Context) { + if redirectURI, ok := controller.auth.ConsumeLogoutCallbackTicket(c.Query("state")); ok { + c.Redirect(http.StatusFound, redirectURI) + return } - http.SetCookie(c.Writer, cookie) + // state is defined by OpenID Connect RP-Initiated Logout 1.0 as an opaque + // RP value that the OP returns unchanged after logout. We use it to carry + // the already-validated Tinyauth application return URI across the OP hop. + redirectURI := controller.safeLogoutRedirect(c.Query("state")) + c.Redirect(http.StatusFound, redirectURI) +} - c.JSON(200, gin.H{ - "status": 200, - "message": "Logout successful", +func (controller *UserController) safeLogoutRedirect(raw string) string { + fallback := controller.runtime.AppURL + if raw == "" { + return fallback + } + + appURL, err := url.Parse(controller.runtime.AppURL) + if err != nil { + return fallback + } + + allowedSchemes := []string{"http", "https"} + if appURL.Scheme == "https" { + allowedSchemes = []string{"https"} + } + + schemeValidator := validators.NewDomainValidator(validators.DomainValidatorOptions{ + WithScheme: true, + AllowedSchemes: allowedSchemes, }) + hostname, err := schemeValidator.SafeHostname(raw) + if err != nil { + return fallback + } + + domainValidator := validators.NewDomainValidator(validators.DomainValidatorOptions{ + WithPort: true, + }) + err = domainValidator.Validate(raw, controller.runtime.AppURL) + if err == nil { + return raw + } + + if !errors.Is(err, validators.ErrHostnameMismatch) || + controller.config == nil || + !controller.config.Auth.SubdomainsEnabled { + return fallback + } + + cookieDomain := strings.ToLower(controller.runtime.CookieDomain) + if hostname == cookieDomain || strings.HasSuffix(hostname, "."+cookieDomain) { + return raw + } + + return fallback } func (controller *UserController) totpHandler(c *gin.Context) { diff --git a/internal/controller/user_controller_sso_logout_test.go b/internal/controller/user_controller_sso_logout_test.go new file mode 100644 index 00000000..c7c9e1cc --- /dev/null +++ b/internal/controller/user_controller_sso_logout_test.go @@ -0,0 +1,183 @@ +package controller + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "net/url" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/steveiliop56/ding" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/tinyauthapp/tinyauth/internal/model" + "github.com/tinyauthapp/tinyauth/internal/repository" + "github.com/tinyauthapp/tinyauth/internal/repository/memory" + "github.com/tinyauthapp/tinyauth/internal/service" + "github.com/tinyauthapp/tinyauth/internal/test" + "github.com/tinyauthapp/tinyauth/internal/utils/logger" +) + +type logoutResponse struct { + RedirectURL string `json:"redirectUrl"` +} + +func TestSSOLogout(t *testing.T) { + gin.SetMode(gin.TestMode) + cfg, runtime := test.CreateTestConfigs(t) + + tests := []struct { + description string + provider model.OAuthServiceConfig + session *repository.CreateSessionParams + userContext *model.UserContext + requestPath string + validate func(t *testing.T, recorder *httptest.ResponseRecorder) + }{ + { + description: "Uses context ID token for provider logout", + provider: model.OAuthServiceConfig{ + ClientID: "client-id", + LogoutURL: "https://id.example.com/api/oidc/end-session", + }, + session: &repository.CreateSessionParams{ + UUID: "oauth-session", + Username: "user@example.com", + Email: "user@example.com", + Name: "Test User", + Provider: "pocketid", + OAuthGroups: "admins", + Expiry: time.Now().Add(time.Hour).Unix(), + CreatedAt: time.Now().Unix(), + OAuthName: "Pocket ID", + OAuthSub: "sub-123", + OAuthIDToken: "id-token", + }, + userContext: newOAuthUserContext("pocketid", "id-token"), + requestPath: "/api/user/logout?login_for=app&redirect_uri=https://app.example.com/", + validate: func(t *testing.T, recorder *httptest.ResponseRecorder) { + require.Equal(t, http.StatusOK, recorder.Code) + require.Len(t, recorder.Result().Cookies(), 1) + assert.Equal(t, "tinyauth-session", recorder.Result().Cookies()[0].Name) + + response := parseLogoutResponse(t, recorder) + require.NotEmpty(t, response.RedirectURL) + + parsed, err := url.Parse(response.RedirectURL) + require.NoError(t, err) + assert.Equal(t, "https", parsed.Scheme) + assert.Equal(t, "id.example.com", parsed.Host) + assert.Equal(t, "id-token", parsed.Query().Get("id_token_hint")) + assert.Equal(t, "client-id", parsed.Query().Get("client_id")) + assert.Equal(t, "https://app.example.com/", parsed.Query().Get("state")) + assert.Equal( + t, + "https://tinyauth.example.com/api/user/logout/callback", + parsed.Query().Get("post_logout_redirect_uri"), + ) + }, + }, + { + description: "Falls back to app redirect when provider logout URL is invalid", + provider: model.OAuthServiceConfig{ + LogoutURL: "http://id.example.com/api/oidc/end-session", + }, + userContext: newOAuthUserContext("pocketid", ""), + requestPath: "/api/user/logout?login_for=app&redirect_uri=https://app.example.com/", + validate: func(t *testing.T, recorder *httptest.ResponseRecorder) { + require.Equal(t, http.StatusOK, recorder.Code) + + response := parseLogoutResponse(t, recorder) + assert.Equal(t, "https://app.example.com/", response.RedirectURL) + }, + }, + } + + for _, test := range tests { + t.Run(test.description, func(t *testing.T) { + log := logger.NewLogger().WithTestConfig() + log.Init() + + runtime.OAuthProviders = map[string]model.OAuthServiceConfig{ + "pocketid": test.provider, + } + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + store := memory.New() + if test.session != nil { + _, err := store.CreateSession(ctx, *test.session) + require.NoError(t, err) + } + + dg := ding.New(ctx) + authService, err := service.NewAuthService(service.AuthServiceInput{ + Log: log, + Config: &cfg, + Runtime: &runtime, + Ctx: ctx, + Ding: dg, + Queries: store, + }) + require.NoError(t, err) + + router := gin.New() + if test.userContext != nil { + router.Use(func(c *gin.Context) { + c.Set("context", test.userContext) + c.Next() + }) + } + + NewUserController(UserControllerInput{ + Log: log, + StaticConfig: &cfg, + RuntimeConfig: &runtime, + RouterGroup: router.Group("/api"), + AuthService: authService, + }) + + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, test.requestPath, nil) + if test.session != nil { + req.AddCookie(&http.Cookie{ + Name: runtime.SessionCookieName, + Value: test.session.UUID, + }) + } + + router.ServeHTTP(recorder, req) + + test.validate(t, recorder) + }) + } +} + +func newOAuthUserContext(providerID, idToken string) *model.UserContext { + return &model.UserContext{ + Authenticated: true, + Provider: model.ProviderOAuth, + OAuth: &model.OAuthContext{ + BaseContext: model.BaseContext{ + Username: "user@example.com", + Name: "Test User", + Email: "user@example.com", + }, + DisplayName: "Pocket ID", + ID: providerID, + IDToken: idToken, + }, + } +} + +func parseLogoutResponse(t *testing.T, recorder *httptest.ResponseRecorder) logoutResponse { + t.Helper() + + var response logoutResponse + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response)) + return response +} diff --git a/internal/controller/user_controller_test.go b/internal/controller/user_controller_test.go index a971a5d8..962f227f 100644 --- a/internal/controller/user_controller_test.go +++ b/internal/controller/user_controller_test.go @@ -577,6 +577,7 @@ func TestUserController(t *testing.T) { NewUserController(UserControllerInput{ Log: log, + StaticConfig: &cfg, RuntimeConfig: &runtime, RouterGroup: group, AuthService: authService, diff --git a/internal/controller/well_known_controller.go b/internal/controller/well_known_controller.go index a32c3a06..30df9ba4 100644 --- a/internal/controller/well_known_controller.go +++ b/internal/controller/well_known_controller.go @@ -29,6 +29,7 @@ type OpenIDConnectConfiguration struct { AuthorizationEndpoint string `json:"authorization_endpoint"` TokenEndpoint string `json:"token_endpoint"` UserinfoEndpoint string `json:"userinfo_endpoint"` + EndSessionEndpoint string `json:"end_session_endpoint"` JwksUri string `json:"jwks_uri"` ScopesSupported []string `json:"scopes_supported"` ResponseTypesSupported []string `json:"response_types_supported"` @@ -80,6 +81,7 @@ func (controller *WellKnownController) OpenIDConnectConfiguration(c *gin.Context AuthorizationEndpoint: fmt.Sprintf("%s/authorize", issuer), TokenEndpoint: fmt.Sprintf("%s/api/oidc/token", issuer), UserinfoEndpoint: fmt.Sprintf("%s/api/oidc/userinfo", issuer), + EndSessionEndpoint: fmt.Sprintf("%s/api/oidc/end-session", issuer), JwksUri: fmt.Sprintf("%s/.well-known/jwks.json", issuer), ScopesSupported: service.SupportedScopes, ResponseTypesSupported: service.SupportedResponseTypes, diff --git a/internal/controller/well_known_controller_test.go b/internal/controller/well_known_controller_test.go index 8a969667..281e9bb4 100644 --- a/internal/controller/well_known_controller_test.go +++ b/internal/controller/well_known_controller_test.go @@ -49,6 +49,7 @@ func TestWellKnownController(t *testing.T) { AuthorizationEndpoint: fmt.Sprintf("%s/authorize", runtime.AppURL), TokenEndpoint: fmt.Sprintf("%s/api/oidc/token", runtime.AppURL), UserinfoEndpoint: fmt.Sprintf("%s/api/oidc/userinfo", runtime.AppURL), + EndSessionEndpoint: fmt.Sprintf("%s/api/oidc/end-session", runtime.AppURL), JwksUri: fmt.Sprintf("%s/.well-known/jwks.json", runtime.AppURL), ScopesSupported: service.SupportedScopes, ResponseTypesSupported: service.SupportedResponseTypes, diff --git a/internal/model/config.go b/internal/model/config.go index 642c00a9..fc807e80 100644 --- a/internal/model/config.go +++ b/internal/model/config.go @@ -253,19 +253,23 @@ type TailscaleConfig struct { // OAuth/OIDC config type OAuthServiceConfig struct { - ClientID string `description:"OAuth client ID." yaml:"clientId,omitempty"` - ClientSecret string `description:"OAuth client secret." yaml:"clientSecret,omitempty"` - ClientSecretFile string `description:"Path to the file containing the OAuth client secret." yaml:"clientSecretFile,omitempty"` - Whitelist []string `description:"Comma-separated list of allowed OAuth domains for this provider." yaml:"whitelist,omitempty"` - WhitelistFile string `description:"Path to the OAuth whitelist file for this provider." yaml:"whitelistFile,omitempty"` - Scopes []string `description:"OAuth scopes." yaml:"scopes,omitempty"` - RedirectURL string `description:"OAuth redirect URL." yaml:"redirectUrl,omitempty"` - AuthURL string `description:"OAuth authorization URL." yaml:"authUrl,omitempty"` - TokenURL string `description:"OAuth token URL." yaml:"tokenUrl,omitempty"` - UserinfoURL string `description:"OAuth userinfo URL." yaml:"userinfoUrl,omitempty"` - Insecure bool `description:"Allow insecure OAuth connections." yaml:"insecure,omitempty"` - Name string `description:"Provider name in UI." yaml:"name,omitempty"` - Claims OAuthServiceClaimsMap `description:"Map of claims to extract from the userinfo response." yaml:"claims,omitempty"` + ClientID string `description:"OAuth client ID." yaml:"clientId,omitempty"` + ClientSecret string `description:"OAuth client secret." yaml:"clientSecret,omitempty"` + ClientSecretFile string `description:"Path to the file containing the OAuth client secret." yaml:"clientSecretFile,omitempty"` + Whitelist []string `description:"Comma-separated list of allowed OAuth domains for this provider." yaml:"whitelist,omitempty"` + WhitelistFile string `description:"Path to the OAuth whitelist file for this provider." yaml:"whitelistFile,omitempty"` + Scopes []string `description:"OAuth scopes." yaml:"scopes,omitempty"` + RedirectURL string `description:"OAuth redirect URL." yaml:"redirectUrl,omitempty"` + AuthURL string `description:"OAuth authorization URL." yaml:"authUrl,omitempty"` + TokenURL string `description:"OAuth token URL." yaml:"tokenUrl,omitempty"` + UserinfoURL string `description:"OAuth userinfo URL." yaml:"userinfoUrl,omitempty"` + // LogoutURL is the OpenID Provider end_session_endpoint used for + // OpenID Connect RP-Initiated Logout 1.0: + // https://openid.net/specs/openid-connect-rpinitiated-1_0-final.html#RPLogout + LogoutURL string `description:"OpenID Connect RP-Initiated Logout end_session_endpoint URL." yaml:"logoutUrl,omitempty"` + Insecure bool `description:"Allow insecure OAuth connections." yaml:"insecure,omitempty"` + Name string `description:"Provider name in UI." yaml:"name,omitempty"` + Claims OAuthServiceClaimsMap `description:"Map of claims to extract from the userinfo response." yaml:"claims,omitempty"` } type OAuthServiceClaimsMap struct { @@ -276,12 +280,13 @@ type OAuthServiceClaimsMap struct { } type OIDCClientConfig struct { - ID string `description:"OIDC client ID." yaml:"-"` - ClientID string `description:"OIDC client ID." yaml:"clientId,omitempty"` - ClientSecret string `description:"OIDC client secret." yaml:"clientSecret,omitempty"` - ClientSecretFile string `description:"Path to the file containing the OIDC client secret." yaml:"clientSecretFile,omitempty"` - TrustedRedirectURIs []string `description:"List of trusted redirect URIs." yaml:"trustedRedirectUris,omitempty"` - Name string `description:"Client name in UI." yaml:"name,omitempty"` + ID string `description:"OIDC client ID." yaml:"-"` + ClientID string `description:"OIDC client ID." yaml:"clientId,omitempty"` + ClientSecret string `description:"OIDC client secret." yaml:"clientSecret,omitempty"` + ClientSecretFile string `description:"Path to the file containing the OIDC client secret." yaml:"clientSecretFile,omitempty"` + TrustedRedirectURIs []string `description:"List of trusted redirect URIs." yaml:"trustedRedirectUris,omitempty"` + TrustedPostLogoutRedirectURIs []string `description:"List of trusted post-logout redirect URIs." yaml:"trustedPostLogoutRedirectUris,omitempty"` + Name string `description:"Client name in UI." yaml:"name,omitempty"` } type ACLsConfig struct { diff --git a/internal/model/context.go b/internal/model/context.go index 03d76941..0574b24c 100644 --- a/internal/model/context.go +++ b/internal/model/context.go @@ -48,6 +48,7 @@ type OAuthContext struct { BaseContext Groups []string Sub string + IDToken string DisplayName string ID string } @@ -159,6 +160,7 @@ func (c *UserContext) NewFromSession(session *repository.Session) (*UserContext, return strings.Split(session.OAuthGroups, ",") }(), Sub: session.OAuthSub, + IDToken: session.OAuthIDToken, DisplayName: session.OAuthName, ID: session.Provider, } diff --git a/internal/model/context_test.go b/internal/model/context_test.go index ab9da7cf..c90b7056 100644 --- a/internal/model/context_test.go +++ b/internal/model/context_test.go @@ -98,12 +98,12 @@ func TestContext(t *testing.T) { run: func(t *testing.T, c *UserContext) any { got, err := c.NewFromSession(&repository.Session{ Username: "dave", Provider: "github", - OAuthGroups: "devs,admins", OAuthSub: "sub-123", OAuthName: "GitHub", + OAuthGroups: "devs,admins", OAuthSub: "sub-123", OAuthIDToken: "id-token", OAuthName: "GitHub", }) require.NoError(t, err) - return [5]any{got.Provider, got.OAuth.ID, got.OAuth.Sub, got.OAuth.DisplayName, got.OAuth.Groups} + return [6]any{got.Provider, got.OAuth.ID, got.OAuth.Sub, got.OAuth.IDToken, got.OAuth.DisplayName, got.OAuth.Groups} }, - expected: [5]any{ProviderOAuth, "github", "sub-123", "GitHub", []string{"devs", "admins"}}, + expected: [6]any{ProviderOAuth, "github", "sub-123", "id-token", "GitHub", []string{"devs", "admins"}}, }, { description: "Local getters return BaseContext fields", diff --git a/internal/repository/memory/session_queries.go b/internal/repository/memory/session_queries.go index 2edde6b1..6c12d1e4 100644 --- a/internal/repository/memory/session_queries.go +++ b/internal/repository/memory/session_queries.go @@ -40,6 +40,7 @@ func (s *Store) UpdateSession(_ context.Context, arg repository.UpdateSessionPar sess.Expiry = arg.Expiry sess.OAuthName = arg.OAuthName sess.OAuthSub = arg.OAuthSub + sess.OAuthIDToken = arg.OAuthIDToken s.sessions[arg.UUID] = sess return sess, nil } diff --git a/internal/repository/models.go b/internal/repository/models.go index 9e356680..df57ecb1 100644 --- a/internal/repository/models.go +++ b/internal/repository/models.go @@ -4,17 +4,18 @@ package repository // sqlc-generated driver packages use these via the conversion layer in their store.go. type Session struct { - UUID string - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - CreatedAt int64 - OAuthName string - OAuthSub string + UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + CreatedAt int64 + OAuthName string + OAuthSub string + OAuthIDToken string } type OidcSession struct { @@ -30,30 +31,32 @@ type OidcSession struct { } type CreateSessionParams struct { - UUID string - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - CreatedAt int64 - OAuthName string - OAuthSub string + UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + CreatedAt int64 + OAuthName string + OAuthSub string + OAuthIDToken string } type UpdateSessionParams struct { - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - OAuthName string - OAuthSub string - UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + OAuthName string + OAuthSub string + OAuthIDToken string + UUID string } type CreateOIDCSessionParams struct { diff --git a/internal/repository/postgres/models.go b/internal/repository/postgres/models.go index ccf7ce62..50f47e07 100644 --- a/internal/repository/postgres/models.go +++ b/internal/repository/postgres/models.go @@ -24,15 +24,16 @@ type OidcSession struct { } type Session struct { - UUID string - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - CreatedAt int64 - OAuthName string - OAuthSub string + UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + CreatedAt int64 + OAuthName string + OAuthSub string + OAuthIDToken string } diff --git a/internal/repository/postgres/session_queries.sql.go b/internal/repository/postgres/session_queries.sql.go index c7ea71d4..77ff97ee 100644 --- a/internal/repository/postgres/session_queries.sql.go +++ b/internal/repository/postgres/session_queries.sql.go @@ -21,25 +21,27 @@ INSERT INTO "sessions" ( "expiry", "created_at", "oauth_name", - "oauth_sub" + "oauth_sub", + "oauth_id_token" ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11 + $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12 ) -RETURNING uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub +RETURNING uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub, oauth_id_token ` type CreateSessionParams struct { - UUID string - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - CreatedAt int64 - OAuthName string - OAuthSub string + UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + CreatedAt int64 + OAuthName string + OAuthSub string + OAuthIDToken string } func (q *Queries) CreateSession(ctx context.Context, arg CreateSessionParams) (Session, error) { @@ -55,6 +57,7 @@ func (q *Queries) CreateSession(ctx context.Context, arg CreateSessionParams) (S arg.CreatedAt, arg.OAuthName, arg.OAuthSub, + arg.OAuthIDToken, ) var i Session err := row.Scan( @@ -69,6 +72,7 @@ func (q *Queries) CreateSession(ctx context.Context, arg CreateSessionParams) (S &i.CreatedAt, &i.OAuthName, &i.OAuthSub, + &i.OAuthIDToken, ) return i, err } @@ -94,7 +98,7 @@ func (q *Queries) DeleteSession(ctx context.Context, uuid string) error { } const getSession = `-- name: GetSession :one -SELECT uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub FROM "sessions" +SELECT uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub, oauth_id_token FROM "sessions" WHERE "uuid" = $1 ` @@ -113,6 +117,7 @@ func (q *Queries) GetSession(ctx context.Context, uuid string) (Session, error) &i.CreatedAt, &i.OAuthName, &i.OAuthSub, + &i.OAuthIDToken, ) return i, err } @@ -127,22 +132,24 @@ UPDATE "sessions" SET "oauth_groups" = $6, "expiry" = $7, "oauth_name" = $8, - "oauth_sub" = $9 -WHERE "uuid" = $10 -RETURNING uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub + "oauth_sub" = $9, + "oauth_id_token" = $10 +WHERE "uuid" = $11 +RETURNING uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub, oauth_id_token ` type UpdateSessionParams struct { - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - OAuthName string - OAuthSub string - UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + OAuthName string + OAuthSub string + OAuthIDToken string + UUID string } func (q *Queries) UpdateSession(ctx context.Context, arg UpdateSessionParams) (Session, error) { @@ -156,6 +163,7 @@ func (q *Queries) UpdateSession(ctx context.Context, arg UpdateSessionParams) (S arg.Expiry, arg.OAuthName, arg.OAuthSub, + arg.OAuthIDToken, arg.UUID, ) var i Session @@ -171,6 +179,7 @@ func (q *Queries) UpdateSession(ctx context.Context, arg UpdateSessionParams) (S &i.CreatedAt, &i.OAuthName, &i.OAuthSub, + &i.OAuthIDToken, ) return i, err } diff --git a/internal/repository/sqlite/models.go b/internal/repository/sqlite/models.go index f30ae672..22697c5b 100644 --- a/internal/repository/sqlite/models.go +++ b/internal/repository/sqlite/models.go @@ -24,15 +24,16 @@ type OidcSession struct { } type Session struct { - UUID string - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - CreatedAt int64 - OAuthName string - OAuthSub string + UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + CreatedAt int64 + OAuthName string + OAuthSub string + OAuthIDToken string } diff --git a/internal/repository/sqlite/session_queries.sql.go b/internal/repository/sqlite/session_queries.sql.go index 7792fc4b..8e9537f9 100644 --- a/internal/repository/sqlite/session_queries.sql.go +++ b/internal/repository/sqlite/session_queries.sql.go @@ -21,25 +21,27 @@ INSERT INTO "sessions" ( "expiry", "created_at", "oauth_name", - "oauth_sub" + "oauth_sub", + "oauth_id_token" ) VALUES ( - ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ? + ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ? ) -RETURNING uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub +RETURNING uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub, oauth_id_token ` type CreateSessionParams struct { - UUID string - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - CreatedAt int64 - OAuthName string - OAuthSub string + UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + CreatedAt int64 + OAuthName string + OAuthSub string + OAuthIDToken string } func (q *Queries) CreateSession(ctx context.Context, arg CreateSessionParams) (Session, error) { @@ -55,6 +57,7 @@ func (q *Queries) CreateSession(ctx context.Context, arg CreateSessionParams) (S arg.CreatedAt, arg.OAuthName, arg.OAuthSub, + arg.OAuthIDToken, ) var i Session err := row.Scan( @@ -69,6 +72,7 @@ func (q *Queries) CreateSession(ctx context.Context, arg CreateSessionParams) (S &i.CreatedAt, &i.OAuthName, &i.OAuthSub, + &i.OAuthIDToken, ) return i, err } @@ -94,7 +98,7 @@ func (q *Queries) DeleteSession(ctx context.Context, uuid string) error { } const getSession = `-- name: GetSession :one -SELECT uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub FROM "sessions" +SELECT uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub, oauth_id_token FROM "sessions" WHERE "uuid" = ? ` @@ -113,6 +117,7 @@ func (q *Queries) GetSession(ctx context.Context, uuid string) (Session, error) &i.CreatedAt, &i.OAuthName, &i.OAuthSub, + &i.OAuthIDToken, ) return i, err } @@ -127,22 +132,24 @@ UPDATE "sessions" SET "oauth_groups" = ?, "expiry" = ?, "oauth_name" = ?, - "oauth_sub" = ? + "oauth_sub" = ?, + "oauth_id_token" = ? WHERE "uuid" = ? -RETURNING uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub +RETURNING uuid, username, email, name, provider, totp_pending, oauth_groups, expiry, created_at, oauth_name, oauth_sub, oauth_id_token ` type UpdateSessionParams struct { - Username string - Email string - Name string - Provider string - TotpPending bool - OAuthGroups string - Expiry int64 - OAuthName string - OAuthSub string - UUID string + Username string + Email string + Name string + Provider string + TotpPending bool + OAuthGroups string + Expiry int64 + OAuthName string + OAuthSub string + OAuthIDToken string + UUID string } func (q *Queries) UpdateSession(ctx context.Context, arg UpdateSessionParams) (Session, error) { @@ -156,6 +163,7 @@ func (q *Queries) UpdateSession(ctx context.Context, arg UpdateSessionParams) (S arg.Expiry, arg.OAuthName, arg.OAuthSub, + arg.OAuthIDToken, arg.UUID, ) var i Session @@ -171,6 +179,7 @@ func (q *Queries) UpdateSession(ctx context.Context, arg UpdateSessionParams) (S &i.CreatedAt, &i.OAuthName, &i.OAuthSub, + &i.OAuthIDToken, ) return i, err } diff --git a/internal/service/auth_service.go b/internal/service/auth_service.go index 0b503e9c..444207d7 100644 --- a/internal/service/auth_service.go +++ b/internal/service/auth_service.go @@ -7,6 +7,7 @@ import ( "fmt" "math/big" "net/http" + "net/url" "strings" "time" @@ -57,6 +58,25 @@ type LoginAttempt struct { LockedUntil time.Time } +// LogoutRequest contains the session and redirect information needed to end a +// local session and optionally cascade logout to its OAuth provider. +type LogoutRequest struct { + SessionID string + UserContext *model.UserContext + ClientIP string + RedirectURI string + ProviderCallbackURL string + ProviderState string +} + +// LogoutResponse contains the local session cookie and the next browser +// location selected by the logout flow. +type LogoutResponse struct { + Cookie *http.Cookie + RedirectURL string + ProviderLogout bool +} + type AuthService struct { log *logger.Logger config *model.Config @@ -72,9 +92,10 @@ type AuthService struct { dummyHash string caches struct { - login *cache.CacheStore[LoginAttempt] - oauth *cache.CacheStore[OAuthPendingSession] - ldap *cache.CacheStore[[]string] + login *cache.CacheStore[LoginAttempt] + oauth *cache.CacheStore[OAuthPendingSession] + ldap *cache.CacheStore[[]string] + logoutCallback *cache.CacheStore[string] } } @@ -119,10 +140,12 @@ func NewAuthService(i AuthServiceInput) (*AuthService, error) { oauthCache := cache.NewCacheStore[OAuthPendingSession](256) loginCache := cache.NewCacheStore[LoginAttempt](service.calculateLockdownLimit()) ldapCache := cache.NewCacheStore[[]string](1024) + logoutCallbackCache := cache.NewCacheStore[string](256) service.caches.oauth = oauthCache service.caches.login = loginCache service.caches.ldap = ldapCache + service.caches.logoutCallback = logoutCallbackCache i.Ding.Go(func(ctx context.Context) { ticker := time.NewTicker(1 * time.Minute) @@ -134,6 +157,7 @@ func NewAuthService(i AuthServiceInput) (*AuthService, error) { service.caches.oauth.Sweep() service.caches.login.Sweep() service.caches.ldap.Sweep() + service.caches.logoutCallback.Sweep() case <-ctx.Done(): return } @@ -363,17 +387,18 @@ func (auth *AuthService) CreateSession(ctx context.Context, data repository.Sess expiresAt := time.Now().Add(time.Duration(expiry) * time.Second) session := repository.CreateSessionParams{ - UUID: u.String(), - Username: data.Username, - Email: data.Email, - Name: data.Name, - Provider: data.Provider, - TotpPending: data.TotpPending, - OAuthGroups: data.OAuthGroups, - Expiry: expiresAt.Unix(), - CreatedAt: time.Now().Unix(), - OAuthName: data.OAuthName, - OAuthSub: data.OAuthSub, + UUID: u.String(), + Username: data.Username, + Email: data.Email, + Name: data.Name, + Provider: data.Provider, + TotpPending: data.TotpPending, + OAuthGroups: data.OAuthGroups, + Expiry: expiresAt.Unix(), + CreatedAt: time.Now().Unix(), + OAuthName: data.OAuthName, + OAuthSub: data.OAuthSub, + OAuthIDToken: data.OAuthIDToken, } _, err = auth.queries.CreateSession(ctx, session) @@ -419,16 +444,17 @@ func (auth *AuthService) RefreshSession(ctx context.Context, uuid string) (*http newExpiry := session.Expiry + refreshThreshold _, err = auth.queries.UpdateSession(ctx, repository.UpdateSessionParams{ - Username: session.Username, - Email: session.Email, - Name: session.Name, - Provider: session.Provider, - TotpPending: session.TotpPending, - OAuthGroups: session.OAuthGroups, - Expiry: newExpiry, - OAuthName: session.OAuthName, - OAuthSub: session.OAuthSub, - UUID: session.UUID, + Username: session.Username, + Email: session.Email, + Name: session.Name, + Provider: session.Provider, + TotpPending: session.TotpPending, + OAuthGroups: session.OAuthGroups, + Expiry: newExpiry, + OAuthName: session.OAuthName, + OAuthSub: session.OAuthSub, + OAuthIDToken: session.OAuthIDToken, + UUID: session.UUID, }) if err != nil { @@ -469,6 +495,104 @@ func (auth *AuthService) DeleteSession(ctx context.Context, uuid string) (*http. }, nil } +// Logout deletes the local session and compiles an upstream provider logout +// request when the session originated from a configured OAuth provider. +func (auth *AuthService) Logout(ctx context.Context, req LogoutRequest) (*LogoutResponse, error) { + providerID := "" + idToken := "" + if req.UserContext != nil && req.UserContext.IsOAuth() { + providerID = req.UserContext.OAuth.ID + idToken = req.UserContext.OAuth.IDToken + } + + response := &LogoutResponse{ + RedirectURL: req.RedirectURI, + } + if req.SessionID != "" { + cookie, err := auth.DeleteSession(ctx, req.SessionID) + if err != nil { + return nil, err + } + response.Cookie = cookie + + if req.UserContext != nil { + auth.log.AuditLogout(req.UserContext.GetUsername(), req.UserContext.GetProviderID(), req.ClientIP) + } else { + auth.log.App.Warn().Msg("Failed to get user context during logout, logging audit with unknown user") + auth.log.AuditLogout("unknown", "unknown", req.ClientIP) + } + } else { + auth.log.App.Warn().Msg("Logout attempt without session cookie, treating as successful logout") + } + + provider, ok := auth.runtime.OAuthProviders[providerID] + if !ok || provider.LogoutURL == "" { + return response, nil + } + + logoutURL, err := buildOAuthLogoutURL( + provider, + req.ProviderCallbackURL, + idToken, + req.ProviderState, + ) + if err != nil { + auth.log.App.Warn().Err(err).Str("provider", providerID).Msg("Invalid OAuth logout URL, skipping provider logout") + return response, nil + } + + response.RedirectURL = logoutURL + response.ProviderLogout = true + return response, nil +} + +// CreateLogoutCallbackTicket stores a previously validated redirect URI under +// an opaque, short-lived identifier for the upstream provider callback. +func (auth *AuthService) CreateLogoutCallbackTicket(redirectURI string) string { + ticket := utils.GenerateString(32) + auth.caches.logoutCallback.Set(ticket, redirectURI, 10*time.Minute) + return ticket +} + +// ConsumeLogoutCallbackTicket returns and removes a redirect URI associated +// with an opaque callback ticket. +func (auth *AuthService) ConsumeLogoutCallbackTicket(ticket string) (string, bool) { + redirectURI := "" + found := false + auth.caches.logoutCallback.WithLock(func(actions cache.CacheStoreActions[string]) { + redirectURI, found = actions.Get(ticket) + if found { + actions.Delete(ticket) + } + }) + return redirectURI, found +} + +func buildOAuthLogoutURL(provider model.OAuthServiceConfig, callbackURL, idToken, state string) (string, error) { + logoutURL, err := url.Parse(provider.LogoutURL) + if err != nil || logoutURL.Host == "" { + return "", fmt.Errorf("invalid logout URL") + } + if logoutURL.Scheme != "https" { + return "", fmt.Errorf("unsupported logout URL scheme") + } + + query := logoutURL.Query() + if provider.ClientID != "" { + query.Set("client_id", provider.ClientID) + } + if idToken != "" { + query.Set("id_token_hint", idToken) + } + query.Set("post_logout_redirect_uri", callbackURL) + if state != "" { + query.Set("state", state) + } + logoutURL.RawQuery = query.Encode() + + return logoutURL.String(), nil +} + func (auth *AuthService) GetSession(ctx context.Context, uuid string) (*repository.Session, error) { session, err := auth.queries.GetSession(ctx, uuid) diff --git a/internal/service/oidc_service.go b/internal/service/oidc_service.go index 3848a6c9..184c6d51 100644 --- a/internal/service/oidc_service.go +++ b/internal/service/oidc_service.go @@ -39,11 +39,13 @@ var ( ) var ( - ErrCodeExpired = errors.New("code_expired") - ErrCodeNotFound = errors.New("code_not_found") - ErrTokenNotFound = errors.New("token_not_found") - ErrTokenExpired = errors.New("token_expired") - ErrInvalidClient = errors.New("invalid_client") + ErrCodeExpired = errors.New("code_expired") + ErrCodeNotFound = errors.New("code_not_found") + ErrTokenNotFound = errors.New("token_not_found") + ErrTokenExpired = errors.New("token_expired") + ErrInvalidClient = errors.New("invalid_client") + ErrEndSessionConfirmationNeeded = errors.New("end_session_confirmation_needed") + ErrInvalidPostLogoutRedirectURI = errors.New("invalid_post_logout_redirect_uri") ) type OIDCPrompt string @@ -133,6 +135,16 @@ type AuthorizeRequest struct { MaxAge string `form:"max_age" json:"max_age" url:"max_age"` } +// EndSessionRequest contains the parameters defined by OpenID Connect +// RP-Initiated Logout 1.0. +type EndSessionRequest struct { + IDTokenHint string `form:"id_token_hint"` + LogoutHint string `form:"logout_hint"` + ClientID string `form:"client_id"` + PostLogoutRedirectURI string `form:"post_logout_redirect_uri"` + State string `form:"state"` +} + type AuthorizeCodeEntry struct { CodeHash string Scope string @@ -430,6 +442,71 @@ func (service *OIDCService) ValidateAuthorizeParams(req AuthorizeRequest) error return nil } +// ValidateEndSessionRequest validates an ID token hint and compiles the trusted +// location to which the user agent should be redirected after logout. +func (service *OIDCService) ValidateEndSessionRequest(ctx context.Context, req EndSessionRequest, userContext *model.UserContext) (string, error) { + if req.IDTokenHint == "" { + return "", ErrEndSessionConfirmationNeeded + } + + object, err := jose.ParseSigned(req.IDTokenHint, []jose.SignatureAlgorithm{jose.RS256}) + if err != nil { + return "", ErrEndSessionConfirmationNeeded + } + + payload, err := object.Verify(service.publicKey) + if err != nil { + return "", ErrEndSessionConfirmationNeeded + } + + claims := ClaimSet{} + if err = json.Unmarshal(payload, &claims); err != nil { + return "", ErrEndSessionConfirmationNeeded + } + + if claims.Iss != service.issuer || claims.Aud == "" || claims.Sub == "" { + return "", ErrEndSessionConfirmationNeeded + } + if req.ClientID != "" && req.ClientID != claims.Aud { + return "", ErrEndSessionConfirmationNeeded + } + + client, ok := service.GetClient(claims.Aud) + if !ok { + return "", ErrEndSessionConfirmationNeeded + } + + session, err := service.queries.GetOIDCSessionBySub(ctx, claims.Sub) + if err != nil || session.ClientID != client.ClientID { + return "", ErrEndSessionConfirmationNeeded + } + + if userContext != nil && userContext.IsAuthenticated() && + claims.Sub != service.CreateSub(*userContext, client.ClientID) { + return "", ErrEndSessionConfirmationNeeded + } + + if req.PostLogoutRedirectURI == "" { + return service.issuer, nil + } + if !slices.Contains(client.TrustedPostLogoutRedirectURIs, req.PostLogoutRedirectURI) { + return "", ErrInvalidPostLogoutRedirectURI + } + if req.State == "" { + return req.PostLogoutRedirectURI, nil + } + + parsedRedirectURI, err := url.Parse(req.PostLogoutRedirectURI) + if err != nil { + return "", ErrInvalidPostLogoutRedirectURI + } + query := parsedRedirectURI.Query() + query.Set("state", req.State) + parsedRedirectURI.RawQuery = query.Encode() + + return parsedRedirectURI.String(), nil +} + func (service *OIDCService) filterScopes(scopes []string) []string { return utils.Filter(scopes, func(scope string) bool { return slices.Contains(SupportedScopes, scope) diff --git a/internal/test/test.go b/internal/test/test.go index 031d5e1f..223e1dff 100644 --- a/internal/test/test.go +++ b/internal/test/test.go @@ -27,10 +27,11 @@ func CreateTestConfigs(t *testing.T) (model.Config, model.RuntimeConfig) { OIDC: model.OIDCConfig{ Clients: map[string]model.OIDCClientConfig{ "test": { - ClientID: "some-client-id", - ClientSecret: "some-client-secret", - TrustedRedirectURIs: []string{"https://test.example.com/callback"}, - Name: "Test Client", + ClientID: "some-client-id", + ClientSecret: "some-client-secret", + TrustedRedirectURIs: []string{"https://test.example.com/callback"}, + TrustedPostLogoutRedirectURIs: []string{"https://rp.example.net/logged-out"}, + Name: "Test Client", }, }, PrivateKeyPath: filepath.Join(tempDir, "key.pem"), diff --git a/pkg/validators/domain_validator.go b/pkg/validators/domain_validator.go index d41d82b3..7e1365b0 100644 --- a/pkg/validators/domain_validator.go +++ b/pkg/validators/domain_validator.go @@ -87,6 +87,10 @@ func (v *DomainValidator) getURL(i string) (*url.URL, error) { return nil, fmt.Errorf("missing host or scheme in url: %s", i) } + if u.User != nil { + return nil, fmt.Errorf("userinfo is not supported") + } + return u, nil } @@ -110,6 +114,10 @@ func (v *DomainValidator) getURL(i string) (*url.URL, error) { return nil, fmt.Errorf("missing host in url: %s", i) } + if u.User != nil { + return nil, fmt.Errorf("userinfo is not supported") + } + return u, nil } diff --git a/pkg/validators/domain_validator_test.go b/pkg/validators/domain_validator_test.go index aa7587f1..89eb0bad 100644 --- a/pkg/validators/domain_validator_test.go +++ b/pkg/validators/domain_validator_test.go @@ -119,6 +119,13 @@ func TestDomainValidator_SafeHostname(t *testing.T) { input: "example.com", expected: "example.com", }, + { + description: "URL with userinfo should fail", + input: "https://evil.example.net@example.com", + errorFunc: func(t *testing.T, e error) { + assert.ErrorContains(t, e, "userinfo is not supported") + }, + }, } for _, test := range tests { @@ -254,6 +261,14 @@ func TestDomainValidator_Validate(t *testing.T) { assert.ErrorIs(t, e, ErrHostnameMismatch) }, }, + { + description: "Hostname ending with expected domain but not subdomain should fail", + expected: "example.com", + actual: "badexample.com", + errorFunc: func(t *testing.T, e error) { + assert.ErrorIs(t, e, ErrHostnameMismatch) + }, + }, } for _, test := range tests { diff --git a/sql/postgres/session_queries.sql b/sql/postgres/session_queries.sql index 22aecd46..79e71de0 100644 --- a/sql/postgres/session_queries.sql +++ b/sql/postgres/session_queries.sql @@ -10,9 +10,10 @@ INSERT INTO "sessions" ( "expiry", "created_at", "oauth_name", - "oauth_sub" + "oauth_sub", + "oauth_id_token" ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11 + $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12 ) RETURNING *; @@ -34,8 +35,9 @@ UPDATE "sessions" SET "oauth_groups" = $6, "expiry" = $7, "oauth_name" = $8, - "oauth_sub" = $9 -WHERE "uuid" = $10 + "oauth_sub" = $9, + "oauth_id_token" = $10 +WHERE "uuid" = $11 RETURNING *; -- name: DeleteExpiredSessions :exec diff --git a/sql/postgres/session_schemas.sql b/sql/postgres/session_schemas.sql index 925bcd74..cc294e4f 100644 --- a/sql/postgres/session_schemas.sql +++ b/sql/postgres/session_schemas.sql @@ -9,5 +9,6 @@ CREATE TABLE IF NOT EXISTS "sessions" ( "expiry" BIGINT NOT NULL, "created_at" BIGINT NOT NULL, "oauth_name" TEXT NOT NULL DEFAULT '', - "oauth_sub" TEXT NOT NULL DEFAULT '' + "oauth_sub" TEXT NOT NULL DEFAULT '', + "oauth_id_token" TEXT NOT NULL DEFAULT '' ); diff --git a/sql/sqlite/session_queries.sql b/sql/sqlite/session_queries.sql index da93126e..bea0c8a8 100644 --- a/sql/sqlite/session_queries.sql +++ b/sql/sqlite/session_queries.sql @@ -10,9 +10,10 @@ INSERT INTO "sessions" ( "expiry", "created_at", "oauth_name", - "oauth_sub" + "oauth_sub", + "oauth_id_token" ) VALUES ( - ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ? + ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ? ) RETURNING *; @@ -34,7 +35,8 @@ UPDATE "sessions" SET "oauth_groups" = ?, "expiry" = ?, "oauth_name" = ?, - "oauth_sub" = ? + "oauth_sub" = ?, + "oauth_id_token" = ? WHERE "uuid" = ? RETURNING *; diff --git a/sql/sqlite/session_schemas.sql b/sql/sqlite/session_schemas.sql index a7f37eb7..61583813 100644 --- a/sql/sqlite/session_schemas.sql +++ b/sql/sqlite/session_schemas.sql @@ -9,5 +9,6 @@ CREATE TABLE IF NOT EXISTS "sessions" ( "expiry" INTEGER NOT NULL, "created_at" INTEGER NOT NULL, "oauth_name" TEXT NULL, - "oauth_sub" TEXT NULL + "oauth_sub" TEXT NULL, + "oauth_id_token" TEXT NOT NULL DEFAULT '' ); diff --git a/sqlc.yml b/sqlc.yml index e4f98a25..b13d8da3 100644 --- a/sqlc.yml +++ b/sqlc.yml @@ -12,6 +12,7 @@ sql: oauth_groups: "OAuthGroups" oauth_name: "OAuthName" oauth_sub: "OAuthSub" + oauth_id_token: "OAuthIDToken" redirect_uri: "RedirectURI" overrides: - column: "sessions.oauth_groups" @@ -20,6 +21,8 @@ sql: go_type: "string" - column: "sessions.oauth_sub" go_type: "string" + - column: "sessions.oauth_id_token" + go_type: "string" - column: "sessions.ldap_groups" go_type: "string" - column: "oidc_sessions.nonce" @@ -36,4 +39,5 @@ sql: oauth_groups: "OAuthGroups" oauth_name: "OAuthName" oauth_sub: "OAuthSub" + oauth_id_token: "OAuthIDToken" redirect_uri: "RedirectURI"