Fix missing cmid scope

Signed-off-by: Bryan Frimin <bryan@probo.com>
This commit is contained in:
Bryan Frimin
2026-06-19 18:49:33 +02:00
parent 8add4713c8
commit 9fd95a0bf9
19 changed files with 611 additions and 646 deletions

View File

@@ -14,20 +14,15 @@
package accessreview package accessreview
import ( import "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/iam"
)
const ( const (
ScopeV1AccessReviewRead coredata.OAuth2Scope = "v1:access-review:read" ScopeV1AccessReviewRead coredata.OAuth2Scope = "v1:access-review:read"
ScopeV1AccessReview coredata.OAuth2Scope = "v1:access-review" ScopeV1AccessReview coredata.OAuth2Scope = "v1:access-review"
) )
// OAuth2ScopeSet returns OAuth2 scope mappings for access-review actions. // OAuth2ScopeMappings maps OAuth2 scopes to access-review actions.
func OAuth2ScopeSet() *iam.ScopeSet { var OAuth2ScopeMappings = map[coredata.OAuth2Scope][]string{
return iam.CreateScopeSet(
map[coredata.OAuth2Scope][]iam.Action{
ScopeV1AccessReviewRead: { ScopeV1AccessReviewRead: {
ActionCampaignGet, ActionCampaignGet,
ActionCampaignList, ActionCampaignList,
@@ -53,6 +48,4 @@ func OAuth2ScopeSet() *iam.ScopeSet {
ActionSourceDelete, ActionSourceDelete,
ActionSourceSync, ActionSourceSync,
}, },
},
)
} }

View File

@@ -14,20 +14,14 @@
package agentrun package agentrun
import ( import "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/iam"
)
const ( const (
ScopeV1AgentRead coredata.OAuth2Scope = "v1:agent:read" ScopeV1AgentRead coredata.OAuth2Scope = "v1:agent:read"
ScopeV1Agent coredata.OAuth2Scope = "v1:agent" ScopeV1Agent coredata.OAuth2Scope = "v1:agent"
) )
// OAuth2ScopeSet returns OAuth2 scope mappings for agent-run actions. var OAuth2ScopeMappings = map[coredata.OAuth2Scope][]string{
func OAuth2ScopeSet() *iam.ScopeSet {
return iam.CreateScopeSet(
map[coredata.OAuth2Scope][]iam.Action{
ScopeV1AgentRead: { ScopeV1AgentRead: {
ActionAgentRunGet, ActionAgentRunGet,
ActionAgentRunList, ActionAgentRunList,
@@ -35,6 +29,4 @@ func OAuth2ScopeSet() *iam.ScopeSet {
ScopeV1Agent: { ScopeV1Agent: {
ActionAgentRunApprove, ActionAgentRunApprove,
}, },
},
)
} }

View File

@@ -29,6 +29,7 @@ import (
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam/oauth2" "go.probo.inc/probo/pkg/iam/oauth2"
"go.probo.inc/probo/pkg/iam/policy" "go.probo.inc/probo/pkg/iam/policy"
"go.probo.inc/probo/pkg/iam/scopeset"
) )
// AuthorizationAttributer is implemented by entities that can provide // AuthorizationAttributer is implemented by entities that can provide
@@ -86,17 +87,17 @@ type Authorizer struct {
pg *pg.Client pg *pg.Client
evaluator *policy.Evaluator evaluator *policy.Evaluator
policySet *PolicySet policySet *PolicySet
oauth2ScopeSet *ScopeSet oauth2ScopeSet *scopeset.ScopeSet
logger *log.Logger logger *log.Logger
} }
// NewAuthorizer creates a new Authorizer instance. // NewAuthorizer creates a new Authorizer instance.
func NewAuthorizer(pgClient *pg.Client, logger *log.Logger) *Authorizer { func NewAuthorizer(pgClient *pg.Client, logger *log.Logger, scopeSet *scopeset.ScopeSet) *Authorizer {
return &Authorizer{ return &Authorizer{
pg: pgClient, pg: pgClient,
evaluator: policy.NewEvaluator(), evaluator: policy.NewEvaluator(),
policySet: NewPolicySet(), policySet: NewPolicySet(),
oauth2ScopeSet: NewScopeSet(), oauth2ScopeSet: scopeSet,
logger: logger, logger: logger,
} }
} }
@@ -106,20 +107,6 @@ func (a *Authorizer) RegisterPolicySet(ps *PolicySet) {
a.policySet.Merge(ps) a.policySet.Merge(ps)
} }
// RegisterScopes merges OAuth2 scope-to-action mappings into the authorizer.
func (a *Authorizer) RegisterScopes(ss *ScopeSet) {
a.oauth2ScopeSet.Merge(ss)
}
// APIScopes returns OAuth2 API scopes advertised in discovery metadata.
func (a *Authorizer) APIScopes() []coredata.OAuth2Scope {
if a.oauth2ScopeSet == nil {
return []coredata.OAuth2Scope{}
}
return a.oauth2ScopeSet.APIScopes()
}
func (a *Authorizer) checkOAuth2Scope( func (a *Authorizer) checkOAuth2Scope(
ctx context.Context, ctx context.Context,
principal gid.GID, principal gid.GID,

View File

@@ -32,6 +32,7 @@ import (
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam" "go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/iam/policy" "go.probo.inc/probo/pkg/iam/policy"
"go.probo.inc/probo/pkg/iam/scopeset"
"go.probo.inc/probo/pkg/mail" "go.probo.inc/probo/pkg/mail"
) )
@@ -216,7 +217,7 @@ func TestAuthorizer_AuthorizeBatch(t *testing.T) {
t.Run("empty input", func(t *testing.T) { t.Run("empty input", func(t *testing.T) {
t.Parallel() t.Parallel()
authorizer := iam.NewAuthorizer(nil, log.NewLogger(log.WithOutput(io.Discard))) authorizer := iam.NewAuthorizer(nil, log.NewLogger(log.WithOutput(io.Discard)), scopeset.New())
_, err := authorizer.AuthorizeBatch( _, err := authorizer.AuthorizeBatch(
context.Background(), context.Background(),
@@ -644,7 +645,7 @@ func TestAuthorizer_AuthorizeMulti(t *testing.T) {
t.Run("rejects empty items", func(t *testing.T) { t.Run("rejects empty items", func(t *testing.T) {
t.Parallel() t.Parallel()
authorizer := iam.NewAuthorizer(nil, log.NewLogger(log.WithOutput(io.Discard))) authorizer := iam.NewAuthorizer(nil, log.NewLogger(log.WithOutput(io.Discard)), scopeset.New())
scope, decisions, err := authorizer.AuthorizeMulti( scope, decisions, err := authorizer.AuthorizeMulti(
context.Background(), context.Background(),
@@ -663,7 +664,7 @@ func TestAuthorizer_AuthorizeMulti(t *testing.T) {
t.Run("rejects unsupported principal type", func(t *testing.T) { t.Run("rejects unsupported principal type", func(t *testing.T) {
t.Parallel() t.Parallel()
authorizer := iam.NewAuthorizer(nil, log.NewLogger(log.WithOutput(io.Discard))) authorizer := iam.NewAuthorizer(nil, log.NewLogger(log.WithOutput(io.Discard)), scopeset.New())
scope, decisions, err := authorizer.AuthorizeMulti( scope, decisions, err := authorizer.AuthorizeMulti(
context.Background(), context.Background(),
@@ -782,7 +783,7 @@ func newTestAuthorizer(client *pg.Client, action string, allowResourceID *gid.GI
} }
func newTestAuthorizerWithStatements(client *pg.Client, statements ...policy.Statement) *iam.Authorizer { func newTestAuthorizerWithStatements(client *pg.Client, statements ...policy.Statement) *iam.Authorizer {
authorizer := iam.NewAuthorizer(client, log.NewLogger(log.WithOutput(io.Discard))) authorizer := iam.NewAuthorizer(client, log.NewLogger(log.WithOutput(io.Discard)), scopeset.New())
authorizer.RegisterPolicySet( authorizer.RegisterPolicySet(
iam.NewPolicySet().AddRolePolicy( iam.NewPolicySet().AddRolePolicy(
string(coredata.MembershipRoleOwner), string(coredata.MembershipRoleOwner),
@@ -794,7 +795,7 @@ func newTestAuthorizerWithStatements(client *pg.Client, statements ...policy.Sta
} }
func newTestAuthorizerWithIdentityScopedStatements(client *pg.Client, statements ...policy.Statement) *iam.Authorizer { func newTestAuthorizerWithIdentityScopedStatements(client *pg.Client, statements ...policy.Statement) *iam.Authorizer {
authorizer := iam.NewAuthorizer(client, log.NewLogger(log.WithOutput(io.Discard))) authorizer := iam.NewAuthorizer(client, log.NewLogger(log.WithOutput(io.Discard)), scopeset.New())
authorizer.RegisterPolicySet( authorizer.RegisterPolicySet(
iam.NewPolicySet().AddIdentityScopedPolicy( iam.NewPolicySet().AddIdentityScopedPolicy(
policy.NewPolicy("batch-authorize-identity-test", "Batch Authorize Identity Test", statements...), policy.NewPolicy("batch-authorize-identity-test", "Batch Authorize Identity Test", statements...),

View File

@@ -30,6 +30,7 @@ import (
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam" "go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/iam/policy" "go.probo.inc/probo/pkg/iam/policy"
"go.probo.inc/probo/pkg/iam/scopeset"
) )
func TestAuthorizer_DecisionLogging(t *testing.T) { func TestAuthorizer_DecisionLogging(t *testing.T) {
@@ -151,7 +152,7 @@ func newTestAuthorizerWithLogger(
statements = append(statements, extraStatements...) statements = append(statements, extraStatements...)
authorizer := iam.NewAuthorizer(client, log.NewLogger(log.WithOutput(logOutput))) authorizer := iam.NewAuthorizer(client, log.NewLogger(log.WithOutput(logOutput)), scopeset.New())
authorizer.RegisterPolicySet( authorizer.RegisterPolicySet(
iam.NewPolicySet().AddRolePolicy( iam.NewPolicySet().AddRolePolicy(
string(coredata.MembershipRoleOwner), string(coredata.MembershipRoleOwner),

View File

@@ -24,6 +24,7 @@ import (
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam/oauth2" "go.probo.inc/probo/pkg/iam/oauth2"
"go.probo.inc/probo/pkg/iam/scopeset"
) )
func TestAuthorizer_checkOAuth2Scope(t *testing.T) { func TestAuthorizer_checkOAuth2Scope(t *testing.T) {
@@ -65,15 +66,14 @@ func TestAuthorizer_checkOAuth2Scope(t *testing.T) {
t.Run("allows when registered scopes authorize the action", func(t *testing.T) { t.Run("allows when registered scopes authorize the action", func(t *testing.T) {
t.Parallel() t.Parallel()
a := NewAuthorizer(nil, nil) scopeSet := scopeset.New().Register(
a.RegisterScopes( map[coredata.OAuth2Scope][]string{
CreateScopeSet(
map[coredata.OAuth2Scope][]Action{
scopeV1OrgRead: {action}, scopeV1OrgRead: {action},
}, },
),
) )
a := NewAuthorizer(nil, nil, scopeSet)
ctx := oauth2.ContextWithAccessToken( ctx := oauth2.ContextWithAccessToken(
context.Background(), context.Background(),
&coredata.OAuth2AccessToken{Scopes: coredata.OAuth2Scopes{scopeV1OrgRead}}, &coredata.OAuth2AccessToken{Scopes: coredata.OAuth2Scopes{scopeV1OrgRead}},

View File

@@ -334,14 +334,6 @@ func (f *cimdFetcher) storeCache(clientIDURL string, doc *ClientMetadataDocument
) )
} }
func (s *Service) ResolveClient(
ctx context.Context,
clientIDRaw string,
redirectURI string,
) (*coredata.OAuth2Client, error) {
return s.resolveClient(ctx, nil, clientIDRaw, redirectURI)
}
func (s *Service) resolveClient( func (s *Service) resolveClient(
ctx context.Context, ctx context.Context,
tx pg.Tx, tx pg.Tx,
@@ -415,7 +407,7 @@ func (s *Service) upsertCIMDClient(
ScopeEmail, ScopeEmail,
ScopeOfflineAccess, ScopeOfflineAccess,
}, },
s.apiScopes, s.scopeSet.APIScopes(),
) )
now := time.Now() now := time.Now()

View File

@@ -33,6 +33,7 @@ import (
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/net" "go.probo.inc/probo/pkg/net"
"go.probo.inc/probo/pkg/page" "go.probo.inc/probo/pkg/page"
"go.probo.inc/probo/pkg/iam/scopeset"
"go.probo.inc/probo/pkg/uri" "go.probo.inc/probo/pkg/uri"
) )
@@ -62,7 +63,7 @@ type (
gc *GarbageCollector gc *GarbageCollector
cimd *cimdFetcher cimd *cimdFetcher
cimdAllowedClientIDs []string cimdAllowedClientIDs []string
apiScopes []coredata.OAuth2Scope scopeSet *scopeset.ScopeSet
accessTokenDuration time.Duration accessTokenDuration time.Duration
refreshTokenDuration time.Duration refreshTokenDuration time.Duration
authorizationCodeDuration time.Duration authorizationCodeDuration time.Duration
@@ -159,9 +160,9 @@ func WithDeviceCodeDuration(d time.Duration) Option {
} }
} }
func WithAPIScopes(scopes []coredata.OAuth2Scope) Option { func WithScopeSet(scopeSet *scopeset.ScopeSet) Option {
return func(s *Service) { return func(s *Service) {
s.apiScopes = scopes s.scopeSet = scopeSet
} }
} }
@@ -1438,6 +1439,8 @@ func (s *Service) Authorize(
return err return err
} }
fmt.Printf("X: %+v\n", client)
if !client.IsRedirectURIAllowed(req.RedirectURI) { if !client.IsRedirectURIAllowed(req.RedirectURI) {
return ErrInvalidRedirectURI return ErrInvalidRedirectURI
} }
@@ -1736,7 +1739,7 @@ func (s *Service) AuthenticateClient(
clientIDRaw string, clientIDRaw string,
clientSecret string, clientSecret string,
) (*coredata.OAuth2Client, error) { ) (*coredata.OAuth2Client, error) {
client, err := s.ResolveClient(ctx, clientIDRaw, "") client, err := s.resolveClient(ctx, nil, clientIDRaw, "")
if err != nil { if err != nil {
return nil, err return nil, err
} }

View File

@@ -22,30 +22,21 @@ import (
"go.probo.inc/probo/pkg/agentrun" "go.probo.inc/probo/pkg/agentrun"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/iam" "go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/iam/scopeset"
"go.probo.inc/probo/pkg/probo" "go.probo.inc/probo/pkg/probo"
) )
func allRegisteredOAuth2ScopeSets() *iam.ScopeSet { func allRegisteredOAuth2ScopeSets() *scopeset.ScopeSet {
return iam.NewScopeSet(). return scopeset.New().
Merge(iam.IAMOAuth2ScopeSet()). Register(iam.IAMOAuth2ScopeMappings).
Merge(probo.OAuth2ScopeSet()). Register(probo.OAuth2ScopeMappings).
Merge(accessreview.OAuth2ScopeSet()). Register(accessreview.OAuth2ScopeMappings).
Merge(agentrun.OAuth2ScopeSet()) Register(agentrun.OAuth2ScopeMappings)
}
func registerAllOAuth2ScopeSets(authorizer *iam.Authorizer) {
authorizer.RegisterScopes(iam.IAMOAuth2ScopeSet())
authorizer.RegisterScopes(probo.OAuth2ScopeSet())
authorizer.RegisterScopes(accessreview.OAuth2ScopeSet())
authorizer.RegisterScopes(agentrun.OAuth2ScopeSet())
} }
func TestRegisteredOAuth2ScopeSets_OrganizationRead(t *testing.T) { func TestRegisteredOAuth2ScopeSets_OrganizationRead(t *testing.T) {
t.Parallel() t.Parallel()
authorizer := iam.NewAuthorizer(nil, nil)
registerAllOAuth2ScopeSets(authorizer)
scopeSet := allRegisteredOAuth2ScopeSets() scopeSet := allRegisteredOAuth2ScopeSets()
tokenScopes := coredata.OAuth2Scopes{probo.ScopeV1OrgRead} tokenScopes := coredata.OAuth2Scopes{probo.ScopeV1OrgRead}
@@ -57,9 +48,6 @@ func TestRegisteredOAuth2ScopeSets_OrganizationRead(t *testing.T) {
func TestRegisteredOAuth2ScopeSets_UnmappedActionDenies(t *testing.T) { func TestRegisteredOAuth2ScopeSets_UnmappedActionDenies(t *testing.T) {
t.Parallel() t.Parallel()
authorizer := iam.NewAuthorizer(nil, nil)
registerAllOAuth2ScopeSets(authorizer)
scopeSet := allRegisteredOAuth2ScopeSets() scopeSet := allRegisteredOAuth2ScopeSets()
tokenScopes := coredata.OAuth2Scopes{ tokenScopes := coredata.OAuth2Scopes{
probo.ScopeV1OrgRead, probo.ScopeV1OrgRead,

View File

@@ -21,10 +21,7 @@ const (
ScopeV1IAM coredata.OAuth2Scope = "v1:iam" ScopeV1IAM coredata.OAuth2Scope = "v1:iam"
) )
// IAMOAuth2ScopeSet returns OAuth2 scope mappings for IAM actions. var IAMOAuth2ScopeMappings = map[coredata.OAuth2Scope][]string{
func IAMOAuth2ScopeSet() *ScopeSet {
return CreateScopeSet(
map[coredata.OAuth2Scope][]Action{
ScopeV1IAMRead: { ScopeV1IAMRead: {
ActionOrganizationGet, ActionOrganizationGet,
ActionOrganizationList, ActionOrganizationList,
@@ -84,6 +81,4 @@ func IAMOAuth2ScopeSet() *ScopeSet {
ActionSCIMBridgeDelete, ActionSCIMBridgeDelete,
ActionOAuth2ConsentApprove, ActionOAuth2ConsentApprove,
}, },
},
)
} }

View File

@@ -20,65 +20,21 @@ import (
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/iam/scopeset"
) )
func TestScopeSet_Allows(t *testing.T) { func TestAuthorizer_UsesOAuth2ScopeSet(t *testing.T) {
t.Parallel() t.Parallel()
const scopeV1OrgRead = coredata.OAuth2Scope("v1:org:read") const scopeV1OrgRead = coredata.OAuth2Scope("v1:org:read")
scopeSet := CreateScopeSet( scopeSet := scopeset.New().Register(
map[coredata.OAuth2Scope][]Action{ map[coredata.OAuth2Scope][]string{
scopeV1OrgRead: {"core:organization:get"}, scopeV1OrgRead: {"core:organization:get"},
}, },
) )
tokenScopes := coredata.OAuth2Scopes{scopeV1OrgRead} authorizer := NewAuthorizer(nil, nil, scopeSet)
assert.True(t, scopeSet.Allows(tokenScopes, "core:organization:get"))
assert.False(t, scopeSet.Allows(tokenScopes, "core:organization:update"))
}
func TestScopeSet_Merge(t *testing.T) {
t.Parallel()
const scopeV1OrgRead = coredata.OAuth2Scope("v1:org:read")
scopeSet := NewScopeSet().
Merge(
CreateScopeSet(
map[coredata.OAuth2Scope][]Action{
scopeV1OrgRead: {"core:organization:get"},
},
),
).
Merge(
CreateScopeSet(
map[coredata.OAuth2Scope][]Action{
scopeV1OrgRead: {"core:organization-context:get"},
},
),
)
tokenScopes := coredata.OAuth2Scopes{scopeV1OrgRead}
assert.True(t, scopeSet.Allows(tokenScopes, "core:organization:get"))
assert.True(t, scopeSet.Allows(tokenScopes, "core:organization-context:get"))
}
func TestAuthorizer_RegisterScopes(t *testing.T) {
t.Parallel()
const scopeV1OrgRead = coredata.OAuth2Scope("v1:org:read")
authorizer := NewAuthorizer(nil, nil)
authorizer.RegisterScopes(
CreateScopeSet(
map[coredata.OAuth2Scope][]Action{
scopeV1OrgRead: {"core:organization:get"},
},
),
)
require.NotNil(t, authorizer.oauth2ScopeSet) require.NotNil(t, authorizer.oauth2ScopeSet)

View File

@@ -12,34 +12,32 @@
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR // OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE. // PERFORMANCE OF THIS SOFTWARE.
package iam package scopeset
import ( import (
"cmp" "cmp"
"maps" "maps"
"slices" "slices"
"sync"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
) )
// ScopeSet holds OAuth2 scope to IAM action mappings. Services create their
// own ScopeSet and register it on the Authorizer at composition time.
type ScopeSet struct { type ScopeSet struct {
scopeActions map[coredata.OAuth2Scope][]Action mu sync.RWMutex
actionScopes map[Action][]coredata.OAuth2Scope scopeActions map[coredata.OAuth2Scope][]string
actionScopes map[string][]coredata.OAuth2Scope
} }
// NewScopeSet creates an empty ScopeSet. func New() *ScopeSet {
func NewScopeSet() *ScopeSet {
return &ScopeSet{ return &ScopeSet{
scopeActions: make(map[coredata.OAuth2Scope][]Action), scopeActions: make(map[coredata.OAuth2Scope][]string),
} }
} }
// CreateScopeSet creates a ScopeSet from scope-to-action mappings. Entries with func (s *ScopeSet) Register(mappings map[coredata.OAuth2Scope][]string) *ScopeSet {
// no actions are skipped. s.mu.Lock()
func CreateScopeSet(mappings map[coredata.OAuth2Scope][]Action) *ScopeSet { defer s.mu.Unlock()
s := NewScopeSet()
for scope, actions := range mappings { for scope, actions := range mappings {
if len(actions) == 0 { if len(actions) == 0 {
@@ -54,24 +52,17 @@ func CreateScopeSet(mappings map[coredata.OAuth2Scope][]Action) *ScopeSet {
return s return s
} }
// Merge combines another ScopeSet into this one.
func (s *ScopeSet) Merge(other *ScopeSet) *ScopeSet {
for scope, actions := range other.scopeActions {
s.scopeActions[scope] = append(s.scopeActions[scope], actions...)
}
s.rebuildActionScopes()
return s
}
// APIScopes returns every registered OAuth2 API scope in this set.
func (s *ScopeSet) APIScopes() []coredata.OAuth2Scope { func (s *ScopeSet) APIScopes() []coredata.OAuth2Scope {
s.mu.RLock()
defer s.mu.RUnlock()
return sortedScopes(slices.Collect(maps.Keys(s.scopeActions))) return sortedScopes(slices.Collect(maps.Keys(s.scopeActions)))
} }
// Allows reports whether tokenScopes authorize action. func (s *ScopeSet) Allows(tokenScopes coredata.OAuth2Scopes, action string) bool {
func (s *ScopeSet) Allows(tokenScopes coredata.OAuth2Scopes, action Action) bool { s.mu.RLock()
defer s.mu.RUnlock()
grantingScopes, ok := s.actionScopes[action] grantingScopes, ok := s.actionScopes[action]
if !ok { if !ok {
return false return false
@@ -81,7 +72,7 @@ func (s *ScopeSet) Allows(tokenScopes coredata.OAuth2Scopes, action Action) bool
} }
func (s *ScopeSet) rebuildActionScopes() { func (s *ScopeSet) rebuildActionScopes() {
actionScopes := make(map[Action][]coredata.OAuth2Scope, len(s.scopeActions)*4) actionScopes := make(map[string][]coredata.OAuth2Scope, len(s.scopeActions)*4)
for scope, actions := range s.scopeActions { for scope, actions := range s.scopeActions {
for _, action := range actions { for _, action := range actions {
@@ -94,9 +85,12 @@ func (s *ScopeSet) rebuildActionScopes() {
func sortedScopes(scopes []coredata.OAuth2Scope) []coredata.OAuth2Scope { func sortedScopes(scopes []coredata.OAuth2Scope) []coredata.OAuth2Scope {
sorted := slices.Clone(scopes) sorted := slices.Clone(scopes)
slices.SortFunc(sorted, func(a, b coredata.OAuth2Scope) int { slices.SortFunc(
sorted,
func(a, b coredata.OAuth2Scope) int {
return cmp.Compare(string(a), string(b)) return cmp.Compare(string(a), string(b))
}) },
)
return sorted return sorted
} }

View File

@@ -0,0 +1,63 @@
// Copyright (c) 2026 Probo Inc <hello@probo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package scopeset_test
import (
"testing"
"github.com/stretchr/testify/assert"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/iam/scopeset"
)
func TestScopeSet_Allows(t *testing.T) {
t.Parallel()
const scopeV1OrgRead = coredata.OAuth2Scope("v1:org:read")
scopeSet := scopeset.New().Register(
map[coredata.OAuth2Scope][]string{
scopeV1OrgRead: {"core:organization:get"},
},
)
tokenScopes := coredata.OAuth2Scopes{scopeV1OrgRead}
assert.True(t, scopeSet.Allows(tokenScopes, "core:organization:get"))
assert.False(t, scopeSet.Allows(tokenScopes, "core:organization:update"))
}
func TestScopeSet_Register(t *testing.T) {
t.Parallel()
const scopeV1OrgRead = coredata.OAuth2Scope("v1:org:read")
scopeSet := scopeset.New().
Register(
map[coredata.OAuth2Scope][]string{
scopeV1OrgRead: {"core:organization:get"},
},
).
Register(
map[coredata.OAuth2Scope][]string{
scopeV1OrgRead: {"core:organization-context:get"},
},
)
tokenScopes := coredata.OAuth2Scopes{scopeV1OrgRead}
assert.True(t, scopeSet.Allows(tokenScopes, "core:organization:get"))
assert.True(t, scopeSet.Allows(tokenScopes, "core:organization-context:get"))
}

View File

@@ -37,6 +37,7 @@ import (
"go.probo.inc/probo/pkg/iam/oidc" "go.probo.inc/probo/pkg/iam/oidc"
"go.probo.inc/probo/pkg/iam/saml" "go.probo.inc/probo/pkg/iam/saml"
"go.probo.inc/probo/pkg/iam/scim" "go.probo.inc/probo/pkg/iam/scim"
"go.probo.inc/probo/pkg/iam/scopeset"
"go.probo.inc/probo/pkg/uri" "go.probo.inc/probo/pkg/uri"
) )
@@ -69,6 +70,7 @@ type (
APIKeyService *APIKeyService APIKeyService *APIKeyService
OAuth2ServerService *oauth2.Service OAuth2ServerService *oauth2.Service
Authorizer *Authorizer Authorizer *Authorizer
OAuth2ScopeSet *scopeset.ScopeSet
samlDomainVerifier *SAMLDomainVerifier samlDomainVerifier *SAMLDomainVerifier
} }
@@ -157,12 +159,15 @@ func NewService(
svc.AuthService = NewAuthService(svc) svc.AuthService = NewAuthService(svc)
svc.APIKeyService = NewAPIKeyService(svc) svc.APIKeyService = NewAPIKeyService(svc)
svc.OAuth2ScopeSet = scopeset.New()
svc.OAuth2ScopeSet.Register(IAMOAuth2ScopeMappings)
svc.Authorizer = NewAuthorizer( svc.Authorizer = NewAuthorizer(
pgClient, pgClient,
cfg.Logger.Named("authorizer"), cfg.Logger.Named("authorizer"),
svc.OAuth2ScopeSet,
) )
svc.Authorizer.RegisterPolicySet(IAMPolicySet()) svc.Authorizer.RegisterPolicySet(IAMPolicySet())
svc.Authorizer.RegisterScopes(IAMOAuth2ScopeSet())
samlService, err := saml.NewService(svc.pg, svc.baseURL, svc.certificate, svc.privateKey, cfg.Logger) samlService, err := saml.NewService(svc.pg, svc.baseURL, svc.certificate, svc.privateKey, cfg.Logger)
if err != nil { if err != nil {
@@ -201,7 +206,7 @@ func NewService(
uri.URI(cfg.BaseURL.String()), uri.URI(cfg.BaseURL.String()),
cfg.Logger.Named("oauth2"), cfg.Logger.Named("oauth2"),
append( append(
[]oauth2.Option{oauth2.WithAPIScopes(svc.Authorizer.APIScopes())}, []oauth2.Option{oauth2.WithScopeSet(svc.OAuth2ScopeSet)},
cfg.OAuth2ServerOptions..., cfg.OAuth2ServerOptions...,
)..., )...,
) )
@@ -219,12 +224,12 @@ func NewService(
// OAuth2ServerMetadata returns the OIDC discovery document. // OAuth2ServerMetadata returns the OIDC discovery document.
func (s *Service) OAuth2ServerMetadata(endpoints oauth2.Endpoints) *oauth2.ServerMetadata { func (s *Service) OAuth2ServerMetadata(endpoints oauth2.Endpoints) *oauth2.ServerMetadata {
return oauth2.NewMetadata(uri.URI(s.baseURL), endpoints, s.Authorizer.APIScopes()) return oauth2.NewMetadata(uri.URI(s.baseURL), endpoints, s.OAuth2ScopeSet.APIScopes())
} }
// OAuth2ProtectedResourceMetadata returns the RFC 9728 protected resource metadata document. // OAuth2ProtectedResourceMetadata returns the RFC 9728 protected resource metadata document.
func (s *Service) OAuth2ProtectedResourceMetadata(resource uri.URI) *oauth2.ProtectedResourceMetadata { func (s *Service) OAuth2ProtectedResourceMetadata(resource uri.URI) *oauth2.ProtectedResourceMetadata {
return oauth2.NewProtectedResourceMetadata(resource, resource, s.Authorizer.APIScopes()) return oauth2.NewProtectedResourceMetadata(resource, resource, s.OAuth2ScopeSet.APIScopes())
} }
func (s *Service) IsSignUpEnabled() bool { func (s *Service) IsSignUpEnabled() bool {

View File

@@ -16,7 +16,6 @@ package probo
import ( import (
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/iam"
) )
const ( const (
@@ -66,10 +65,8 @@ const (
ScopeV1Webhook coredata.OAuth2Scope = "v1:webhook" ScopeV1Webhook coredata.OAuth2Scope = "v1:webhook"
) )
// OAuth2ScopeSet returns OAuth2 scope mappings for core probo actions. // OAuth2ScopeMappings maps OAuth2 scopes to core probo actions.
func OAuth2ScopeSet() *iam.ScopeSet { var OAuth2ScopeMappings = map[coredata.OAuth2Scope][]string{
return iam.CreateScopeSet(
map[coredata.OAuth2Scope][]iam.Action{
ScopeV1AssetRead: { ScopeV1AssetRead: {
ActionAssetGet, ActionAssetGet,
@@ -443,6 +440,4 @@ func OAuth2ScopeSet() *iam.ScopeSet {
ActionWebhookSubscriptionUpdate, ActionWebhookSubscriptionUpdate,
ActionWebhookSubscriptionDelete, ActionWebhookSubscriptionDelete,
}, },
},
)
} }

View File

@@ -149,7 +149,7 @@ func NewService(
} }
iamService.Authorizer.RegisterPolicySet(ProboPolicySet()) iamService.Authorizer.RegisterPolicySet(ProboPolicySet())
iamService.Authorizer.RegisterScopes(OAuth2ScopeSet()) iamService.OAuth2ScopeSet.Register(OAuth2ScopeMappings)
svc := &Service{ svc := &Service{
pg: pgClient, pg: pgClient,

View File

@@ -601,8 +601,8 @@ func (impl *Implm) Run(
iamService.Authorizer.RegisterPolicySet(agentrun.PolicySet()) iamService.Authorizer.RegisterPolicySet(agentrun.PolicySet())
iamService.Authorizer.RegisterPolicySet(accessreview.PolicySet()) iamService.Authorizer.RegisterPolicySet(accessreview.PolicySet())
iamService.Authorizer.RegisterScopes(agentrun.OAuth2ScopeSet()) iamService.OAuth2ScopeSet.Register(agentrun.OAuth2ScopeMappings)
iamService.Authorizer.RegisterScopes(accessreview.OAuth2ScopeSet()) iamService.OAuth2ScopeSet.Register(accessreview.OAuth2ScopeMappings)
thirdPartyService := thirdparty.NewService(pgClient, fileManagerService, thirdPartyVetter) thirdPartyService := thirdparty.NewService(pgClient, fileManagerService, thirdPartyVetter)
riskManagementService := riskmanagement.NewService(pgClient) riskManagementService := riskmanagement.NewService(pgClient)

View File

@@ -268,7 +268,7 @@ func (r *queryResolver) SignUpEnabled(ctx context.Context) (bool, error) {
// Oauth2ScopesSupported is the resolver for the oauth2ScopesSupported field. // Oauth2ScopesSupported is the resolver for the oauth2ScopesSupported field.
func (r *queryResolver) Oauth2ScopesSupported(ctx context.Context) ([]string, error) { func (r *queryResolver) Oauth2ScopesSupported(ctx context.Context) ([]string, error) {
apiScopes := r.iam.Authorizer.APIScopes() apiScopes := r.iam.OAuth2ScopeSet.APIScopes()
scopes := make([]string, len(apiScopes)) scopes := make([]string, len(apiScopes))
for i, scope := range apiScopes { for i, scope := range apiScopes {

View File

@@ -40,7 +40,7 @@ func (r *mutationResolver) CreateOAuth2AccessToken(ctx context.Context, input ty
Name: strings.TrimSpace(input.Name), Name: strings.TrimSpace(input.Name),
ExpiresAt: input.ExpiresAt, ExpiresAt: input.ExpiresAt,
Scopes: scopes, Scopes: scopes,
AllowedAPIScopes: r.iam.Authorizer.APIScopes(), AllowedAPIScopes: r.iam.OAuth2ScopeSet.APIScopes(),
}, },
) )
if err != nil { if err != nil {