mirror of
https://github.com/rcourtman/Pulse.git
synced 2026-09-10 02:25:56 +00:00
@@ -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
|
||||
|
||||
@@ -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(`<html><head>` +
|
||||
`<script type="importmap">{"integrity":{}}</script>` +
|
||||
`<script src="/assets/index.js"></script>` +
|
||||
`<style>body{color:#111}</style>` +
|
||||
`<script nonce="__CSP_NONCE__">window.__theme="dark"</script>` +
|
||||
`</head><body></body></html>`)
|
||||
|
||||
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, `<script type="importmap" nonce="test-nonce">`) {
|
||||
t.Fatalf("expected import map script to receive nonce, got body %q", body)
|
||||
}
|
||||
if !strings.Contains(body, `<style nonce="test-nonce">body{color:#111}</style>`) {
|
||||
t.Fatalf("expected inline style to receive nonce, got body %q", body)
|
||||
}
|
||||
if !strings.Contains(body, `<script src="/assets/index.js"></script>`) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 ""
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user