diff --git a/opentdf-dev.yaml b/opentdf-dev.yaml index 6f7754f7c9..93a471101f 100644 --- a/opentdf-dev.yaml +++ b/opentdf-dev.yaml @@ -75,6 +75,10 @@ server: groups_claim: # realm_access.roles # Claim the represents the idP client ID client_id_claim: # azp + # Optional external role provider (name is resolved via StartOptions) + # roles_provider: + # name: external + # config: {} # provider-specific (any object) ## Extends the builtin policy extension: | g, opentdf-admin, role:admin diff --git a/service/internal/auth/authn.go b/service/internal/auth/authn.go index 24c72b56ea..016875b589 100644 --- a/service/internal/auth/authn.go +++ b/service/internal/auth/authn.go @@ -28,6 +28,7 @@ import ( "google.golang.org/grpc/metadata" ctxAuth "github.com/opentdf/platform/service/pkg/auth" + "github.com/opentdf/platform/service/pkg/authz" ) var ( @@ -143,8 +144,13 @@ func NewAuthenticator(ctx context.Context, cfg Config, logger *logger.Logger, we return nil, err } + roleProvider, err := resolveRoleProvider(ctx, cfg, logger) + if err != nil { + return nil, err + } casbinConfig := CasbinConfig{ PolicyConfig: cfg.Policy, + RoleProvider: roleProvider, } logger.Info("initializing casbin enforcer") if a.enforcer, err = NewCasbinEnforcer(casbinConfig, a.logger); err != nil { @@ -276,8 +282,13 @@ func (a Authentication) MuxHandler(handler http.Handler) http.Handler { default: action = ActionUnsafe } - if allow, err := a.enforcer.Enforce(accessTok, r.URL.Path, action); err != nil { - if err.Error() == "permission denied" { + roleReq := authz.RoleRequest{ + Issuer: a.oidcConfiguration.Issuer, + Resource: r.URL.Path, + Action: action, + } + if allow, err := a.enforcer.Enforce(ctx, accessTok, roleReq); err != nil { + if errors.Is(err, ErrPermissionDenied) { log.WarnContext( ctx, "permission denied", @@ -366,8 +377,13 @@ func (a Authentication) ConnectUnaryServerInterceptor() connect.UnaryInterceptor } // Check if the token is allowed to access the resource - if allowed, err := a.enforcer.Enforce(token, resource, action); err != nil { - if err.Error() == "permission denied" { + roleReq := authz.RoleRequest{ + Issuer: a.oidcConfiguration.Issuer, + Resource: resource, + Action: action, + } + if allowed, err := a.enforcer.Enforce(ctxWithJWK, token, roleReq); err != nil { + if errors.Is(err, ErrPermissionDenied) { log.WarnContext( ctxWithJWK, "permission denied", diff --git a/service/internal/auth/casbin.go b/service/internal/auth/casbin.go index ac9a40f598..f54e9d6ed4 100644 --- a/service/internal/auth/casbin.go +++ b/service/internal/auth/casbin.go @@ -1,6 +1,7 @@ package auth import ( + "context" "errors" "fmt" "log/slog" @@ -11,13 +12,15 @@ import ( stringadapter "github.com/casbin/casbin/v2/persist/string-adapter" "github.com/lestrrat-go/jwx/v2/jwt" "github.com/opentdf/platform/service/logger" + "github.com/opentdf/platform/service/pkg/authz" _ "embed" ) var ( - rolePrefix = "role:" - defaultRole = "unknown" + rolePrefix = "role:" + defaultRole = "unknown" + ErrPermissionDenied = errors.New("permission denied") ) //go:embed casbin_policy.csv @@ -34,12 +37,14 @@ type Enforcer struct { isDefaultPolicy bool isDefaultModel bool + roleProvider authz.RoleProvider } type casbinSubject []string type CasbinConfig struct { PolicyConfig + RoleProvider authz.RoleProvider } // newCasbinEnforcer creates a new casbin enforcer @@ -112,6 +117,11 @@ func NewCasbinEnforcer(c CasbinConfig, logger *logger.Logger) (*Enforcer, error) return nil, fmt.Errorf("failed to create casbin enforcer: %w", err) } + roleProvider := c.RoleProvider + if roleProvider == nil { + roleProvider = newJWTClaimsRoleProvider(c.GroupsClaim, logger) + } + return &Enforcer{ Enforcer: e, Config: c, @@ -119,16 +129,23 @@ func NewCasbinEnforcer(c CasbinConfig, logger *logger.Logger) (*Enforcer, error) isDefaultPolicy: isDefaultPolicy, isDefaultModel: isDefaultModel, logger: logger, + roleProvider: roleProvider, }, nil } // casbinEnforce is a helper function to enforce the policy with casbin // TODO implement a common type so this can be used for both http and grpc -func (e *Enforcer) Enforce(token jwt.Token, resource, action string) (bool, error) { +func (e *Enforcer) Enforce(ctx context.Context, token jwt.Token, req authz.RoleRequest) (bool, error) { // extract the role claim from the token - s := e.buildSubjectFromToken(token) + s, err := e.buildSubjectFromToken(ctx, token, req) + if err != nil { + e.logger.Warn("role provider error", slog.Any("error", err)) + return false, ErrPermissionDenied + } s = append(s, rolePrefix+defaultRole) + resource := req.Resource + action := req.Action for _, info := range s { allowed, err := e.Enforcer.Enforce(info, resource, action) if err != nil { @@ -153,15 +170,18 @@ func (e *Enforcer) Enforce(token jwt.Token, resource, action string) (bool, erro slog.String("action", action), slog.String("resource", resource), ) - return false, errors.New("permission denied") + return false, ErrPermissionDenied } -func (e *Enforcer) buildSubjectFromToken(t jwt.Token) casbinSubject { +func (e *Enforcer) buildSubjectFromToken(ctx context.Context, t jwt.Token, req authz.RoleRequest) (casbinSubject, error) { var subject string info := casbinSubject{} e.logger.Debug("building subject from token") - roles := e.extractRolesFromToken(t) + roles, err := e.roleProvider.Roles(ctx, t, req) + if err != nil { + return nil, err + } if claim, found := t.Get(e.Config.UserNameClaim); found { sub, ok := claim.(string) @@ -176,66 +196,5 @@ func (e *Enforcer) buildSubjectFromToken(t jwt.Token) casbinSubject { } info = append(info, roles...) info = append(info, subject) - return info -} - -func (e *Enforcer) extractRolesFromToken(t jwt.Token) []string { - e.logger.Debug("extracting roles from token") - roles := []string{} - - roleClaim := e.Config.GroupsClaim - // roleMap := e.Config.RoleMap - - selectors := strings.Split(roleClaim, ".") - claim, exists := t.Get(selectors[0]) - if !exists { - e.logger.Warn("claim not found", - slog.String("claim", roleClaim), - slog.Any("claims", claim), - ) - return nil - } - e.logger.Debug("root claim found", - slog.String("claim", roleClaim), - slog.Any("claims", claim), - ) - // use dotnotation if the claim is nested - if len(selectors) > 1 { - claimMap, ok := claim.(map[string]interface{}) - if !ok { - e.logger.Warn("claim is not of type map[string]interface{}", - slog.String("claim", roleClaim), - slog.Any("claims", claim), - ) - return nil - } - claim = dotNotation(claimMap, strings.Join(selectors[1:], ".")) - if claim == nil { - e.logger.Warn("claim not found", - slog.String("claim", roleClaim), - slog.Any("claims", claim), - ) - return nil - } - } - - // check the type of the role claim - switch v := claim.(type) { - case string: - roles = append(roles, v) - case []interface{}: - for _, rr := range v { - if r, ok := rr.(string); ok { - roles = append(roles, r) - } - } - default: - e.logger.Warn("could not get claim type", - slog.String("selector", roleClaim), - slog.Any("claims", claim), - ) - return nil - } - - return roles + return info, nil } diff --git a/service/internal/auth/casbin_test.go b/service/internal/auth/casbin_test.go index a67f45fb0b..c22acc36db 100644 --- a/service/internal/auth/casbin_test.go +++ b/service/internal/auth/casbin_test.go @@ -1,6 +1,7 @@ package auth import ( + "context" "fmt" "log/slog" "strings" @@ -10,6 +11,7 @@ import ( "github.com/creasty/defaults" "github.com/lestrrat-go/jwx/v2/jwt" "github.com/opentdf/platform/service/logger" + "github.com/opentdf/platform/service/pkg/authz" "github.com/stretchr/testify/suite" ) @@ -80,7 +82,7 @@ func (s *AuthnCasbinSuite) Test_NewEnforcerWithCustomModel() { }) s.Require().NoError(err) - allowed, err := enforcer.Enforce(tok, "", "") + allowed, err := s.enforce(enforcer, tok, "", "") s.Require().NoError(err) s.True(allowed) } @@ -244,7 +246,7 @@ func (s *AuthnCasbinSuite) Test_Enforcement() { enforcer, err := NewCasbinEnforcer(CasbinConfig{PolicyConfig: policyCfg}, logger.CreateTestLogger()) s.Require().NoError(err, name) tok := s.newTokWithDefaultClaim(test.roles[0], test.roles[1], "", "") - allowed, err := enforcer.Enforce(tok, test.resource, test.action) + allowed, err := s.enforce(enforcer, tok, test.resource, test.action) if !test.allowed { s.Require().Error(err, name) } else { @@ -261,7 +263,7 @@ func (s *AuthnCasbinSuite) Test_Enforcement() { }, logger.CreateTestLogger()) s.Require().NoError(err, name) _, tok = s.newTokenWithCustomClaim(test.roles[0], test.roles[1]) - allowed, err = enforcer.Enforce(tok, test.resource, test.action) + allowed, err = s.enforce(enforcer, tok, test.resource, test.action) if !test.allowed { s.Require().Error(err, name) } else { @@ -282,7 +284,7 @@ func (s *AuthnCasbinSuite) Test_Enforcement() { }, logger.CreateTestLogger()) s.Require().NoError(err, name) _, tok = s.newTokenWithCustomRoleMap(test.roles[0], test.roles[1]) - allowed, err = enforcer.Enforce(tok, test.resource, test.action) + allowed, err = s.enforce(enforcer, tok, test.resource, test.action) if !test.allowed { s.Require().Error(err, name) } else { @@ -306,8 +308,8 @@ func (s *AuthnCasbinSuite) Test_Enforcement() { PolicyConfig: policyCfg, }, logger.CreateTestLogger()) s.Require().NoError(err, name) - _, tok = s.newTokenWithCilentID() - allowed, err = enforcer.Enforce(tok, test.resource, test.action) + _, tok = s.newTokenWithClientID() + allowed, err = s.enforce(enforcer, tok, test.resource, test.action) if !test.allowed { s.Require().Error(err, name) } else { @@ -332,19 +334,19 @@ func (s *AuthnCasbinSuite) Test_ExtendDefaultPolicies() { s.Require().NoError(err) // other roles denied new policy: admin tok := s.newTokWithDefaultClaim(true, false, "", "") - allowed, err := enforcer.Enforce(tok, "new.service.DoSomething", "read") + allowed, err := s.enforce(enforcer, tok, "new.service.DoSomething", "read") s.Require().NoError(err) s.True(allowed) - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "write") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "write") s.Require().NoError(err) s.True(allowed) // other roles denied new policy: standard tok = s.newTokWithDefaultClaim(false, true, "", "") - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "read") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "read") s.Require().NoError(err) s.True(allowed) - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "write") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "write") s.Require().Error(err) s.False(allowed) } @@ -357,7 +359,7 @@ func (s *AuthnCasbinSuite) Test_ExtendDefaultPolicies_MalformedErrors() { enforcer, err := NewCasbinEnforcer(CasbinConfig{PolicyConfig: policyCfg}, logger.CreateTestLogger()) s.Require().NoError(err) tok := s.newTokWithDefaultClaim(true, false, "", "") - allowed, err := enforcer.Enforce(tok, "policy.attributes.DoSomething", "read") + allowed, err := s.enforce(enforcer, tok, "policy.attributes.DoSomething", "read") s.Require().NoError(err) s.True(allowed) @@ -372,7 +374,7 @@ func (s *AuthnCasbinSuite) Test_ExtendDefaultPolicies_MalformedErrors() { }, logger.CreateTestLogger()) s.Require().NoError(err) tok = s.newTokWithDefaultClaim(true, false, "", "") - allowed, err = enforcer.Enforce(tok, "policy.attributes.DoSomething", "read") + allowed, err = s.enforce(enforcer, tok, "policy.attributes.DoSomething", "read") s.Require().NoError(err) s.True(allowed) @@ -387,7 +389,7 @@ func (s *AuthnCasbinSuite) Test_ExtendDefaultPolicies_MalformedErrors() { }, logger.CreateTestLogger()) s.Require().NoError(err) tok = s.newTokWithDefaultClaim(true, false, "", "") - allowed, err = enforcer.Enforce(tok, "policy.attributes.DoSomething", "read") + allowed, err = s.enforce(enforcer, tok, "policy.attributes.DoSomething", "read") s.Require().NoError(err) s.True(allowed) @@ -400,7 +402,7 @@ func (s *AuthnCasbinSuite) Test_ExtendDefaultPolicies_MalformedErrors() { }, logger.CreateTestLogger()) s.Require().NoError(err) tok = s.newTokWithDefaultClaim(true, false, "", "") - allowed, err = enforcer.Enforce(tok, "policy.attributes.DoSomething", "read") + allowed, err = s.enforce(enforcer, tok, "policy.attributes.DoSomething", "read") s.Require().NoError(err) s.True(allowed) @@ -415,7 +417,7 @@ func (s *AuthnCasbinSuite) Test_ExtendDefaultPolicies_MalformedErrors() { }, logger.CreateTestLogger()) s.Require().NoError(err) tok = s.newTokWithDefaultClaim(true, false, "", "") - allowed, err = enforcer.Enforce(tok, "policy.attributes.DoSomething", "read") + allowed, err = s.enforce(enforcer, tok, "policy.attributes.DoSomething", "read") s.Require().NoError(err) s.True(allowed) } @@ -438,46 +440,46 @@ func (s *AuthnCasbinSuite) Test_SetBuiltinPolicy() { // unauthorized role tok := s.newTokWithDefaultClaim(false, false, "", "") - allowed, err := enforcer.Enforce(tok, "new.hello.World", "read") + allowed, err := s.enforce(enforcer, tok, "new.hello.World", "read") s.Require().Error(err) s.False(allowed) - allowed, err = enforcer.Enforce(tok, "new.hello.World", "write") + allowed, err = s.enforce(enforcer, tok, "new.hello.World", "write") s.Require().Error(err) s.False(allowed) - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "read") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "read") s.Require().Error(err) s.False(allowed) - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "write") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "write") s.Require().Error(err) s.False(allowed) // other roles denied new policy: admin tok = s.newTokWithDefaultClaim(true, false, "", "") - allowed, err = enforcer.Enforce(tok, "new.hello.World", "read") + allowed, err = s.enforce(enforcer, tok, "new.hello.World", "read") s.Require().NoError(err) s.True(allowed) - allowed, err = enforcer.Enforce(tok, "new.hello.World", "write") + allowed, err = s.enforce(enforcer, tok, "new.hello.World", "write") s.Require().NoError(err) s.True(allowed) - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "read") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "read") s.Require().Error(err) s.False(allowed) - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "write") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "write") s.Require().Error(err) s.False(allowed) // other roles denied new policy: standard tok = s.newTokWithDefaultClaim(false, true, "", "") - allowed, err = enforcer.Enforce(tok, "new.hello.World", "read") + allowed, err = s.enforce(enforcer, tok, "new.hello.World", "read") s.Require().NoError(err) s.True(allowed) - allowed, err = enforcer.Enforce(tok, "new.hello.World", "write") + allowed, err = s.enforce(enforcer, tok, "new.hello.World", "write") s.Require().Error(err) s.False(allowed) - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "read") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "read") s.Require().Error(err) s.False(allowed) - allowed, err = enforcer.Enforce(tok, "new.service.DoSomething", "write") + allowed, err = s.enforce(enforcer, tok, "new.service.DoSomething", "write") s.Require().Error(err) s.False(allowed) } @@ -495,15 +497,45 @@ func (s *AuthnCasbinSuite) Test_Username_Policy() { s.Require().NoError(err) tok := s.newTokWithDefaultClaim(true, false, "preferred_username", "") - allowed, err := enforcer.Enforce(tok, "new.service.DoSomething", "read") + allowed, err := s.enforce(enforcer, tok, "new.service.DoSomething", "read") s.Require().NoError(err) s.True(allowed) - allowed, err = enforcer.Enforce(tok, "policy.attributes.List", "read") + allowed, err = s.enforce(enforcer, tok, "policy.attributes.List", "read") s.Require().Error(err) s.False(allowed) } +type staticProvider struct { + roles []string + err error +} + +func (p staticProvider) Roles(_ context.Context, _ jwt.Token, _ authz.RoleRequest) ([]string, error) { + return p.roles, p.err +} + +func (s *AuthnCasbinSuite) Test_ExternalRoleProvider() { + policyCfg := PolicyConfig{} + err := defaults.Set(&policyCfg) + s.Require().NoError(err) + + policyCfg.Extension = strings.Join([]string{ + "p, role:admin, policy.attributes.*, read, allow", + }, "\n") + + enforcer, err := NewCasbinEnforcer(CasbinConfig{ + PolicyConfig: policyCfg, + RoleProvider: staticProvider{roles: []string{"role:admin"}}, + }, logger.CreateTestLogger()) + s.Require().NoError(err) + + tok := jwt.New() + allowed, err := s.enforce(enforcer, tok, "policy.attributes.List", "read") + s.Require().NoError(err) + s.True(allowed) +} + func (s *AuthnCasbinSuite) Test_Override_Of_Username_Claim() { policyCfg := PolicyConfig{} err := defaults.Set(&policyCfg) @@ -518,11 +550,11 @@ func (s *AuthnCasbinSuite) Test_Override_Of_Username_Claim() { s.Require().NoError(err) tok := s.newTokWithDefaultClaim(true, false, "username", "") - allowed, err := enforcer.Enforce(tok, "new.service.DoSomething", "read") + allowed, err := s.enforce(enforcer, tok, "new.service.DoSomething", "read") s.Require().NoError(err) s.True(allowed) - allowed, err = enforcer.Enforce(tok, "policy.attributes.List", "read") + allowed, err = s.enforce(enforcer, tok, "policy.attributes.List", "read") s.Require().Error(err) s.False(allowed) } @@ -538,15 +570,26 @@ func (s *AuthnCasbinSuite) Test_Override_Of_Groups_Claim() { s.Require().NoError(err) tok := s.newTokWithDefaultClaim(false, true, "", "groups") - allowed, err := enforcer.Enforce(tok, "new.service.DoSomething", "read") + allowed, err := s.enforce(enforcer, tok, "new.service.DoSomething", "read") s.Require().Error(err) s.False(allowed) - allowed, err = enforcer.Enforce(tok, "policy.attributes.List", "read") + allowed, err = s.enforce(enforcer, tok, "policy.attributes.List", "read") s.Require().NoError(err) s.True(allowed) } +func (s *AuthnCasbinSuite) enforce(enforcer *Enforcer, tok jwt.Token, resource, action string) (bool, error) { + return enforcer.Enforce( + context.Background(), + tok, + authz.RoleRequest{ + Resource: resource, + Action: action, + }, + ) +} + func (s *AuthnCasbinSuite) buildTokenRoles(admin bool, standard bool, roleMaps []string) []interface{} { adminRole := "opentdf-admin" if len(roleMaps) > 0 { @@ -610,7 +653,7 @@ func (s *AuthnCasbinSuite) newTokenWithCustomRoleMap(admin bool, standard bool) return "", tok } -func (s *AuthnCasbinSuite) newTokenWithCilentID() (string, jwt.Token) { +func (s *AuthnCasbinSuite) newTokenWithClientID() (string, jwt.Token) { tok := jwt.New() if err := tok.Set("client_id", "test"); err != nil { s.T().Fatal(err) diff --git a/service/internal/auth/config.go b/service/internal/auth/config.go index 5e48877cf1..3a449e30d9 100644 --- a/service/internal/auth/config.go +++ b/service/internal/auth/config.go @@ -6,6 +6,7 @@ import ( "github.com/casbin/casbin/v2/persist" "github.com/opentdf/platform/service/logger" + "github.com/opentdf/platform/service/pkg/authz" ) // AuthConfig pulls AuthN and AuthZ together @@ -15,6 +16,10 @@ type Config struct { // Used for re-authentication of IPC connections IPCReauthRoutes []string `mapstructure:"-" json:"-"` AuthNConfig `mapstructure:",squash"` + + // Programmatic role provider overrides (not loaded from config) + RoleProvider authz.RoleProvider `mapstructure:"-" json:"-"` + RoleProviderFactories map[string]authz.RoleProviderFactory `mapstructure:"-" json:"-"` } // AuthNConfig is the configuration need for the platform to validate tokens @@ -34,6 +39,8 @@ type PolicyConfig struct { UserNameClaim string `mapstructure:"username_claim" json:"username_claim" default:"preferred_username"` // Claim to use for group/role information GroupsClaim string `mapstructure:"groups_claim" json:"groups_claim" default:"realm_access.roles"` + // Role provider configuration (resolved via StartOptions) + RolesProvider RolesProviderConfig `mapstructure:"roles_provider" json:"roles_provider"` // Claim to use to reference idP clientID ClientIDClaim string `mapstructure:"client_id_claim" json:"client_id_claim" default:"azp"` // Deprecated: Use GroupClain instead @@ -49,6 +56,11 @@ type PolicyConfig struct { Adapter persist.Adapter `mapstructure:"-" json:"-"` } +type RolesProviderConfig struct { + Name string `mapstructure:"name" json:"name"` + Config map[string]any `mapstructure:"config" json:"config"` +} + func (c AuthNConfig) validateAuthNConfig(logger *logger.Logger) error { if c.Issuer == "" { return errors.New("config Auth.Issuer is required") diff --git a/service/internal/auth/role_provider.go b/service/internal/auth/role_provider.go new file mode 100644 index 0000000000..c00811b23b --- /dev/null +++ b/service/internal/auth/role_provider.go @@ -0,0 +1,118 @@ +package auth + +import ( + "context" + "fmt" + "log/slog" + "strings" + + "github.com/lestrrat-go/jwx/v2/jwt" + "github.com/opentdf/platform/service/logger" + "github.com/opentdf/platform/service/pkg/authz" +) + +type jwtClaimsRoleProvider struct { + groupsClaim string + logger *logger.Logger +} + +func newJWTClaimsRoleProvider(groupsClaim string, logger *logger.Logger) authz.RoleProvider { + return &jwtClaimsRoleProvider{ + groupsClaim: groupsClaim, + logger: logger, + } +} + +func (p *jwtClaimsRoleProvider) Roles(_ context.Context, token jwt.Token, _ authz.RoleRequest) ([]string, error) { + p.logger.Debug("extracting roles from token") + if p.groupsClaim == "" { + p.logger.Warn("groups claim not configured") + return nil, nil + } + + selectors := strings.Split(p.groupsClaim, ".") + claim, exists := token.Get(selectors[0]) + if !exists { + p.logger.Warn("claim not found", + slog.String("claim", p.groupsClaim), + slog.Any("claims", claim), + ) + return nil, nil + } + p.logger.Debug("root claim found", + slog.String("claim", p.groupsClaim), + slog.Any("claims", claim), + ) + + if len(selectors) > 1 { + claimMap, ok := claim.(map[string]interface{}) + if !ok { + p.logger.Warn("claim is not of type map[string]interface{}", + slog.String("claim", p.groupsClaim), + slog.Any("claims", claim), + ) + return nil, nil + } + claim = dotNotation(claimMap, strings.Join(selectors[1:], ".")) + if claim == nil { + p.logger.Warn("claim not found", + slog.String("claim", p.groupsClaim), + slog.Any("claims", claim), + ) + return nil, nil + } + } + + roles := []string{} + switch v := claim.(type) { + case string: + roles = append(roles, v) + case []interface{}: + for _, rr := range v { + if r, ok := rr.(string); ok { + roles = append(roles, r) + } + } + default: + p.logger.Warn("could not get claim type", + slog.String("selector", p.groupsClaim), + slog.Any("claims", claim), + ) + return nil, nil + } + + return roles, nil +} + +func resolveRoleProvider(ctx context.Context, cfg Config, logger *logger.Logger) (authz.RoleProvider, error) { + if cfg.Policy.RolesProvider.Name != "" { + if cfg.RoleProvider != nil && cfg.RoleProviderFactories != nil { + logger.Warn( + "role provider configured in start options is ignored because roles_provider is set", + slog.String("roles_provider", cfg.Policy.RolesProvider.Name), + ) + } + if cfg.RoleProviderFactories == nil { + return nil, fmt.Errorf("no role provider factories are registered, cannot create provider %q", cfg.Policy.RolesProvider.Name) + } + factory, ok := cfg.RoleProviderFactories[cfg.Policy.RolesProvider.Name] + if !ok { + return nil, fmt.Errorf("role provider factory not registered: %s", cfg.Policy.RolesProvider.Name) + } + providerCfg := authz.ProviderConfig{ + Config: cfg.Policy.RolesProvider.Config, + UsernameClaim: cfg.Policy.UserNameClaim, + GroupsClaim: cfg.Policy.GroupsClaim, + ClientIDClaim: cfg.Policy.ClientIDClaim, + } + provider, err := factory(ctx, providerCfg) + if err != nil { + return nil, fmt.Errorf("role provider factory failed: %w", err) + } + return provider, nil + } + if cfg.RoleProvider != nil { + return cfg.RoleProvider, nil + } + return newJWTClaimsRoleProvider(cfg.Policy.GroupsClaim, logger), nil +} diff --git a/service/internal/auth/role_provider_test.go b/service/internal/auth/role_provider_test.go new file mode 100644 index 0000000000..7c0114f7d3 --- /dev/null +++ b/service/internal/auth/role_provider_test.go @@ -0,0 +1,56 @@ +package auth + +import ( + "context" + "testing" + + "github.com/opentdf/platform/service/logger" + "github.com/opentdf/platform/service/pkg/authz" + "github.com/stretchr/testify/require" +) + +func TestResolveRoleProviderDefault(t *testing.T) { + logger := logger.CreateTestLogger() + cfg := Config{} + provider, err := resolveRoleProvider(context.Background(), cfg, logger) + require.NoError(t, err) + require.NotNil(t, provider) + require.IsType(t, &jwtClaimsRoleProvider{}, provider) +} + +func TestResolveRoleProviderNamed(t *testing.T) { + logger := logger.CreateTestLogger() + cfg := Config{ + AuthNConfig: AuthNConfig{ + Policy: PolicyConfig{ + RolesProvider: RolesProviderConfig{ + Name: "mock", + }, + }, + }, + RoleProviderFactories: map[string]authz.RoleProviderFactory{ + "mock": func(_ context.Context, _ authz.ProviderConfig) (authz.RoleProvider, error) { + return staticProvider{roles: []string{"role:admin"}}, nil + }, + }, + } + provider, err := resolveRoleProvider(context.Background(), cfg, logger) + require.NoError(t, err) + require.NotNil(t, provider) +} + +func TestResolveRoleProviderMissingName(t *testing.T) { + logger := logger.CreateTestLogger() + cfg := Config{ + AuthNConfig: AuthNConfig{ + Policy: PolicyConfig{ + RolesProvider: RolesProviderConfig{ + Name: "missing", + }, + }, + }, + } + provider, err := resolveRoleProvider(context.Background(), cfg, logger) + require.Error(t, err) + require.Nil(t, provider) +} diff --git a/service/pkg/authz/role_provider.go b/service/pkg/authz/role_provider.go new file mode 100644 index 0000000000..0146abf549 --- /dev/null +++ b/service/pkg/authz/role_provider.go @@ -0,0 +1,30 @@ +package authz + +import ( + "context" + + "github.com/lestrrat-go/jwx/v2/jwt" +) + +// RoleProvider returns role/group identifiers used as Casbin subjects. +type RoleProvider interface { + Roles(ctx context.Context, token jwt.Token, req RoleRequest) ([]string, error) +} + +// RoleProviderFactory constructs a RoleProvider at startup. +type RoleProviderFactory func(ctx context.Context, cfg ProviderConfig) (RoleProvider, error) + +// ProviderConfig carries provider-specific configuration and claim selectors. +type ProviderConfig struct { + Config map[string]any + UsernameClaim string + GroupsClaim string + ClientIDClaim string +} + +// RoleRequest provides request context to role providers. +type RoleRequest struct { + Issuer string + Resource string + Action string +} diff --git a/service/pkg/server/options.go b/service/pkg/server/options.go index 4cca108124..db3952ceef 100644 --- a/service/pkg/server/options.go +++ b/service/pkg/server/options.go @@ -5,6 +5,7 @@ import ( "connectrpc.com/connect" "github.com/casbin/casbin/v2/persist" + "github.com/opentdf/platform/service/pkg/authz" "github.com/opentdf/platform/service/pkg/config" "github.com/opentdf/platform/service/pkg/serviceregistry" "github.com/opentdf/platform/service/trust" @@ -30,6 +31,9 @@ type StartConfig struct { trustKeyManagerCtxs []trust.NamedKeyManagerCtxFactory + authzRoleProvider authz.RoleProvider + authzRoleProviderFactories map[string]authz.RoleProviderFactory + // CORS additive configuration - appended to YAML/env config values additionalCORSHeaders []string additionalCORSMethods []string @@ -131,6 +135,25 @@ func WithCasbinAdapter(adapter persist.Adapter) StartOptions { } } +// WithAuthZRoleProvider option sets a role provider directly. +func WithAuthZRoleProvider(provider authz.RoleProvider) StartOptions { + return func(c StartConfig) StartConfig { + c.authzRoleProvider = provider + return c + } +} + +// WithAuthZRoleProviderFactory option registers a named role provider factory. +func WithAuthZRoleProviderFactory(name string, factory authz.RoleProviderFactory) StartOptions { + return func(c StartConfig) StartConfig { + if c.authzRoleProviderFactories == nil { + c.authzRoleProviderFactories = make(map[string]authz.RoleProviderFactory) + } + c.authzRoleProviderFactories[name] = factory + return c + } +} + // WithAdditionalConfigLoader option adds an additional configuration loader to the server. func WithAdditionalConfigLoader(loader config.Loader) StartOptions { return func(c StartConfig) StartConfig { diff --git a/service/pkg/server/options_test.go b/service/pkg/server/options_test.go index dd56b17441..ff17bb4b12 100644 --- a/service/pkg/server/options_test.go +++ b/service/pkg/server/options_test.go @@ -5,6 +5,8 @@ import ( "testing" "connectrpc.com/connect" + "github.com/lestrrat-go/jwx/v2/jwt" + "github.com/opentdf/platform/service/pkg/authz" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -139,3 +141,27 @@ func TestWithConnectAndIPCInterceptorsTogether(t *testing.T) { "connect and IPC interceptor slices must be independent", ) } + +type noopRoleProvider struct{} + +func (noopRoleProvider) Roles(_ context.Context, _ jwt.Token, _ authz.RoleRequest) ([]string, error) { + return nil, nil +} + +func TestWithAuthZRoleProvider(t *testing.T) { + var cfg StartConfig + cfg = WithAuthZRoleProvider(noopRoleProvider{})(cfg) + + require.NotNil(t, cfg.authzRoleProvider) + assert.Nil(t, cfg.authzRoleProviderFactories) +} + +func TestWithAuthZRoleProviderFactory(t *testing.T) { + var cfg StartConfig + cfg = WithAuthZRoleProviderFactory("mock", func(_ context.Context, _ authz.ProviderConfig) (authz.RoleProvider, error) { + return noopRoleProvider{}, nil + })(cfg) + + require.NotNil(t, cfg.authzRoleProviderFactories) + require.Contains(t, cfg.authzRoleProviderFactories, "mock") +} diff --git a/service/pkg/server/start.go b/service/pkg/server/start.go index 2efcc47c38..f7c8498ce8 100644 --- a/service/pkg/server/start.go +++ b/service/pkg/server/start.go @@ -165,6 +165,14 @@ func Start(f ...StartOptions) error { cfg.Server.Auth.Policy.Adapter = startConfig.casbinAdapter } + // Set AuthZ role provider overrides + if startConfig.authzRoleProvider != nil { + cfg.Server.Auth.RoleProvider = startConfig.authzRoleProvider + } + if startConfig.authzRoleProviderFactories != nil { + cfg.Server.Auth.RoleProviderFactories = startConfig.authzRoleProviderFactories + } + // Apply additional CORS configuration from programmatic options // These are appended to the YAML config values; deduplication happens in Effective*() methods if len(startConfig.additionalCORSHeaders) > 0 {