package provisionerdserver import ( "context" "testing" "time" "github.com/google/uuid" "github.com/stretchr/testify/require" "golang.org/x/oauth2" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/database/dbtime" "github.com/coder/coder/v2/testutil" ) func TestShouldRefreshOIDCToken(t *testing.T) { t.Parallel() now := dbtime.Now() testCases := []struct { name string link database.UserLink want bool }{ { name: "NoRefreshToken", link: database.UserLink{OAuthExpiry: now.Add(-time.Hour)}, want: false, }, { name: "ZeroExpiry", link: database.UserLink{OAuthRefreshToken: "refresh"}, want: false, }, { name: "ExpiredBeyondAssumedWindow", link: database.UserLink{ OAuthRefreshToken: "refresh", OAuthExpiry: now.Add(-20 * time.Minute), }, want: true, }, { name: "ExpiredWithinAssumedWindow", link: database.UserLink{ OAuthRefreshToken: "refresh", OAuthExpiry: now.Add(-5 * time.Minute), }, want: false, }, { name: "NotExpired", link: database.UserLink{ OAuthRefreshToken: "refresh", OAuthExpiry: now.Add(time.Hour), }, want: false, }, } for _, tc := range testCases { tc := tc t.Run(tc.name, func(t *testing.T) { t.Parallel() require.Equal(t, tc.want, shouldRefreshOIDCToken(tc.link)) }) } } func TestObtainOIDCAccessToken(t *testing.T) { t.Parallel() ctx := context.Background() t.Run("NoToken", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) _, err := obtainOIDCAccessToken(ctx, testutil.Logger(t), db, nil, uuid.Nil) require.NoError(t, err) }) t.Run("InvalidConfig", func(t *testing.T) { // We still want OIDC to succeed even if exchanging the token fails. t.Parallel() db, _ := dbtestutil.NewDB(t) user := dbgen.User(t, db, database.User{}) dbgen.UserLink(t, db, database.UserLink{ UserID: user.ID, LoginType: database.LoginTypeOIDC, OAuthExpiry: dbtime.Now().Add(-time.Hour), }) _, err := obtainOIDCAccessToken(ctx, testutil.Logger(t), db, &oauth2.Config{}, user.ID) require.NoError(t, err) }) t.Run("MissingLink", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) user := dbgen.User(t, db, database.User{ LoginType: database.LoginTypeOIDC, }) tok, err := obtainOIDCAccessToken(ctx, testutil.Logger(t), db, &oauth2.Config{}, user.ID) require.Empty(t, tok) require.NoError(t, err) }) t.Run("Exchange", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) user := dbgen.User(t, db, database.User{}) dbgen.UserLink(t, db, database.UserLink{ UserID: user.ID, LoginType: database.LoginTypeOIDC, OAuthExpiry: dbtime.Now().Add(-time.Hour), }) _, err := obtainOIDCAccessToken(ctx, testutil.Logger(t), db, &testutil.OAuth2Config{ Token: &oauth2.Token{ AccessToken: "token", }, }, user.ID) require.NoError(t, err) link, err := db.GetUserLinkByUserIDLoginType(ctx, database.GetUserLinkByUserIDLoginTypeParams{ UserID: user.ID, LoginType: database.LoginTypeOIDC, }) require.NoError(t, err) require.Equal(t, "token", link.OAuthAccessToken) }) }