From 1c8a9346ef076b4bbf9f04458dc579f7ba6ad922 Mon Sep 17 00:00:00 2001 From: rcourtman Date: Tue, 7 Jul 2026 09:54:14 +0100 Subject: [PATCH] Fix legacy OIDC SSO discovery and CSP nonce Refs #1533 --- internal/api/frontend_embed.go | 94 +++++++++ internal/api/frontend_embed_test.go | 33 +++ internal/api/identity_sso_handlers.go | 16 ++ internal/api/oidc_handlers.go | 8 +- .../api/oidc_handlers_callback_url_test.go | 16 ++ .../api/security_status_additional_test.go | 44 ++++ internal/config/config_load_test.go | 20 +- internal/config/oidc.go | 198 ++++++++++++++++++ internal/config/persistence.go | 60 +++++- .../config/persistence_sso_coverage_test.go | 41 +++- internal/config/sso.go | 32 +++ 11 files changed, 546 insertions(+), 16 deletions(-) diff --git a/internal/api/frontend_embed.go b/internal/api/frontend_embed.go index 63aa4280b..84efd3f6b 100644 --- a/internal/api/frontend_embed.go +++ b/internal/api/frontend_embed.go @@ -3,6 +3,7 @@ package api import ( "bytes" "embed" + "html" "io" "io/fs" "net/http" @@ -23,6 +24,7 @@ var cspNoncePlaceholder = []byte("__CSP_NONCE__") func serveIndexWithNonce(w http.ResponseWriter, r *http.Request, content []byte) { if nonce := CSPNonceFromContext(r.Context()); nonce != "" { content = bytes.ReplaceAll(content, cspNoncePlaceholder, []byte(nonce)) + content = addNonceToInlineHTMLTags(content, nonce) } w.Header().Set("Content-Type", "text/html; charset=utf-8") w.Header().Set("Cache-Control", "no-cache, no-store, must-revalidate") @@ -31,6 +33,98 @@ func serveIndexWithNonce(w http.ResponseWriter, r *http.Request, content []byte) w.Write(content) } +func addNonceToInlineHTMLTags(content []byte, nonce string) []byte { + if len(content) == 0 || nonce == "" { + return content + } + nonceAttr := []byte(` nonce="` + html.EscapeString(nonce) + `"`) + content = addNonceToInlineHTMLTag(content, "script", "src", nonceAttr) + content = addNonceToInlineHTMLTag(content, "style", "", nonceAttr) + return content +} + +func addNonceToInlineHTMLTag(content []byte, tagName string, skipAttr string, nonceAttr []byte) []byte { + lowerContent := bytes.ToLower(content) + needle := []byte("<" + strings.ToLower(tagName)) + searchAt := 0 + copiedUntil := 0 + var out []byte + + for searchAt < len(content) { + relativeStart := bytes.Index(lowerContent[searchAt:], needle) + if relativeStart == -1 { + break + } + start := searchAt + relativeStart + tagNameEnd := start + len(needle) + if tagNameEnd < len(content) && !isHTMLTagBoundary(lowerContent[tagNameEnd]) { + searchAt = tagNameEnd + continue + } + + relativeEnd := bytes.IndexByte(content[tagNameEnd:], '>') + if relativeEnd == -1 { + break + } + end := tagNameEnd + relativeEnd + tag := lowerContent[start : end+1] + if htmlStartTagHasAttribute(tag, "nonce") || (skipAttr != "" && htmlStartTagHasAttribute(tag, skipAttr)) { + searchAt = end + 1 + continue + } + + if out == nil { + out = make([]byte, 0, len(content)+len(nonceAttr)) + } + out = append(out, content[copiedUntil:end]...) + out = append(out, nonceAttr...) + out = append(out, content[end]) + copiedUntil = end + 1 + searchAt = end + 1 + } + + if out == nil { + return content + } + out = append(out, content[copiedUntil:]...) + return out +} + +func htmlStartTagHasAttribute(tag []byte, attr string) bool { + if len(tag) == 0 || attr == "" { + return false + } + attrBytes := []byte(strings.ToLower(attr)) + searchAt := 0 + for searchAt < len(tag) { + idx := bytes.Index(tag[searchAt:], attrBytes) + if idx == -1 { + return false + } + start := searchAt + idx + end := start + len(attrBytes) + beforeOK := start == 0 || isHTMLAttributeBoundaryBefore(tag[start-1]) + afterOK := end >= len(tag) || isHTMLAttributeBoundaryAfter(tag[end]) + if beforeOK && afterOK { + return true + } + searchAt = end + } + return false +} + +func isHTMLTagBoundary(b byte) bool { + return b == ' ' || b == '\t' || b == '\n' || b == '\r' || b == '/' || b == '>' +} + +func isHTMLAttributeBoundaryBefore(b byte) bool { + return b == '<' || b == '/' || b == ' ' || b == '\t' || b == '\n' || b == '\r' +} + +func isHTMLAttributeBoundaryAfter(b byte) bool { + return b == '=' || b == '/' || b == '>' || b == ' ' || b == '\t' || b == '\n' || b == '\r' +} + // Embed the entire frontend dist directory // //go:embed all:frontend-modern/dist diff --git a/internal/api/frontend_embed_test.go b/internal/api/frontend_embed_test.go index b3b875d77..e3634a480 100644 --- a/internal/api/frontend_embed_test.go +++ b/internal/api/frontend_embed_test.go @@ -1,6 +1,7 @@ package api import ( + "context" "io" "net/http" "net/http/httptest" @@ -135,3 +136,35 @@ func TestServeFrontendHandler_StaticAndSPA(t *testing.T) { t.Fatalf("api route status = %d", rec.Code) } } + +func TestServeIndexWithNonceAddsNonceToGeneratedInlineTags(t *testing.T) { + content := []byte(`` + + `` + + `` + + `` + + `` + + ``) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + req = req.WithContext(context.WithValue(req.Context(), cspNonceKey{}, "test-nonce")) + rec := httptest.NewRecorder() + + serveIndexWithNonce(rec, req, content) + + body := rec.Body.String() + if strings.Contains(body, "__CSP_NONCE__") { + t.Fatalf("expected nonce placeholder to be replaced, got body %q", body) + } + if !strings.Contains(body, ``) { + t.Fatalf("expected external script to remain without nonce, got body %q", body) + } + if got := strings.Count(body, `nonce="test-nonce"`); got != 3 { + t.Fatalf("nonce count = %d, want 3 in body %q", got, body) + } +} diff --git a/internal/api/identity_sso_handlers.go b/internal/api/identity_sso_handlers.go index 0fa585730..9359bfa78 100644 --- a/internal/api/identity_sso_handlers.go +++ b/internal/api/identity_sso_handlers.go @@ -248,6 +248,14 @@ func (r *Router) ensureSSOConfig() *config.SSOConfig { } } + publicURL := "" + if r.config != nil { + publicURL = r.config.PublicURL + } + if config.ApplyLegacyOIDCEnvProvider(r.ssoConfig, publicURL) { + log.Info().Str("provider_id", config.LegacyOIDCProviderID).Msg("Loaded legacy OIDC environment SSO provider") + } + setSSOAuthSnapshot(r.config, r.ssoConfig) return r.ssoConfig } @@ -421,6 +429,10 @@ func (r *Router) handleUpdateSSOProvider(w http.ResponseWriter, req *http.Reques writeErrorResponse(w, http.StatusNotFound, "not_found", "Provider not found", nil) return } + if existing.RuntimeManaged { + writeErrorResponse(w, http.StatusConflict, "provider_managed_by_environment", "Provider is managed by OIDC environment variables", nil) + return + } body, err := io.ReadAll(io.LimitReader(req.Body, maxRequestBodySize)) if err != nil { @@ -581,6 +593,10 @@ func (r *Router) handleDeleteSSOProvider(w http.ResponseWriter, req *http.Reques writeErrorResponse(w, http.StatusNotFound, "not_found", "Provider not found", nil) return } + if existing.RuntimeManaged { + writeErrorResponse(w, http.StatusConflict, "provider_managed_by_environment", "Provider is managed by OIDC environment variables", nil) + return + } // Remove provider if err := r.ssoConfig.RemoveProvider(providerID); err != nil { diff --git a/internal/api/oidc_handlers.go b/internal/api/oidc_handlers.go index ed8602fd6..9ca939cfd 100644 --- a/internal/api/oidc_handlers.go +++ b/internal/api/oidc_handlers.go @@ -234,13 +234,19 @@ func ssoProviderToOIDCConfig(provider *config.SSOProvider, redirectURL string) * } // extractOIDCProviderID extracts the provider ID from an OIDC endpoint path. -// Expected paths: /api/oidc/{providerID}/login, /api/oidc/{providerID}/callback +// Expected paths: /api/oidc/{providerID}/login, /api/oidc/{providerID}/callback. +// Legacy v5 paths /api/oidc/login and /api/oidc/callback map to the migrated +// legacy provider ID. func extractOIDCProviderID(urlPath, endpoint string) string { parts := strings.Split(strings.TrimPrefix(urlPath, "/"), "/") // parts: ["api", "oidc", "{providerID}", "{endpoint}"] if len(parts) >= 4 && parts[0] == "api" && parts[1] == "oidc" && parts[3] == endpoint { return parts[2] } + // parts: ["api", "oidc", "{endpoint}"] + if len(parts) == 3 && parts[0] == "api" && parts[1] == "oidc" && parts[2] == endpoint { + return config.LegacyOIDCProviderID + } return "" } diff --git a/internal/api/oidc_handlers_callback_url_test.go b/internal/api/oidc_handlers_callback_url_test.go index 9f0c1f3ca..31078c3d0 100644 --- a/internal/api/oidc_handlers_callback_url_test.go +++ b/internal/api/oidc_handlers_callback_url_test.go @@ -3,6 +3,8 @@ package api import ( "net/http/httptest" "testing" + + "github.com/rcourtman/pulse-go-rewrite/internal/config" ) func TestBuildSSOOIDCCallbackURL(t *testing.T) { @@ -28,3 +30,17 @@ func TestBuildSSOOIDCCallbackURL(t *testing.T) { } }) } + +func TestExtractOIDCProviderIDSupportsLegacyPaths(t *testing.T) { + t.Parallel() + + if got := extractOIDCProviderID("/api/oidc/legacy-oidc/login", "login"); got != "legacy-oidc" { + t.Fatalf("provider scoped login id = %q, want legacy-oidc", got) + } + if got := extractOIDCProviderID("/api/oidc/callback", "callback"); got != config.LegacyOIDCProviderID { + t.Fatalf("legacy callback id = %q, want %s", got, config.LegacyOIDCProviderID) + } + if got := extractOIDCProviderID("/api/oidc/login", "login"); got != config.LegacyOIDCProviderID { + t.Fatalf("legacy login id = %q, want %s", got, config.LegacyOIDCProviderID) + } +} diff --git a/internal/api/security_status_additional_test.go b/internal/api/security_status_additional_test.go index d37fc4cc0..63f1d6899 100644 --- a/internal/api/security_status_additional_test.go +++ b/internal/api/security_status_additional_test.go @@ -132,6 +132,50 @@ func TestSecurityStatusExposesPersistedSSOProvider(t *testing.T) { } } +func TestSecurityStatusExposesLegacyOIDCEnvProvider(t *testing.T) { + t.Setenv("OIDC_ENABLED", "true") + t.Setenv("OIDC_ISSUER_URL", "https://id.example.test") + t.Setenv("OIDC_CLIENT_ID", "pulse-client") + t.Setenv("OIDC_CLIENT_SECRET", "secret") + + cfg := newTestConfigWithTokens(t) + cfg.PublicURL = "https://pulse.example.test" + router := NewRouter(cfg, nil, nil, nil, nil, "1.0.0") + + req := httptest.NewRequest(http.MethodGet, "/api/security/status", nil) + rec := httptest.NewRecorder() + router.Handler().ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("expected 200 for security status, got %d", rec.Code) + } + + var payload map[string]interface{} + if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { + t.Fatalf("decode response: %v", err) + } + + rawProviders, ok := payload["ssoProviders"].([]interface{}) + if !ok || len(rawProviders) != 1 { + t.Fatalf("expected one legacy OIDC provider in security status, got %#v", payload["ssoProviders"]) + } + + firstProvider, ok := rawProviders[0].(map[string]interface{}) + if !ok { + t.Fatalf("expected provider object, got %#v", rawProviders[0]) + } + + if got := firstProvider["id"]; got != config.LegacyOIDCProviderID { + t.Fatalf("provider id = %v, want %s", got, config.LegacyOIDCProviderID) + } + if got := firstProvider["displayName"]; got != "Single Sign-On" { + t.Fatalf("provider displayName = %v, want Single Sign-On", got) + } + if got := firstProvider["loginUrl"]; got != "/api/oidc/legacy-oidc/login" { + t.Fatalf("provider loginUrl = %v, want /api/oidc/legacy-oidc/login", got) + } +} + func TestSecurityStatusExposesSettingsCapabilitiesForScopedToken(t *testing.T) { prevAuthorizer := auth.GetAuthorizer() auth.SetAuthorizer(&allowRulesAuthorizer{ diff --git a/internal/config/config_load_test.go b/internal/config/config_load_test.go index 3203e14cc..a7ff7dd0d 100644 --- a/internal/config/config_load_test.go +++ b/internal/config/config_load_test.go @@ -345,16 +345,28 @@ func TestLoad_ProxyAuth(t *testing.T) { assert.Equal(t, "X-User", cfg.ProxyAuthUserHeader) } -func TestLoad_OIDCEnvIgnored(t *testing.T) { +func TestLegacyOIDCEnvProvider(t *testing.T) { t.Setenv("PULSE_DATA_DIR", t.TempDir()) t.Setenv("OIDC_ENABLED", "true") t.Setenv("OIDC_ISSUER_URL", "https://issuer.com") t.Setenv("OIDC_CLIENT_ID", "client-id") t.Setenv("OIDC_CLIENT_SECRET", "client-secret") + t.Setenv("OIDC_ALLOWED_GROUPS", "admins, operators") + t.Setenv("OIDC_GROUP_ROLE_MAPPINGS", "admins=admin,operators=viewer") - cfg, err := Load() - require.NoError(t, err) - require.NotNil(t, cfg) + provider, ok := LegacyOIDCEnvProvider("https://pulse.example.com") + require.True(t, ok) + require.NotNil(t, provider) + require.NotNil(t, provider.OIDC) + assert.Equal(t, LegacyOIDCProviderID, provider.ID) + assert.True(t, provider.RuntimeManaged) + assert.Equal(t, "https://issuer.com", provider.OIDC.IssuerURL) + assert.Equal(t, "client-id", provider.OIDC.ClientID) + assert.Equal(t, "client-secret", provider.OIDC.ClientSecret) + assert.Equal(t, "https://pulse.example.com/api/oidc/callback", provider.OIDC.RedirectURL) + assert.Equal(t, []string{"admins", "operators"}, provider.AllowedGroups) + assert.Equal(t, map[string]string{"admins": "admin", "operators": "viewer"}, provider.GroupRoleMappings) + assert.True(t, provider.OIDC.EnvOverrides["clientSecret"]) } func TestLoad_AuthPass_AutoHash(t *testing.T) { diff --git a/internal/config/oidc.go b/internal/config/oidc.go index 89e50ad3e..7f17cbcc9 100644 --- a/internal/config/oidc.go +++ b/internal/config/oidc.go @@ -3,6 +3,7 @@ package config import ( "fmt" "net/url" + "os" "strings" ) @@ -12,6 +13,10 @@ var defaultOIDCScopes = []string{"openid", "profile", "email"} // DefaultOIDCCallbackPath is the path we expose for the OIDC redirect handler. const DefaultOIDCCallbackPath = "/api/oidc/callback" +// LegacyOIDCProviderID is the synthetic SSO provider ID used for v5 OIDC +// configuration migrated from oidc.enc or OIDC_* environment variables. +const LegacyOIDCProviderID = "legacy-oidc" + // OIDCConfig captures configuration required to integrate with an OpenID Connect provider. type OIDCConfig struct { Enabled bool `json:"enabled"` @@ -114,6 +119,116 @@ func DefaultRedirectURL(publicURL string) string { return base + DefaultOIDCCallbackPath } +// LegacyOIDCConfigToSSOConfig converts the v5 single-provider OIDC +// configuration into the v6 multi-provider SSO model. +func LegacyOIDCConfigToSSOConfig(legacy *OIDCConfig, publicURL string) (*SSOConfig, bool) { + provider, ok := LegacyOIDCProviderFromConfig(legacy, publicURL, false) + if !ok { + return nil, false + } + cfg := NewSSOConfig() + cfg.Providers = append(cfg.Providers, *provider) + cfg.DefaultProviderID = provider.ID + return cfg, true +} + +// LegacyOIDCProviderFromConfig converts a v5 OIDC configuration into a single +// v6 SSO OIDC provider. +func LegacyOIDCProviderFromConfig(legacy *OIDCConfig, publicURL string, runtimeManaged bool) (*SSOProvider, bool) { + if legacy == nil || !legacy.Enabled { + return nil, false + } + + cfg := legacy.Clone() + cfg.ApplyDefaults(publicURL) + if strings.TrimSpace(cfg.IssuerURL) == "" || strings.TrimSpace(cfg.ClientID) == "" { + return nil, false + } + + provider := &SSOProvider{ + ID: LegacyOIDCProviderID, + Name: "Single Sign-On", + Type: SSOProviderTypeOIDC, + Enabled: true, + DisplayName: "Single Sign-On", + AllowedGroups: append([]string{}, cfg.AllowedGroups...), + AllowedDomains: append([]string{}, cfg.AllowedDomains...), + AllowedEmails: append([]string{}, cfg.AllowedEmails...), + GroupsClaim: cfg.GroupsClaim, + GroupRoleMappings: cloneStringMap(cfg.GroupRoleMappings), + RuntimeManaged: runtimeManaged, + OIDC: &OIDCProviderConfig{ + IssuerURL: strings.TrimSpace(cfg.IssuerURL), + ClientID: strings.TrimSpace(cfg.ClientID), + ClientSecret: cfg.ClientSecret, + RedirectURL: cfg.RedirectURL, + LogoutURL: cfg.LogoutURL, + Scopes: append([]string{}, cfg.Scopes...), + UsernameClaim: cfg.UsernameClaim, + EmailClaim: cfg.EmailClaim, + CABundle: cfg.CABundle, + ClientSecretSet: cfg.ClientSecret != "", + EnvOverrides: cloneBoolMap(cfg.EnvOverrides), + }, + } + return provider, true +} + +// LegacyOIDCEnvProvider converts OIDC_* environment variables into a +// runtime-managed SSO provider. +func LegacyOIDCEnvProvider(publicURL string) (*SSOProvider, bool) { + legacy, ok := LegacyOIDCConfigFromEnv(publicURL) + if !ok { + return nil, false + } + return LegacyOIDCProviderFromConfig(legacy, publicURL, true) +} + +// LegacyOIDCConfigFromEnv reads the v5 OIDC_* environment contract. +func LegacyOIDCConfigFromEnv(publicURL string) (*OIDCConfig, bool) { + enabledRaw, enabledSet := envValue("OIDC_ENABLED") + if !enabledSet || !truthyEnv(enabledRaw) { + return nil, false + } + + cfg := &OIDCConfig{ + Enabled: true, + IssuerURL: envTrim("OIDC_ISSUER_URL"), + ClientID: envTrim("OIDC_CLIENT_ID"), + ClientSecret: envTrim("OIDC_CLIENT_SECRET"), + RedirectURL: envTrim("OIDC_REDIRECT_URL"), + LogoutURL: envTrim("OIDC_LOGOUT_URL"), + UsernameClaim: envTrim("OIDC_USERNAME_CLAIM"), + EmailClaim: envTrim("OIDC_EMAIL_CLAIM"), + GroupsClaim: envTrim("OIDC_GROUPS_CLAIM"), + CABundle: envTrim("OIDC_CA_BUNDLE"), + Scopes: splitOIDCEnvList(envTrim("OIDC_SCOPES")), + AllowedGroups: splitOIDCEnvList(envTrim("OIDC_ALLOWED_GROUPS")), + AllowedDomains: splitOIDCEnvList(envTrim("OIDC_ALLOWED_DOMAINS")), + AllowedEmails: splitOIDCEnvList(envTrim("OIDC_ALLOWED_EMAILS")), + GroupRoleMappings: parseOIDCGroupRoleMappings(envTrim("OIDC_GROUP_ROLE_MAPPINGS")), + EnvOverrides: map[string]bool{"enabled": true}, + } + + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_ISSUER_URL", "issuerUrl") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_CLIENT_ID", "clientId") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_CLIENT_SECRET", "clientSecret") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_REDIRECT_URL", "redirectUrl") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_LOGOUT_URL", "logoutUrl") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_SCOPES", "scopes") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_USERNAME_CLAIM", "usernameClaim") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_EMAIL_CLAIM", "emailClaim") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_GROUPS_CLAIM", "groupsClaim") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_ALLOWED_GROUPS", "allowedGroups") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_ALLOWED_DOMAINS", "allowedDomains") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_ALLOWED_EMAILS", "allowedEmails") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_GROUP_ROLE_MAPPINGS", "groupRoleMappings") + markOIDCEnvOverride(cfg.EnvOverrides, "OIDC_CA_BUNDLE", "caBundle") + + cfg.ApplyDefaults(publicURL) + return cfg, true +} + // Validate performs sanity checks and returns the first error encountered. func (c *OIDCConfig) Validate() error { if c == nil { @@ -167,3 +282,86 @@ func normaliseList(values []string) []string { } return result } + +func envTrim(name string) string { + return strings.TrimSpace(os.Getenv(name)) +} + +func envValue(name string) (string, bool) { + value, ok := os.LookupEnv(name) + return strings.TrimSpace(value), ok +} + +func truthyEnv(value string) bool { + switch strings.ToLower(strings.TrimSpace(value)) { + case "1", "true", "yes", "y", "on": + return true + default: + return false + } +} + +func splitOIDCEnvList(value string) []string { + if strings.TrimSpace(value) == "" { + return []string{} + } + parts := strings.FieldsFunc(value, func(r rune) bool { + return r == ',' || r == ' ' || r == '\t' || r == '\n' || r == '\r' + }) + return normaliseList(parts) +} + +func parseOIDCGroupRoleMappings(value string) map[string]string { + parts := splitOIDCEnvList(value) + if len(parts) == 0 { + return nil + } + mappings := make(map[string]string, len(parts)) + for _, part := range parts { + group, role, ok := strings.Cut(part, "=") + if !ok { + continue + } + group = strings.TrimSpace(group) + role = strings.TrimSpace(role) + if group == "" || role == "" { + continue + } + mappings[group] = role + } + if len(mappings) == 0 { + return nil + } + return mappings +} + +func markOIDCEnvOverride(overrides map[string]bool, envName string, field string) { + if overrides == nil { + return + } + if _, ok := os.LookupEnv(envName); ok { + overrides[field] = true + } +} + +func cloneStringMap(values map[string]string) map[string]string { + if len(values) == 0 { + return nil + } + clone := make(map[string]string, len(values)) + for key, value := range values { + clone[key] = value + } + return clone +} + +func cloneBoolMap(values map[string]bool) map[string]bool { + if len(values) == 0 { + return nil + } + clone := make(map[string]bool, len(values)) + for key, value := range values { + clone[key] = value + } + return clone +} diff --git a/internal/config/persistence.go b/internal/config/persistence.go index 33ee0561a..a73501da9 100644 --- a/internal/config/persistence.go +++ b/internal/config/persistence.go @@ -38,6 +38,7 @@ type ConfigPersistence struct { availabilityFile string systemFile string ssoFile string + oidcFile string apiTokensFile string aiFile string aiFindingsFile string @@ -114,6 +115,7 @@ type resolvedConfigPersistencePaths struct { availabilityFile string systemFile string ssoFile string + oidcFile string apiTokensFile string aiFile string aiFindingsFile string @@ -178,6 +180,10 @@ func resolveConfigPersistencePaths(configDir string) (string, resolvedConfigPers if err != nil { return "", resolvedConfigPersistencePaths{}, fmt.Errorf("resolve sso.enc: %w", err) } + oidcFile, err := resolveLeaf("oidc.enc") + if err != nil { + return "", resolvedConfigPersistencePaths{}, fmt.Errorf("resolve oidc.enc: %w", err) + } apiTokensFile, err := resolveLeaf("api_tokens.json") if err != nil { return "", resolvedConfigPersistencePaths{}, fmt.Errorf("resolve api_tokens.json: %w", err) @@ -238,6 +244,7 @@ func resolveConfigPersistencePaths(configDir string) (string, resolvedConfigPers availabilityFile: availabilityFile, systemFile: systemFile, ssoFile: ssoFile, + oidcFile: oidcFile, apiTokensFile: apiTokensFile, aiFile: aiFile, aiFindingsFile: aiFindingsFile, @@ -291,6 +298,7 @@ func newConfigPersistence(configDir string) (*ConfigPersistence, error) { availabilityFile: resolvedPaths.availabilityFile, systemFile: resolvedPaths.systemFile, ssoFile: resolvedPaths.ssoFile, + oidcFile: resolvedPaths.oidcFile, apiTokensFile: resolvedPaths.apiTokensFile, aiFile: resolvedPaths.aiFile, aiFindingsFile: resolvedPaths.aiFindingsFile, @@ -1972,11 +1980,19 @@ func (c *ConfigPersistence) SaveSSOConfig(settings *SSOConfig) error { // Clone to avoid modifying the original clone := settings.Clone() - // Clear sensitive data from OIDC providers (env overrides) - for i := range clone.Providers { - if clone.Providers[i].OIDC != nil { - clone.Providers[i].OIDC.EnvOverrides = nil + persistedProviders := clone.Providers[:0] + for _, provider := range clone.Providers { + if provider.RuntimeManaged { + continue } + if provider.OIDC != nil { + provider.OIDC.EnvOverrides = nil + } + persistedProviders = append(persistedProviders, provider) + } + clone.Providers = persistedProviders + if clone.DefaultProviderID != "" && clone.GetProvider(clone.DefaultProviderID) == nil { + clone.DefaultProviderID = "" } data, err := json.Marshal(clone) @@ -2009,7 +2025,7 @@ func (c *ConfigPersistence) LoadSSOConfig() (*SSOConfig, error) { migratedPlaintext, _, err := loadEncryptedJSONLocked(c, c.ssoFile, &settings, "sso config") if err != nil { if os.IsNotExist(err) { - return nil, nil + return c.loadLegacyOIDCConfigAsSSOLocked() } return nil, err } @@ -2028,6 +2044,40 @@ func (c *ConfigPersistence) LoadSSOConfig() (*SSOConfig, error) { return &settings, nil } +func (c *ConfigPersistence) loadLegacyOIDCConfigAsSSOLocked() (*SSOConfig, error) { + var legacy OIDCConfig + _, _, err := loadEncryptedJSONLocked(c, c.oidcFile, &legacy, "legacy oidc config") + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, err + } + + settings, ok := LegacyOIDCConfigToSSOConfig(&legacy, "") + if !ok { + if legacy.Enabled { + log.Warn().Str("file", c.oidcFile).Msg("Legacy OIDC config is enabled but incomplete; skipping SSO migration") + } + return nil, nil + } + + jsonData, err := json.Marshal(settings) + if err != nil { + return nil, fmt.Errorf("marshal legacy oidc migration: %w", err) + } + if err := rewriteEncryptedJSONLocked(c, c.ssoFile, jsonData, "legacy oidc migration to sso config"); err != nil { + return nil, err + } + + log.Info(). + Str("from", c.oidcFile). + Str("to", c.ssoFile). + Int("providers", len(settings.Providers)). + Msg("Migrated legacy OIDC configuration to SSO configuration") + return settings, nil +} + // SaveAIConfig stores AI settings, encrypting them when a crypto manager is available. func (c *ConfigPersistence) SaveAIConfig(settings AIConfig) error { c.mu.Lock() diff --git a/internal/config/persistence_sso_coverage_test.go b/internal/config/persistence_sso_coverage_test.go index 3104d9320..f44f82f86 100644 --- a/internal/config/persistence_sso_coverage_test.go +++ b/internal/config/persistence_sso_coverage_test.go @@ -114,17 +114,46 @@ func TestSaveSSOConfig_ErrorPaths(t *testing.T) { }) } -func TestLoadSSOConfig_DoesNotMigrateLegacyOIDC(t *testing.T) { +func TestLoadSSOConfig_MigratesLegacyOIDC(t *testing.T) { tempDir := t.TempDir() cp := NewConfigPersistence(tempDir) cp.crypto = nil - require.NoError(t, os.WriteFile(filepath.Join(tempDir, "oidc.enc"), []byte(`{"enabled":true}`), 0600)) + legacy := OIDCConfig{ + Enabled: true, + IssuerURL: "https://issuer.example.com", + ClientID: "pulse-client", + ClientSecret: "client-secret", + RedirectURL: "https://pulse.example.com/api/oidc/callback", + Scopes: []string{"openid", "email"}, + GroupsClaim: "groups", + AllowedGroups: []string{"admins"}, + GroupRoleMappings: map[string]string{ + "admins": "admin", + }, + } + raw, err := json.Marshal(legacy) + require.NoError(t, err) + require.NoError(t, os.WriteFile(filepath.Join(tempDir, "oidc.enc"), raw, 0600)) cfg, err := cp.LoadSSOConfig() require.NoError(t, err) - assert.Nil(t, cfg) - assert.NoFileExists(t, filepath.Join(tempDir, "sso.enc")) + require.NotNil(t, cfg) + provider := cfg.GetProvider(LegacyOIDCProviderID) + require.NotNil(t, provider) + require.NotNil(t, provider.OIDC) + assert.Equal(t, "Single Sign-On", provider.DisplayName) + assert.Equal(t, SSOProviderTypeOIDC, provider.Type) + assert.True(t, provider.Enabled) + assert.Equal(t, "https://issuer.example.com", provider.OIDC.IssuerURL) + assert.Equal(t, "pulse-client", provider.OIDC.ClientID) + assert.Equal(t, "client-secret", provider.OIDC.ClientSecret) + assert.Equal(t, "https://pulse.example.com/api/oidc/callback", provider.OIDC.RedirectURL) + assert.Equal(t, []string{"openid", "email"}, provider.OIDC.Scopes) + assert.Equal(t, "groups", provider.GroupsClaim) + assert.Equal(t, []string{"admins"}, provider.AllowedGroups) + assert.Equal(t, map[string]string{"admins": "admin"}, provider.GroupRoleMappings) + assert.FileExists(t, filepath.Join(tempDir, "sso.enc")) } func TestLoadSSOConfig_FallbackAndErrors(t *testing.T) { @@ -137,12 +166,12 @@ func TestLoadSSOConfig_FallbackAndErrors(t *testing.T) { assert.Nil(t, cfg) }) - t.Run("legacy oidc present still returns nil config", func(t *testing.T) { + t.Run("disabled legacy oidc returns nil config", func(t *testing.T) { tempDir := t.TempDir() cp := NewConfigPersistence(tempDir) cp.crypto = nil - require.NoError(t, os.WriteFile(filepath.Join(tempDir, "oidc.enc"), []byte(`{"enabled":true}`), 0600)) + require.NoError(t, os.WriteFile(filepath.Join(tempDir, "oidc.enc"), []byte(`{"enabled":false,"issuerUrl":"https://issuer.example.com","clientId":"pulse-client"}`), 0600)) cfg, err := cp.LoadSSOConfig() require.NoError(t, err) diff --git a/internal/config/sso.go b/internal/config/sso.go index d855ac0fc..d1b0f2b41 100644 --- a/internal/config/sso.go +++ b/internal/config/sso.go @@ -42,6 +42,10 @@ type SSOProvider struct { // SAML-specific configuration SAML *SAMLProviderConfig `json:"saml,omitempty"` + + // RuntimeManaged marks providers sourced from process environment. Runtime + // providers are usable for auth discovery but must not be persisted. + RuntimeManaged bool `json:"-"` } // OIDCProviderConfig contains OIDC-specific settings @@ -121,6 +125,34 @@ func NewSSOConfig() *SSOConfig { } } +// ApplyLegacyOIDCEnvProvider adds the legacy OIDC_* environment configuration +// as a runtime-managed SSO provider. It returns true when a provider is newly +// added. Existing persisted providers with the same ID are left untouched. +func ApplyLegacyOIDCEnvProvider(c *SSOConfig, publicURL string) bool { + if c == nil { + return false + } + + provider, ok := LegacyOIDCEnvProvider(publicURL) + if !ok { + return false + } + + existing := c.GetProvider(provider.ID) + if existing != nil { + if existing.RuntimeManaged { + *existing = *provider + } + return false + } + + c.Providers = append(c.Providers, *provider) + if c.DefaultProviderID == "" || c.GetProvider(c.DefaultProviderID) == nil { + c.DefaultProviderID = provider.ID + } + return true +} + // GetProvider returns a provider by ID func (c *SSOConfig) GetProvider(id string) *SSOProvider { if c == nil {