Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
181 changes: 140 additions & 41 deletions authentication/oidc.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,15 @@ import (
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"net"
"net/http"
"os"
"path"
"reflect"
"strings"
"time"

Expand All @@ -31,6 +33,68 @@ import (
"github.com/observatorium/api/httperr"
)

// StringOrSlice is a []string that can be unmarshaled from either a single
// JSON string or a JSON array of strings. When a single string is provided
// it becomes a one-element slice, preserving backward compatibility with
// configs that specify a single claim name.
type StringOrSlice []string

// UnmarshalJSON implements json.Unmarshaler.
// It accepts both `"claim"` and `["claim1","claim2"]`.
func (s *StringOrSlice) UnmarshalJSON(data []byte) error {
var str string
if err := json.Unmarshal(data, &str); err == nil {
*s = StringOrSlice{str}

return nil
}

var arr []string
if err := json.Unmarshal(data, &arr); err != nil {
return fmt.Errorf("StringOrSlice: expected string or []string: %w", err)
}

*s = StringOrSlice(arr)

return nil
}

// MarshalJSON implements json.Marshaler.
func (s StringOrSlice) MarshalJSON() ([]byte, error) {
return json.Marshal([]string(s))
}

// stringOrSliceDecodeHook is a mapstructure DecodeHookFunc that converts
// a string or []interface{} to StringOrSlice when the target type matches.
func stringOrSliceDecodeHook() mapstructure.DecodeHookFunc {
return func(f reflect.Type, t reflect.Type, data interface{}) (interface{}, error) {
if t != reflect.TypeOf(StringOrSlice{}) {
return data, nil
}

switch v := data.(type) {
case string:
return StringOrSlice{v}, nil
case []interface{}:
result := make([]string, len(v))
for i, elem := range v {
s, ok := elem.(string)
if !ok {
return nil, fmt.Errorf("StringOrSlice element %d is not a string", i)
}

result[i] = s
}

return StringOrSlice(result), nil
case []string:
return v, nil
default:
return data, nil
}
}
}

// OIDCAuthenticatorType represents the oidc authentication provider type.
const OIDCAuthenticatorType = "oidc"

Expand All @@ -40,15 +104,15 @@ func init() {

// oidcConfig represents the oidc authenticator config.
type oidcConfig struct {
ClientID string `json:"clientID"`
ClientSecret string `json:"clientSecret"`
GroupClaim string `json:"groupClaim"`
IssuerRawCA []byte `json:"issuerCA"`
IssuerCAPath string `json:"issuerCAPath"`
ClientID string `json:"clientID"`
ClientSecret string `json:"clientSecret"`
GroupClaim StringOrSlice `json:"groupClaim"`
IssuerRawCA []byte `json:"issuerCA"`
IssuerCAPath string `json:"issuerCAPath"`
issuerCA *x509.Certificate
IssuerURL string `json:"issuerURL"`
RedirectURL string `json:"redirectURL"`
UsernameClaim string `json:"usernameClaim"`
IssuerURL string `json:"issuerURL"`
RedirectURL string `json:"redirectURL"`
UsernameClaim StringOrSlice `json:"usernameClaim"`
}

type oidcAuthenticator struct {
Expand Down Expand Up @@ -77,8 +141,15 @@ func newOIDCAuthenticator(c map[string]interface{}, tenant string,

ctx := context.Background()

err := mapstructure.Decode(c, &config)
decoder, err := mapstructure.NewDecoder(&mapstructure.DecoderConfig{
DecodeHook: stringOrSliceDecodeHook(),
Result: &config,
})
if err != nil {
return nil, fmt.Errorf("create config decoder: %w", err)
}

if err := decoder.Decode(c); err != nil {
return nil, err
}

Expand Down Expand Up @@ -371,69 +442,73 @@ func (a oidcAuthenticator) checkAuth(ctx context.Context, token string) (context

sub := idToken.Subject
usernameExtracted := false
hasUsernameClaims := len(a.config.UsernameClaim) > 0
hasGroupClaims := len(a.config.GroupClaim) > 0

// Extract claims once for both username and group lookups.
var claims map[string]interface{}

if a.config.UsernameClaim != "" {
claims := map[string]interface{}{}
if hasUsernameClaims || hasGroupClaims {
claims = map[string]interface{}{}
if err := idToken.Claims(&claims); err != nil {
const msg = "failed to read claims"

level.Warn(a.logger).Log("msg", msg, "err", err)

return ctx, msg, http.StatusInternalServerError, codes.Internal
}
}

rawUsername, ok := claims[a.config.UsernameClaim]
if !ok {
if a.config.GroupClaim == "" {
const msg = "username cannot be empty"

level.Debug(a.logger).Log("msg", msg)

return ctx, msg, http.StatusBadRequest, codes.PermissionDenied
// Try username claims in order of preference — first match wins.
if hasUsernameClaims {
for _, claimName := range a.config.UsernameClaim {
rawUsername, ok := claims[claimName]
if !ok {
continue
}

level.Debug(a.logger).Log("msg", "username claim not found in token, proceeding with group claim")
} else {
username, ok := rawUsername.(string)
if !ok || username == "" {
const msg = "invalid username claim value"

level.Debug(a.logger).Log("msg", msg)
level.Debug(a.logger).Log("msg", msg, "claim", claimName)

return ctx, msg, http.StatusBadRequest, codes.PermissionDenied
}

sub = username
usernameExtracted = true
}
}

ctx = context.WithValue(ctx, subjectKey, sub)
break
}

if a.config.GroupClaim != "" {
var groups []string
if !usernameExtracted && !hasGroupClaims {
const msg = "username cannot be empty"

claims := map[string]interface{}{}
if err := idToken.Claims(&claims); err != nil {
const msg = "failed to read claims"
level.Debug(a.logger).Log("msg", msg)

level.Warn(a.logger).Log("msg", msg, "err", err)
return ctx, msg, http.StatusBadRequest, codes.PermissionDenied
}

return ctx, msg, http.StatusInternalServerError, codes.Internal
if !usernameExtracted {
level.Debug(a.logger).Log("msg", "username claim not found in token, proceeding with group claim")
}
}

rawGroup, ok := claims[a.config.GroupClaim]
if !ok {
if usernameExtracted {
level.Debug(a.logger).Log("msg", "group claim not found, proceeding with username")
} else {
const msg = "group cannot be empty"
ctx = context.WithValue(ctx, subjectKey, sub)

level.Debug(a.logger).Log("msg", msg)
// Try group claims in order of preference — first match wins.
groupExtracted := false

return ctx, msg, http.StatusBadRequest, codes.PermissionDenied
if hasGroupClaims {
for _, claimName := range a.config.GroupClaim {
rawGroup, ok := claims[claimName]
if !ok {
continue
}
} else {

var groups []string

switch v := rawGroup.(type) {
case string:
groups = append(groups, v)
Expand All @@ -447,7 +522,31 @@ func (a oidcAuthenticator) checkAuth(ctx context.Context, token string) (context
}

ctx = context.WithValue(ctx, groupsKey, groups)
groupExtracted = true

break
}

if !groupExtracted {
if usernameExtracted {
level.Debug(a.logger).Log("msg", "group claim not found, proceeding with username")
} else {
const msg = "group cannot be empty"

level.Debug(a.logger).Log("msg", msg)

return ctx, msg, http.StatusBadRequest, codes.PermissionDenied
}
}
}

// Safety gate: when both claim lists are configured, at least one must match.
if hasUsernameClaims && hasGroupClaims && !usernameExtracted && !groupExtracted {
const msg = "no matching claims found"

level.Debug(a.logger).Log("msg", msg)

return ctx, msg, http.StatusBadRequest, codes.PermissionDenied
}

return ctx, "", http.StatusOK, codes.OK
Expand Down
Loading
Loading