feat: implement package and cli tool for repairing oidc links (#26418)

This commit is contained in:
Steven Masley
2026-06-16 12:46:10 -07:00
committed by GitHub
parent b71bc31eec
commit 1d03e63f4f
19 changed files with 1116 additions and 1 deletions
+15
View File
@@ -1928,6 +1928,14 @@ func (q *querier) CountInProgressPrebuilds(ctx context.Context) ([]database.Coun
return q.db.CountInProgressPrebuilds(ctx)
}
func (q *querier) CountOIDCLinkedIDsByIssuer(ctx context.Context) ([]database.CountOIDCLinkedIDsByIssuerRow, error) {
// Requires the ability to read all user's personal data.
if err := q.authorizeContext(ctx, policy.ActionReadPersonal, rbac.ResourceUser); err != nil {
return nil, err
}
return q.db.CountOIDCLinkedIDsByIssuer(ctx)
}
func (q *querier) CountPendingNonActivePrebuilds(ctx context.Context) ([]database.CountPendingNonActivePrebuildsRow, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceWorkspace.All()); err != nil {
return nil, err
@@ -6927,6 +6935,13 @@ func (q *querier) UnfavoriteWorkspace(ctx context.Context, id uuid.UUID) error {
return update(q.log, q.auth, fetch, q.db.UnfavoriteWorkspace)(ctx, id)
}
func (q *querier) UnlinkOIDCUsersByIssuerMismatch(ctx context.Context, expectedPrefix string) (int64, error) {
if err := q.authorizeContext(ctx, policy.ActionUpdatePersonal, rbac.ResourceUser); err != nil {
return 0, err
}
return q.db.UnlinkOIDCUsersByIssuerMismatch(ctx, expectedPrefix)
}
func (q *querier) UnpinChatByID(ctx context.Context, id uuid.UUID) error {
chat, err := q.db.GetChatByID(ctx, id)
if err != nil {
+9
View File
@@ -4761,6 +4761,15 @@ func (s *MethodTestSuite) TestSystemFunctions() {
dbm.EXPECT().GetUserLinkByLinkedID(gomock.Any(), l.LinkedID).Return(l, nil).AnyTimes()
check.Args(l.LinkedID).Asserts(rbac.ResourceSystem, policy.ActionRead).Returns(l)
}))
s.Run("CountOIDCLinkedIDsByIssuer", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
dbm.EXPECT().CountOIDCLinkedIDsByIssuer(gomock.Any()).Return([]database.CountOIDCLinkedIDsByIssuerRow{}, nil).AnyTimes()
check.Args().Asserts(rbac.ResourceUser, policy.ActionReadPersonal).Returns([]database.CountOIDCLinkedIDsByIssuerRow{})
}))
s.Run("UnlinkOIDCUsersByIssuerMismatch", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
dbm.EXPECT().UnlinkOIDCUsersByIssuerMismatch(gomock.Any(), "issuer||").Return(int64(0), nil).AnyTimes()
check.Args("issuer||").Asserts(rbac.ResourceUser, policy.ActionUpdatePersonal).Returns(int64(0))
}))
s.Run("GetUserLinkByUserIDLoginType", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
l := testutil.Fake(s.T(), faker, database.UserLink{})
arg := database.GetUserLinkByUserIDLoginTypeParams{UserID: l.UserID, LoginType: l.LoginType}
+16
View File
@@ -370,6 +370,14 @@ func (m queryMetricsStore) CountInProgressPrebuilds(ctx context.Context) ([]data
return r0, r1
}
func (m queryMetricsStore) CountOIDCLinkedIDsByIssuer(ctx context.Context) ([]database.CountOIDCLinkedIDsByIssuerRow, error) {
start := time.Now()
r0, r1 := m.s.CountOIDCLinkedIDsByIssuer(ctx)
m.queryLatencies.WithLabelValues("CountOIDCLinkedIDsByIssuer").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "CountOIDCLinkedIDsByIssuer").Inc()
return r0, r1
}
func (m queryMetricsStore) CountPendingNonActivePrebuilds(ctx context.Context) ([]database.CountPendingNonActivePrebuildsRow, error) {
start := time.Now()
r0, r1 := m.s.CountPendingNonActivePrebuilds(ctx)
@@ -4994,6 +5002,14 @@ func (m queryMetricsStore) UnfavoriteWorkspace(ctx context.Context, id uuid.UUID
return r0
}
func (m queryMetricsStore) UnlinkOIDCUsersByIssuerMismatch(ctx context.Context, expectedPrefix string) (int64, error) {
start := time.Now()
r0, r1 := m.s.UnlinkOIDCUsersByIssuerMismatch(ctx, expectedPrefix)
m.queryLatencies.WithLabelValues("UnlinkOIDCUsersByIssuerMismatch").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UnlinkOIDCUsersByIssuerMismatch").Inc()
return r0, r1
}
func (m queryMetricsStore) UnpinChatByID(ctx context.Context, id uuid.UUID) error {
start := time.Now()
r0 := m.s.UnpinChatByID(ctx, id)
+30
View File
@@ -573,6 +573,21 @@ func (mr *MockStoreMockRecorder) CountInProgressPrebuilds(ctx any) *gomock.Call
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountInProgressPrebuilds", reflect.TypeOf((*MockStore)(nil).CountInProgressPrebuilds), ctx)
}
// CountOIDCLinkedIDsByIssuer mocks base method.
func (m *MockStore) CountOIDCLinkedIDsByIssuer(ctx context.Context) ([]database.CountOIDCLinkedIDsByIssuerRow, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "CountOIDCLinkedIDsByIssuer", ctx)
ret0, _ := ret[0].([]database.CountOIDCLinkedIDsByIssuerRow)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// CountOIDCLinkedIDsByIssuer indicates an expected call of CountOIDCLinkedIDsByIssuer.
func (mr *MockStoreMockRecorder) CountOIDCLinkedIDsByIssuer(ctx any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountOIDCLinkedIDsByIssuer", reflect.TypeOf((*MockStore)(nil).CountOIDCLinkedIDsByIssuer), ctx)
}
// CountPendingNonActivePrebuilds mocks base method.
func (m *MockStore) CountPendingNonActivePrebuilds(ctx context.Context) ([]database.CountPendingNonActivePrebuildsRow, error) {
m.ctrl.T.Helper()
@@ -9417,6 +9432,21 @@ func (mr *MockStoreMockRecorder) UnfavoriteWorkspace(ctx, id any) *gomock.Call {
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UnfavoriteWorkspace", reflect.TypeOf((*MockStore)(nil).UnfavoriteWorkspace), ctx, id)
}
// UnlinkOIDCUsersByIssuerMismatch mocks base method.
func (m *MockStore) UnlinkOIDCUsersByIssuerMismatch(ctx context.Context, expectedPrefix string) (int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "UnlinkOIDCUsersByIssuerMismatch", ctx, expectedPrefix)
ret0, _ := ret[0].(int64)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// UnlinkOIDCUsersByIssuerMismatch indicates an expected call of UnlinkOIDCUsersByIssuerMismatch.
func (mr *MockStoreMockRecorder) UnlinkOIDCUsersByIssuerMismatch(ctx, expectedPrefix any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UnlinkOIDCUsersByIssuerMismatch", reflect.TypeOf((*MockStore)(nil).UnlinkOIDCUsersByIssuerMismatch), ctx, expectedPrefix)
}
// UnpinChatByID mocks base method.
func (m *MockStore) UnpinChatByID(ctx context.Context, id uuid.UUID) error {
m.ctrl.T.Helper()
+9
View File
@@ -107,6 +107,11 @@ type sqlcQuerier interface {
// CountInProgressPrebuilds returns the number of in-progress prebuilds, grouped by preset ID and transition.
// Prebuild considered in-progress if it's in the "pending", "starting", "stopping", or "deleting" state.
CountInProgressPrebuilds(ctx context.Context) ([]CountInProgressPrebuildsRow, error)
// Groups OIDC user links by their issuer prefix (the part before "||" in
// linked_id) and returns a count for each. Empty linked_ids are reported
// with an empty issuer_prefix. Used for analysis before resetting
// mismatched links.
CountOIDCLinkedIDsByIssuer(ctx context.Context) ([]CountOIDCLinkedIDsByIssuerRow, error)
// CountPendingNonActivePrebuilds returns the number of pending prebuilds for non-active template versions
CountPendingNonActivePrebuilds(ctx context.Context) ([]CountPendingNonActivePrebuildsRow, error)
CountUnreadInboxNotificationsByUserID(ctx context.Context, userID uuid.UUID) (int64, error)
@@ -1299,6 +1304,10 @@ type sqlcQuerier interface {
// This will always work regardless of the current state of the template version.
UnarchiveTemplateVersion(ctx context.Context, arg UnarchiveTemplateVersionParams) error
UnfavoriteWorkspace(ctx context.Context, id uuid.UUID) error
// Resets linked_id to '' for OIDC links where the linked_id is non-empty
// and does not begin with the expected issuer prefix. This allows users to
// re-authenticate under a new OIDC provider.
UnlinkOIDCUsersByIssuerMismatch(ctx context.Context, expectedPrefix string) (int64, error)
UnpinChatByID(ctx context.Context, id uuid.UUID) error
UnsetDefaultChatModelConfigs(ctx context.Context) error
UpdateAIBridgeInterceptionEnded(ctx context.Context, arg UpdateAIBridgeInterceptionEndedParams) (AIBridgeInterception, error)
+71
View File
@@ -28866,6 +28866,55 @@ func (q *sqlQuerier) UpsertUserAIProviderKey(ctx context.Context, arg UpsertUser
return i, err
}
const countOIDCLinkedIDsByIssuer = `-- name: CountOIDCLinkedIDsByIssuer :many
SELECT
(CASE
WHEN user_links.linked_id = '' THEN ''
ELSE split_part(user_links.linked_id, '||', 1)
END)::text AS issuer_prefix,
COUNT(*)::int AS count
FROM
user_links
INNER JOIN
users ON user_links.user_id = users.id
WHERE
user_links.login_type = 'oidc'
AND users.deleted = false
GROUP BY issuer_prefix
`
type CountOIDCLinkedIDsByIssuerRow struct {
IssuerPrefix string `db:"issuer_prefix" json:"issuer_prefix"`
Count int32 `db:"count" json:"count"`
}
// Groups OIDC user links by their issuer prefix (the part before "||" in
// linked_id) and returns a count for each. Empty linked_ids are reported
// with an empty issuer_prefix. Used for analysis before resetting
// mismatched links.
func (q *sqlQuerier) CountOIDCLinkedIDsByIssuer(ctx context.Context) ([]CountOIDCLinkedIDsByIssuerRow, error) {
rows, err := q.db.QueryContext(ctx, countOIDCLinkedIDsByIssuer)
if err != nil {
return nil, err
}
defer rows.Close()
var items []CountOIDCLinkedIDsByIssuerRow
for rows.Next() {
var i CountOIDCLinkedIDsByIssuerRow
if err := rows.Scan(&i.IssuerPrefix, &i.Count); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getUserLinkByLinkedID = `-- name: GetUserLinkByLinkedID :one
SELECT
user_links.user_id, user_links.login_type, user_links.linked_id, user_links.oauth_access_token, user_links.oauth_refresh_token, user_links.oauth_expiry, user_links.oauth_access_token_key_id, user_links.oauth_refresh_token_key_id, user_links.claims
@@ -29124,6 +29173,28 @@ func (q *sqlQuerier) OIDCClaimFields(ctx context.Context, organizationID uuid.UU
return items, nil
}
const unlinkOIDCUsersByIssuerMismatch = `-- name: UnlinkOIDCUsersByIssuerMismatch :execrows
UPDATE user_links
SET linked_id = ''
FROM users
WHERE user_links.user_id = users.id
AND user_links.login_type = 'oidc'
AND user_links.linked_id != ''
AND NOT starts_with(user_links.linked_id, $1)
AND users.deleted = false
`
// Resets linked_id to ” for OIDC links where the linked_id is non-empty
// and does not begin with the expected issuer prefix. This allows users to
// re-authenticate under a new OIDC provider.
func (q *sqlQuerier) UnlinkOIDCUsersByIssuerMismatch(ctx context.Context, expectedPrefix string) (int64, error) {
result, err := q.db.ExecContext(ctx, unlinkOIDCUsersByIssuerMismatch, expectedPrefix)
if err != nil {
return 0, err
}
return result.RowsAffected()
}
const updateUserLink = `-- name: UpdateUserLink :one
UPDATE
user_links
+33
View File
@@ -113,3 +113,36 @@ WHERE
ELSE true
END
;
-- name: CountOIDCLinkedIDsByIssuer :many
-- Groups OIDC user links by their issuer prefix (the part before "||" in
-- linked_id) and returns a count for each. Empty linked_ids are reported
-- with an empty issuer_prefix. Used for analysis before resetting
-- mismatched links.
SELECT
(CASE
WHEN user_links.linked_id = '' THEN ''
ELSE split_part(user_links.linked_id, '||', 1)
END)::text AS issuer_prefix,
COUNT(*)::int AS count
FROM
user_links
INNER JOIN
users ON user_links.user_id = users.id
WHERE
user_links.login_type = 'oidc'
AND users.deleted = false
GROUP BY issuer_prefix;
-- name: UnlinkOIDCUsersByIssuerMismatch :execrows
-- Resets linked_id to '' for OIDC links where the linked_id is non-empty
-- and does not begin with the expected issuer prefix. This allows users to
-- re-authenticate under a new OIDC provider.
UPDATE user_links
SET linked_id = ''
FROM users
WHERE user_links.user_id = users.id
AND user_links.login_type = 'oidc'
AND user_links.linked_id != ''
AND NOT starts_with(user_links.linked_id, @expected_prefix)
AND users.deleted = false;