mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
chore: add query to fetch top level idp claim fields (#15525)
Adds an api endpoint to grab all available sync field options for IDP sync. This is for autocomplete on idp sync forms. This is required for organization admins to have some insight into the claim fields available when configuring group/role sync.
This commit is contained in:
@@ -3283,6 +3283,18 @@ func (q *querier) ListWorkspaceAgentPortShares(ctx context.Context, workspaceID
|
||||
return q.db.ListWorkspaceAgentPortShares(ctx, workspaceID)
|
||||
}
|
||||
|
||||
func (q *querier) OIDCClaimFields(ctx context.Context, organizationID uuid.UUID) ([]string, error) {
|
||||
resource := rbac.ResourceIdpsyncSettings
|
||||
if organizationID != uuid.Nil {
|
||||
resource = resource.InOrg(organizationID)
|
||||
}
|
||||
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, resource); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.OIDCClaimFields(ctx, organizationID)
|
||||
}
|
||||
|
||||
func (q *querier) OrganizationMembers(ctx context.Context, arg database.OrganizationMembersParams) ([]database.OrganizationMembersRow, error) {
|
||||
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.OrganizationMembers)(ctx, arg)
|
||||
}
|
||||
|
||||
@@ -626,6 +626,13 @@ func (s *MethodTestSuite) TestLicense() {
|
||||
}
|
||||
|
||||
func (s *MethodTestSuite) TestOrganization() {
|
||||
s.Run("Deployment/OIDCClaimFields", s.Subtest(func(db database.Store, check *expects) {
|
||||
check.Args(uuid.Nil).Asserts(rbac.ResourceIdpsyncSettings, policy.ActionRead).Returns([]string{})
|
||||
}))
|
||||
s.Run("Organization/OIDCClaimFields", s.Subtest(func(db database.Store, check *expects) {
|
||||
id := uuid.New()
|
||||
check.Args(id).Asserts(rbac.ResourceIdpsyncSettings.InOrg(id), policy.ActionRead).Returns([]string{})
|
||||
}))
|
||||
s.Run("ByOrganization/GetGroups", s.Subtest(func(db database.Store, check *expects) {
|
||||
o := dbgen.Organization(s.T(), db, database.Organization{})
|
||||
a := dbgen.Group(s.T(), db, database.Group{OrganizationID: o.ID})
|
||||
|
||||
@@ -8409,6 +8409,35 @@ func (q *FakeQuerier) ListWorkspaceAgentPortShares(_ context.Context, workspaceI
|
||||
return shares, nil
|
||||
}
|
||||
|
||||
func (q *FakeQuerier) OIDCClaimFields(_ context.Context, organizationID uuid.UUID) ([]string, error) {
|
||||
orgMembers := q.getOrganizationMemberNoLock(organizationID)
|
||||
|
||||
var fields []string
|
||||
for _, link := range q.userLinks {
|
||||
if organizationID != uuid.Nil {
|
||||
inOrg := slices.ContainsFunc(orgMembers, func(organizationMember database.OrganizationMember) bool {
|
||||
return organizationMember.UserID == link.UserID
|
||||
})
|
||||
if !inOrg {
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
if link.LoginType != database.LoginTypeOIDC {
|
||||
continue
|
||||
}
|
||||
|
||||
for k := range link.Claims.IDTokenClaims {
|
||||
fields = append(fields, k)
|
||||
}
|
||||
for k := range link.Claims.UserInfoClaims {
|
||||
fields = append(fields, k)
|
||||
}
|
||||
}
|
||||
|
||||
return slice.Unique(fields), nil
|
||||
}
|
||||
|
||||
func (q *FakeQuerier) OrganizationMembers(_ context.Context, arg database.OrganizationMembersParams) ([]database.OrganizationMembersRow, error) {
|
||||
if err := validateDatabaseType(arg); err != nil {
|
||||
return []database.OrganizationMembersRow{}, err
|
||||
|
||||
@@ -2058,6 +2058,13 @@ func (m queryMetricsStore) ListWorkspaceAgentPortShares(ctx context.Context, wor
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) OIDCClaimFields(ctx context.Context, organizationID uuid.UUID) ([]string, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.OIDCClaimFields(ctx, organizationID)
|
||||
m.queryLatencies.WithLabelValues("OIDCClaimFields").Observe(time.Since(start).Seconds())
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) OrganizationMembers(ctx context.Context, arg database.OrganizationMembersParams) ([]database.OrganizationMembersRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.OrganizationMembers(ctx, arg)
|
||||
|
||||
@@ -4359,6 +4359,21 @@ func (mr *MockStoreMockRecorder) ListWorkspaceAgentPortShares(arg0, arg1 any) *g
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListWorkspaceAgentPortShares", reflect.TypeOf((*MockStore)(nil).ListWorkspaceAgentPortShares), arg0, arg1)
|
||||
}
|
||||
|
||||
// OIDCClaimFields mocks base method.
|
||||
func (m *MockStore) OIDCClaimFields(arg0 context.Context, arg1 uuid.UUID) ([]string, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "OIDCClaimFields", arg0, arg1)
|
||||
ret0, _ := ret[0].([]string)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// OIDCClaimFields indicates an expected call of OIDCClaimFields.
|
||||
func (mr *MockStoreMockRecorder) OIDCClaimFields(arg0, arg1 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "OIDCClaimFields", reflect.TypeOf((*MockStore)(nil).OIDCClaimFields), arg0, arg1)
|
||||
}
|
||||
|
||||
// OrganizationMembers mocks base method.
|
||||
func (m *MockStore) OrganizationMembers(arg0 context.Context, arg1 database.OrganizationMembersParams) ([]database.OrganizationMembersRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -3,6 +3,7 @@ package database
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
@@ -527,3 +528,9 @@ func insertAuthorizedFilter(query string, replaceWith string) (string, error) {
|
||||
filtered := strings.Replace(query, authorizedQueryPlaceholder, replaceWith, 1)
|
||||
return filtered, nil
|
||||
}
|
||||
|
||||
// UpdateUserLinkRawJSON is a custom query for unit testing. Do not ever expose this
|
||||
func (q *sqlQuerier) UpdateUserLinkRawJSON(ctx context.Context, userID uuid.UUID, data json.RawMessage) error {
|
||||
_, err := q.sdb.ExecContext(ctx, "UPDATE user_links SET claims = $2 WHERE user_id = $1", userID, data)
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,219 @@
|
||||
package database_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbfake"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
type extraKeys struct {
|
||||
database.UserLinkClaims
|
||||
Foo string `json:"foo"`
|
||||
}
|
||||
|
||||
func TestOIDCClaims(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
toJSON := func(a any) json.RawMessage {
|
||||
b, _ := json.Marshal(a)
|
||||
return b
|
||||
}
|
||||
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
g := userGenerator{t: t, db: db}
|
||||
|
||||
// https://en.wikipedia.org/wiki/Alice_and_Bob#Cast_of_characters
|
||||
alice := g.withLink(database.LoginTypeOIDC, toJSON(extraKeys{
|
||||
UserLinkClaims: database.UserLinkClaims{
|
||||
IDTokenClaims: map[string]interface{}{
|
||||
"sub": "alice",
|
||||
"alice-id": "from-bob",
|
||||
},
|
||||
UserInfoClaims: nil,
|
||||
MergedClaims: map[string]interface{}{
|
||||
"sub": "alice",
|
||||
"alice-id": "from-bob",
|
||||
},
|
||||
},
|
||||
// Always should be a no-op
|
||||
Foo: "bar",
|
||||
}))
|
||||
bob := g.withLink(database.LoginTypeOIDC, toJSON(database.UserLinkClaims{
|
||||
IDTokenClaims: map[string]interface{}{
|
||||
"sub": "bob",
|
||||
"bob-id": "from-bob",
|
||||
"array": []string{
|
||||
"a", "b", "c",
|
||||
},
|
||||
"map": map[string]interface{}{
|
||||
"key": "value",
|
||||
"foo": "bar",
|
||||
},
|
||||
"nil": nil,
|
||||
},
|
||||
UserInfoClaims: map[string]interface{}{
|
||||
"sub": "bob",
|
||||
"bob-info": []string{},
|
||||
"number": 42,
|
||||
},
|
||||
MergedClaims: map[string]interface{}{
|
||||
"sub": "bob",
|
||||
"bob-info": []string{},
|
||||
"number": 42,
|
||||
"bob-id": "from-bob",
|
||||
"array": []string{
|
||||
"a", "b", "c",
|
||||
},
|
||||
"map": map[string]interface{}{
|
||||
"key": "value",
|
||||
"foo": "bar",
|
||||
},
|
||||
"nil": nil,
|
||||
},
|
||||
}))
|
||||
charlie := g.withLink(database.LoginTypeOIDC, toJSON(database.UserLinkClaims{
|
||||
IDTokenClaims: map[string]interface{}{
|
||||
"sub": "charlie",
|
||||
"charlie-id": "charlie",
|
||||
},
|
||||
UserInfoClaims: map[string]interface{}{
|
||||
"sub": "charlie",
|
||||
"charlie-info": "charlie",
|
||||
},
|
||||
MergedClaims: map[string]interface{}{
|
||||
"sub": "charlie",
|
||||
"charlie-id": "charlie",
|
||||
"charlie-info": "charlie",
|
||||
},
|
||||
}))
|
||||
|
||||
// users that just try to cause problems, but should not affect the output of
|
||||
// queries.
|
||||
problematics := []database.User{
|
||||
g.withLink(database.LoginTypeOIDC, toJSON(database.UserLinkClaims{})), // null claims
|
||||
g.withLink(database.LoginTypeOIDC, []byte(`{}`)), // empty claims
|
||||
g.withLink(database.LoginTypeOIDC, []byte(`{"foo": "bar"}`)), // random keys
|
||||
g.noLink(database.LoginTypeOIDC), // no link
|
||||
|
||||
g.withLink(database.LoginTypeGithub, toJSON(database.UserLinkClaims{
|
||||
IDTokenClaims: map[string]interface{}{
|
||||
"not": "allowed",
|
||||
},
|
||||
UserInfoClaims: map[string]interface{}{
|
||||
"do-not": "look",
|
||||
},
|
||||
MergedClaims: map[string]interface{}{
|
||||
"not": "allowed",
|
||||
"do-not": "look",
|
||||
},
|
||||
})), // github should be omitted
|
||||
|
||||
// extra random users
|
||||
g.noLink(database.LoginTypeGithub),
|
||||
g.noLink(database.LoginTypePassword),
|
||||
}
|
||||
|
||||
// Insert some orgs, users, and links
|
||||
orgA := dbfake.Organization(t, db).Members(
|
||||
append(problematics,
|
||||
alice,
|
||||
bob,
|
||||
)...,
|
||||
).Do()
|
||||
orgB := dbfake.Organization(t, db).Members(
|
||||
append(problematics,
|
||||
bob,
|
||||
charlie,
|
||||
)...,
|
||||
).Do()
|
||||
orgC := dbfake.Organization(t, db).Members().Do()
|
||||
|
||||
// Verify the OIDC claim fields
|
||||
always := []string{"array", "map", "nil", "number"}
|
||||
expectA := append([]string{"sub", "alice-id", "bob-id", "bob-info"}, always...)
|
||||
expectB := append([]string{"sub", "bob-id", "bob-info", "charlie-id", "charlie-info"}, always...)
|
||||
requireClaims(t, db, orgA.Org.ID, expectA)
|
||||
requireClaims(t, db, orgB.Org.ID, expectB)
|
||||
requireClaims(t, db, orgC.Org.ID, []string{})
|
||||
requireClaims(t, db, uuid.Nil, slice.Unique(append(expectA, expectB...)))
|
||||
}
|
||||
|
||||
func requireClaims(t *testing.T, db database.Store, orgID uuid.UUID, want []string) {
|
||||
t.Helper()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
got, err := db.OIDCClaimFields(ctx, orgID)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.ElementsMatch(t, want, got)
|
||||
}
|
||||
|
||||
type userGenerator struct {
|
||||
t *testing.T
|
||||
db database.Store
|
||||
}
|
||||
|
||||
func (g userGenerator) noLink(lt database.LoginType) database.User {
|
||||
t := g.t
|
||||
db := g.db
|
||||
|
||||
t.Helper()
|
||||
|
||||
u := dbgen.User(t, db, database.User{
|
||||
LoginType: lt,
|
||||
})
|
||||
return u
|
||||
}
|
||||
|
||||
func (g userGenerator) withLink(lt database.LoginType, rawJSON json.RawMessage) database.User {
|
||||
t := g.t
|
||||
db := g.db
|
||||
|
||||
user := g.noLink(lt)
|
||||
|
||||
link := dbgen.UserLink(t, db, database.UserLink{
|
||||
UserID: user.ID,
|
||||
LoginType: lt,
|
||||
})
|
||||
|
||||
if sql, ok := db.(rawUpdater); ok {
|
||||
// The only way to put arbitrary json into the db for testing edge cases.
|
||||
// Making this a public API would be a mistake.
|
||||
err := sql.UpdateUserLinkRawJSON(context.Background(), user.ID, rawJSON)
|
||||
require.NoError(t, err)
|
||||
} else {
|
||||
// no need to test the json key logic in dbmem. Everything is type safe.
|
||||
var claims database.UserLinkClaims
|
||||
err := json.Unmarshal(rawJSON, &claims)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = db.UpdateUserLink(context.Background(), database.UpdateUserLinkParams{
|
||||
OAuthAccessToken: link.OAuthAccessToken,
|
||||
OAuthAccessTokenKeyID: link.OAuthAccessTokenKeyID,
|
||||
OAuthRefreshToken: link.OAuthRefreshToken,
|
||||
OAuthRefreshTokenKeyID: link.OAuthRefreshTokenKeyID,
|
||||
OAuthExpiry: link.OAuthExpiry,
|
||||
UserID: link.UserID,
|
||||
LoginType: link.LoginType,
|
||||
// The new claims
|
||||
Claims: claims,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
return user
|
||||
}
|
||||
|
||||
type rawUpdater interface {
|
||||
UpdateUserLinkRawJSON(ctx context.Context, userID uuid.UUID, data json.RawMessage) error
|
||||
}
|
||||
@@ -413,6 +413,9 @@ type sqlcQuerier interface {
|
||||
ListProvisionerKeysByOrganization(ctx context.Context, organizationID uuid.UUID) ([]ProvisionerKey, error)
|
||||
ListProvisionerKeysByOrganizationExcludeReserved(ctx context.Context, organizationID uuid.UUID) ([]ProvisionerKey, error)
|
||||
ListWorkspaceAgentPortShares(ctx context.Context, workspaceID uuid.UUID) ([]WorkspaceAgentPortShare, error)
|
||||
// OIDCClaimFields returns a list of distinct keys in the the merged_claims fields.
|
||||
// This query is used to generate the list of available sync fields for idp sync settings.
|
||||
OIDCClaimFields(ctx context.Context, organizationID uuid.UUID) ([]string, error)
|
||||
// Arguments are optional with uuid.Nil to ignore.
|
||||
// - Use just 'organization_id' to get all members of an org
|
||||
// - Use just 'user_id' to get all orgs a user is a member of
|
||||
|
||||
@@ -9846,6 +9846,49 @@ func (q *sqlQuerier) InsertUserLink(ctx context.Context, arg InsertUserLinkParam
|
||||
return i, err
|
||||
}
|
||||
|
||||
const oIDCClaimFields = `-- name: OIDCClaimFields :many
|
||||
SELECT
|
||||
DISTINCT jsonb_object_keys(claims->'merged_claims')
|
||||
FROM
|
||||
user_links
|
||||
WHERE
|
||||
-- Only return rows where the top level key exists
|
||||
claims ? 'merged_claims' AND
|
||||
-- 'null' is the default value for the id_token_claims field
|
||||
-- jsonb 'null' is not the same as SQL NULL. Strip these out.
|
||||
jsonb_typeof(claims->'merged_claims') != 'null' AND
|
||||
login_type = 'oidc'
|
||||
AND CASE WHEN $1 :: uuid != '00000000-0000-0000-0000-000000000000'::uuid THEN
|
||||
user_links.user_id = ANY(SELECT organization_members.user_id FROM organization_members WHERE organization_id = $1)
|
||||
ELSE true
|
||||
END
|
||||
`
|
||||
|
||||
// OIDCClaimFields returns a list of distinct keys in the the merged_claims fields.
|
||||
// This query is used to generate the list of available sync fields for idp sync settings.
|
||||
func (q *sqlQuerier) OIDCClaimFields(ctx context.Context, organizationID uuid.UUID) ([]string, error) {
|
||||
rows, err := q.db.QueryContext(ctx, oIDCClaimFields, organizationID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []string
|
||||
for rows.Next() {
|
||||
var jsonb_object_keys string
|
||||
if err := rows.Scan(&jsonb_object_keys); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, jsonb_object_keys)
|
||||
}
|
||||
if err := rows.Close(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const updateUserLink = `-- name: UpdateUserLink :one
|
||||
UPDATE
|
||||
user_links
|
||||
|
||||
@@ -57,3 +57,24 @@ SET
|
||||
claims = $6
|
||||
WHERE
|
||||
user_id = $7 AND login_type = $8 RETURNING *;
|
||||
|
||||
|
||||
-- name: OIDCClaimFields :many
|
||||
-- OIDCClaimFields returns a list of distinct keys in the the merged_claims fields.
|
||||
-- This query is used to generate the list of available sync fields for idp sync settings.
|
||||
SELECT
|
||||
DISTINCT jsonb_object_keys(claims->'merged_claims')
|
||||
FROM
|
||||
user_links
|
||||
WHERE
|
||||
-- Only return rows where the top level key exists
|
||||
claims ? 'merged_claims' AND
|
||||
-- 'null' is the default value for the id_token_claims field
|
||||
-- jsonb 'null' is not the same as SQL NULL. Strip these out.
|
||||
jsonb_typeof(claims->'merged_claims') != 'null' AND
|
||||
login_type = 'oidc'
|
||||
AND CASE WHEN @organization_id :: uuid != '00000000-0000-0000-0000-000000000000'::uuid THEN
|
||||
user_links.user_id = ANY(SELECT organization_members.user_id FROM organization_members WHERE organization_id = @organization_id)
|
||||
ELSE true
|
||||
END
|
||||
;
|
||||
|
||||
Reference in New Issue
Block a user