diff --git a/.env.example b/.env.example index baf02c49..c554e9fa 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. 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..1411ea20 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,13 +126,24 @@ 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; + if (typeof redirectUrl === "string" && redirectUrl.length > 0) { + window.location.replace(redirectUrl); + return; + } + redirectTimer.current = window.setTimeout(() => { window.location.replace(`/login${compiledParams}`); }, 500); diff --git a/frontend/src/pages/logout-page.tsx b/frontend/src/pages/logout-page.tsx index 78ef0555..c535ad36 100644 --- a/frontend/src/pages/logout-page.tsx +++ b/frontend/src/pages/logout-page.tsx @@ -36,15 +36,30 @@ 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; + if (typeof redirectUrl === "string" && redirectUrl.length > 0) { + window.location.replace(redirectUrl); + return; + } + redirectTimer.current = window.setTimeout(() => { window.location.replace(`/login${compiledParams}`); }, 500); 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/user_controller.go b/internal/controller/user_controller.go index 65b12de0..573fdfb3 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,168 @@ 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 by buildOAuthLogoutURL. + 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", + userContext = nil + } + + providerID := "" + idToken := "" + if userContext != nil && userContext.IsOAuth() { + providerID = userContext.OAuth.ID + idToken = userContext.OAuth.IDToken + } + + uuid, err := c.Cookie(controller.runtime.SessionCookieName) + if err == nil { + cookie, err := controller.auth.DeleteSession(c, uuid) + if err != nil { + controller.log.App.Error().Err(err).Msg("Error deleting session on logout") + c.JSON(http.StatusInternalServerError, gin.H{ + "status": http.StatusInternalServerError, + "message": "Internal Server Error", }) return } + + http.SetCookie(c.Writer, cookie) + + if userContext != nil { + controller.log.AuditLogout(userContext.GetUsername(), userContext.GetProviderID(), c.ClientIP()) + } else { + controller.log.App.Warn().Msg("Failed to get user context during logout, logging audit with unknown user") + controller.log.AuditLogout("unknown", "unknown", c.ClientIP()) + } + } else if errors.Is(err, http.ErrNoCookie) { + controller.log.App.Warn().Msg("Logout attempt without session cookie, treating as successful logout") + } else { 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) + response := gin.H{ + "status": http.StatusOK, + "message": "Logout successful", + } + + provider, ok := controller.runtime.OAuthProviders[providerID] + if ok && provider.LogoutURL != "" { + // OpenID Connect RP-Initiated Logout 1.0: + // https://openid.net/specs/openid-connect-rpinitiated-1_0-final.html#RPLogout + // + // OP-facing standardized parameters: + // id_token_hint + // post_logout_redirect_uri + // state + callbackURL := controller.runtime.AppURL + "/api/user/logout/callback" + logoutURL, buildErr := buildOAuthLogoutURL(provider, callbackURL, idToken, redirectURI) + if buildErr != nil { + controller.log.App.Warn().Err(buildErr).Str("provider", providerID).Msg("Invalid OAuth logout URL, skipping provider logout") + if requestedRedirectURI != "" { + response["redirectUrl"] = redirectURI + } + } else { + response["redirectUrl"] = logoutURL + } + } else if requestedRedirectURI != "" { + // Non-OIDC/local logout can still return to the validated application. + response["redirectUrl"] = redirectURI + } + + c.JSON(http.StatusOK, response) +} + +func (controller *UserController) ssoLogoutCallbackHandler(c *gin.Context) { + // 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) +} +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 { - controller.log.App.Error().Err(err).Msg("Error deleting session on logout") - c.JSON(500, gin.H{ - "status": 500, - "message": "Internal Server Error", - }) - return + return fallback } - context, err := new(model.UserContext).NewFromGin(c) + 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 { - 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()) + return raw } - http.SetCookie(c.Writer, cookie) + if !errors.Is(err, validators.ErrHostnameMismatch) || + controller.config == nil || + !controller.config.Auth.SubdomainsEnabled { + return fallback + } - c.JSON(200, gin.H{ - "status": 200, - "message": "Logout successful", - }) + cookieDomain := strings.ToLower(controller.runtime.CookieDomain) + if hostname == cookieDomain || strings.HasSuffix(hostname, "."+cookieDomain) { + return raw + } + + return fallback +} + +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 (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..4e965d55 --- /dev/null +++ b/internal/controller/user_controller_sso_logout_test.go @@ -0,0 +1,241 @@ +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) + + 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 _, tc := range tests { + t.Run(tc.description, func(t *testing.T) { + log := logger.NewLogger().WithTestConfig() + log.Init() + + cfg, runtime := test.CreateTestConfigs(t) + runtime.OAuthProviders = map[string]model.OAuthServiceConfig{ + "pocketid": tc.provider, + } + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + store := memory.New() + var authService *service.AuthService + if tc.session != nil { + _, err := store.CreateSession(ctx, *tc.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 tc.userContext != nil { + router.Use(func(c *gin.Context) { + c.Set("context", tc.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, tc.requestPath, nil) + if tc.session != nil { + req.AddCookie(&http.Cookie{ + Name: runtime.SessionCookieName, + Value: tc.session.UUID, + }) + } + + router.ServeHTTP(recorder, req) + + tc.validate(t, recorder) + }) + } +} + +func TestBuildOAuthLogoutURL(t *testing.T) { + tests := []struct { + description string + provider model.OAuthServiceConfig + expectError bool + validate func(t *testing.T, parsed *url.URL) + }{ + { + description: "Builds provider logout URL with OIDC logout parameters", + provider: model.OAuthServiceConfig{ + ClientID: "client-id", + LogoutURL: "https://id.example.com/api/oidc/end-session", + }, + validate: func(t *testing.T, parsed *url.URL) { + assert.Equal(t, "https", parsed.Scheme) + assert.Equal(t, "id.example.com", parsed.Host) + assert.Equal(t, "/api/oidc/end-session", parsed.Path) + assert.Equal(t, "client-id", parsed.Query().Get("client_id")) + assert.Equal(t, "id-token", parsed.Query().Get("id_token_hint")) + assert.Equal(t, "https://app.example.com/", parsed.Query().Get("state")) + assert.Equal( + t, + "https://auth.example.com/api/user/logout/callback", + parsed.Query().Get("post_logout_redirect_uri"), + ) + }, + }, + { + description: "Rejects HTTP provider logout URL", + provider: model.OAuthServiceConfig{ + LogoutURL: "http://id.example.com/api/oidc/end-session", + }, + expectError: true, + }, + } + + for _, tc := range tests { + t.Run(tc.description, func(t *testing.T) { + got, err := buildOAuthLogoutURL( + tc.provider, + "https://auth.example.com/api/user/logout/callback", + "id-token", + "https://app.example.com/", + ) + if tc.expectError { + require.Error(t, err) + return + } + + require.NoError(t, err) + parsed, err := url.Parse(got) + require.NoError(t, err) + tc.validate(t, parsed) + }) + } +} + +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/model/config.go b/internal/model/config.go index 642c00a9..1ecf064d 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 { 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..c37278bd 100644 --- a/internal/service/auth_service.go +++ b/internal/service/auth_service.go @@ -363,17 +363,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 +420,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 { 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"