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
4 changes: 4 additions & 0 deletions opentdf-dev.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
24 changes: 20 additions & 4 deletions service/internal/auth/authn.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down
97 changes: 28 additions & 69 deletions service/internal/auth/casbin.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package auth

import (
"context"
"errors"
"fmt"
"log/slog"
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -112,23 +117,35 @@ 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,
Policy: c.Csv,
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 {
Expand All @@ -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)
Expand All @@ -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
}
Loading
Loading