chore: add --force-reset-all flag to oidc link repair cli (#26534)

Useful when the issuer is unchanged, but oidc subject claims have
changed.
This commit is contained in:
Steven Masley
2026-06-23 11:37:33 -05:00
committed by GitHub
parent bdf0e417b1
commit 854d280834
7 changed files with 257 additions and 12 deletions
+32 -12
View File
@@ -21,10 +21,11 @@ import (
func (r *RootCmd) newFixOIDCLinksCommand() *serpent.Command {
var (
pgURL string
pgAuth string
issuerURL string
dryRun bool
pgURL string
pgAuth string
issuerURL string
dryRun bool
forceResetAll bool
)
fixOIDCLinksCmd := &serpent.Command{
Use: "fix-oidc-links",
@@ -40,17 +41,29 @@ func (r *RootCmd) newFixOIDCLinksCommand() *serpent.Command {
defer cancel()
issuerURL = strings.TrimSpace(issuerURL)
if issuerURL == "" {
if forceResetAll && issuerURL != "" {
return xerrors.New("--force-reset-all and --issuer-url are mutually exclusive")
}
if !forceResetAll && issuerURL == "" {
return xerrors.Errorf("the --%s flag is required, set it to the OIDC issuer URL (e.g. https://accounts.google.com)", "issuer-url")
}
// Resolve the canonical issuer from OIDC discovery.
cliui.Infof(inv.Stdout, "Resolving OIDC issuer from %q...", issuerURL)
// TODO: The default client might not be configured with the right certs to make this request.
issuer, err := authlink.ResolveIssuer(ctx, http.DefaultClient, issuerURL)
if err != nil {
return xerrors.Errorf("resolve issuer: %w", err)
var issuer string
if forceResetAll {
// Use an unmatchable issuer so the existing analysis shows
// all links as "mismatched" and the reset clears everything.
issuer = authlink.UnmatchableIssuer
} else {
// Resolve the canonical issuer from OIDC discovery.
cliui.Infof(inv.Stdout, "Resolving OIDC issuer from %q...", issuerURL)
// TODO: The default client might not be configured with the right certs to make this request.
resolved, err := authlink.ResolveIssuer(ctx, http.DefaultClient, issuerURL)
if err != nil {
return xerrors.Errorf("resolve issuer: %w", err)
}
issuer = resolved
_, _ = fmt.Fprintf(inv.Stdout, "Resolved OIDC issuer: %q\n\n", issuer)
}
_, _ = fmt.Fprintf(inv.Stdout, "Resolved OIDC issuer: %q\n\n", issuer)
// Connect to the database.
if pgURL == "" {
@@ -59,6 +72,7 @@ func (r *RootCmd) newFixOIDCLinksCommand() *serpent.Command {
sqlDriver := "postgres"
if codersdk.PostgresAuth(pgAuth) == codersdk.PostgresAuthAWSIAMRDS {
var err error
sqlDriver, err = awsiamrds.Register(inv.Context(), sqlDriver)
if err != nil {
return xerrors.Errorf("register aws rds iam auth: %w", err)
@@ -150,6 +164,12 @@ func (r *RootCmd) newFixOIDCLinksCommand() *serpent.Command {
Description: "Print analysis only, do not modify the database.",
Value: serpent.BoolOf(&dryRun),
},
serpent.Option{
Flag: "force-reset-all",
Env: "CODER_FIX_OIDC_LINKS_FORCE_RESET_ALL",
Description: "Reset all OIDC linked IDs, not just those with a mismatched issuer. Mutually exclusive with --issuer-url.",
Value: serpent.BoolOf(&forceResetAll),
},
)
return fixOIDCLinksCmd
+120
View File
@@ -161,6 +161,126 @@ func TestFixOIDCLinks(t *testing.T) {
require.Equal(t, expectedIssuer+"||sub-correct", link.LinkedID)
})
t.Run("ForceResetAll", func(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitMedium)
t.Cleanup(cancel)
connectionURL, err := dbtestutil.Open(t)
require.NoError(t, err)
sqlDB, err := sql.Open("postgres", connectionURL)
require.NoError(t, err)
defer sqlDB.Close()
db := database.New(sqlDB)
// Seed users with different issuers.
user1 := dbgen.User(t, db, database.User{LoginType: database.LoginTypeOIDC})
dbgen.UserLink(t, db, database.UserLink{
UserID: user1.ID,
LoginType: database.LoginTypeOIDC,
LinkedID: "https://accounts.google.com||sub-1",
})
user2 := dbgen.User(t, db, database.User{LoginType: database.LoginTypeOIDC})
dbgen.UserLink(t, db, database.UserLink{
UserID: user2.ID,
LoginType: database.LoginTypeOIDC,
LinkedID: "https://old-issuer.example.com||sub-2",
})
inv, _ := clitest.New(t,
"server", "fix-oidc-links",
"--postgres-url", connectionURL,
"--force-reset-all",
"--yes",
)
stdout := expecter.NewAttachedToInvocation(t, inv)
w := clitest.StartWithWaiter(t, inv)
stdout.ExpectMatch(ctx, "Linked to other issuers:")
stdout.ExpectMatch(ctx, "Reset 2 linked IDs.")
w.RequireSuccess()
// Verify both links were reset.
link, err := db.GetUserLinkByUserIDLoginType(ctx, database.GetUserLinkByUserIDLoginTypeParams{
UserID: user1.ID,
LoginType: database.LoginTypeOIDC,
})
require.NoError(t, err)
require.Equal(t, "", link.LinkedID)
link, err = db.GetUserLinkByUserIDLoginType(ctx, database.GetUserLinkByUserIDLoginTypeParams{
UserID: user2.ID,
LoginType: database.LoginTypeOIDC,
})
require.NoError(t, err)
require.Equal(t, "", link.LinkedID)
})
t.Run("ForceResetAllDryRun", func(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitMedium)
t.Cleanup(cancel)
connectionURL, err := dbtestutil.Open(t)
require.NoError(t, err)
sqlDB, err := sql.Open("postgres", connectionURL)
require.NoError(t, err)
defer sqlDB.Close()
db := database.New(sqlDB)
// Seed users with different issuers.
user1 := dbgen.User(t, db, database.User{LoginType: database.LoginTypeOIDC})
dbgen.UserLink(t, db, database.UserLink{
UserID: user1.ID,
LoginType: database.LoginTypeOIDC,
LinkedID: "https://accounts.google.com||sub-1",
})
user2 := dbgen.User(t, db, database.User{LoginType: database.LoginTypeOIDC})
dbgen.UserLink(t, db, database.UserLink{
UserID: user2.ID,
LoginType: database.LoginTypeOIDC,
LinkedID: "https://old-issuer.example.com||sub-2",
})
inv, _ := clitest.New(t,
"server", "fix-oidc-links",
"--postgres-url", connectionURL,
"--force-reset-all",
"--dry-run",
)
stdout := expecter.NewAttachedToInvocation(t, inv)
w := clitest.StartWithWaiter(t, inv)
stdout.ExpectMatch(ctx, "Total OIDC users:")
stdout.ExpectMatch(ctx, "Linked to other issuers:")
w.RequireSuccess()
// Verify no changes were made.
link, err := db.GetUserLinkByUserIDLoginType(ctx, database.GetUserLinkByUserIDLoginTypeParams{
UserID: user1.ID,
LoginType: database.LoginTypeOIDC,
})
require.NoError(t, err)
require.Equal(t, "https://accounts.google.com||sub-1", link.LinkedID, "dry-run must not modify the database")
link, err = db.GetUserLinkByUserIDLoginType(ctx, database.GetUserLinkByUserIDLoginTypeParams{
UserID: user2.ID,
LoginType: database.LoginTypeOIDC,
})
require.NoError(t, err)
require.Equal(t, "https://old-issuer.example.com||sub-2", link.LinkedID, "dry-run must not modify the database")
})
t.Run("NothingToDo", func(t *testing.T) {
t.Parallel()
@@ -13,6 +13,10 @@ OPTIONS:
-n, --dry-run bool, $CODER_FIX_OIDC_LINKS_DRY_RUN
Print analysis only, do not modify the database.
--force-reset-all bool, $CODER_FIX_OIDC_LINKS_FORCE_RESET_ALL
Reset all OIDC linked IDs, not just those with a mismatched issuer.
Mutually exclusive with --issuer-url.
--issuer-url string, $CODER_OIDC_ISSUER_URL
The OIDC issuer URL. The canonical issuer is resolved via OIDC
discovery.