[client] Pass the login hint to GetOAuthFlow at construction

GetOAuthFlow was the only flow factory without a hint parameter, which
forced its callers to apply the hint afterwards through a local setter
interface and a type assertion. Give it the same constructor-style hint
as NewOAuthFlow and set the hint on the concrete flows before they are
handed out as the interface, so a flow is always complete when built
and the caller-side ordering constraint disappears.

An empty hint is a valid value meaning the IdP chooses the account, so
the flows set it unconditionally.
This commit is contained in:
Zoltan Papp
2026-08-14 23:09:43 +02:00
parent c28cf2fa61
commit ee3aeadf8f
5 changed files with 35 additions and 41 deletions
+11 -24
View File
@@ -191,20 +191,13 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
return nil return nil
} }
// loginHintSetter is implemented by both concrete flows (PKCE and device code)
// but absent from the OAuthFlow interface, hence the assertion below — the same
// way internal/auth wires it in authenticateWithPKCEFlow.
type loginHintSetter interface {
SetLoginHint(hint string)
}
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) { func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) {
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV) oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV, profileLoginHint(a.cfgPath))
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to get OAuth flow: %v", err) return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
} }
return runOAuthFlow(a.ctx, oAuthFlow, profileLoginHint(a.cfgPath), urlOpener, nil) return runOAuthFlow(a.ctx, oAuthFlow, urlOpener, nil)
} }
// profileLoginHint returns the stored account email for the profile at cfgPath. // profileLoginHint returns the stored account email for the profile at cfgPath.
@@ -217,21 +210,15 @@ func profileLoginHint(cfgPath string) string {
return readProfileEmail(cfgPath) return readProfileEmail(cfgPath)
} }
// runOAuthFlow drives an already acquired OAuth flow to a token: applies the // runOAuthFlow drives an already acquired OAuth flow to a token: requests the
// login hint, requests the flow info, presents the verification URL through // flow info, presents the verification URL through the opener and waits for
// the opener and waits for the browser round-trip. Open is called // the browser round-trip. Open is called synchronously — it is what marks the
// synchronously — it is what marks the surface as opened on the client side, // surface as opened on the client side, and a fast token's OnLoginSuccess is
// and a fast token's OnLoginSuccess is a no-op until it has, so the dismissal // a no-op until it has, so the dismissal would be dropped rather than
// would be dropped rather than delayed. Openers must therefore not block: // delayed. Openers must therefore not block: they post their UI work and
// they post their UI work and return. onWaiting, when set, runs after the URL // return. onWaiting, when set, runs after the URL is shown, right before the
// is shown, right before the blocking wait. // blocking wait.
func runOAuthFlow(ctx context.Context, flow auth.OAuthFlow, hint string, urlOpener URLOpener, onWaiting func()) (*auth.TokenInfo, error) { func runOAuthFlow(ctx context.Context, flow auth.OAuthFlow, urlOpener URLOpener, onWaiting func()) (*auth.TokenInfo, error) {
if hint != "" {
if setter, ok := flow.(loginHintSetter); ok {
setter.SetLoginHint(hint)
}
}
flowInfo, err := flow.RequestAuthInfo(ctx) flowInfo, err := flow.RequestAuthInfo(ctx)
if err != nil { if err != nil {
return nil, fmt.Errorf("request auth info: %w", err) return nil, fmt.Errorf("request auth info: %w", err)
+2 -2
View File
@@ -476,14 +476,14 @@ func (s *SSHClient) requestJWTToken(cfg *profilemanager.Config, cfgPath string)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel() defer cancel()
flow, err := auth.NewOAuthFlow(ctx, cfg, false, true, "") flow, err := auth.NewOAuthFlow(ctx, cfg, false, true, profileLoginHint(cfgPath))
if err != nil { if err != nil {
return "", fmt.Errorf("create oauth flow: %w", err) return "", fmt.Errorf("create oauth flow: %w", err)
} }
// The status callback covers the browser round-trip, which would // The status callback covers the browser round-trip, which would
// otherwise leave the terminal blank. // otherwise leave the terminal blank.
tokenInfo, err := runOAuthFlow(ctx, flow, profileLoginHint(cfgPath), urlOpener, func() { tokenInfo, err := runOAuthFlow(ctx, flow, urlOpener, func() {
s.notifyStatus("Waiting for browser authentication...") s.notifyStatus("Waiting for browser authentication...")
}) })
if err != nil { if err != nil {
+19 -8
View File
@@ -138,26 +138,37 @@ func (a *Auth) IsSSOSupported(ctx context.Context) (bool, error) {
// GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection // GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection
// This avoids creating a new connection to the management server // This avoids creating a new connection to the management server
func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool) (OAuthFlow, error) { func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool, hint string) (OAuthFlow, error) {
var flow OAuthFlow var flow OAuthFlow
var err error
err = a.withRetry(ctx, func(client *mgm.GrpcClient) error { err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
if forceDeviceAuth { if forceDeviceAuth {
flow, err = a.getDeviceFlow(client) deviceFlow, err := a.getDeviceFlow(client)
return err if err != nil {
return err
}
deviceFlow.SetLoginHint(hint)
flow = deviceFlow
return nil
} }
// Try PKCE flow first // Try PKCE flow first
flow, err = a.getPKCEFlow(client) pkceFlow, err := a.getPKCEFlow(client)
if err != nil { if err != nil {
// If PKCE not supported, try Device flow // If PKCE not supported, try Device flow
if s, ok := status.FromError(err); ok && (s.Code() == codes.NotFound || s.Code() == codes.Unimplemented) { if s, ok := status.FromError(err); ok && (s.Code() == codes.NotFound || s.Code() == codes.Unimplemented) {
flow, err = a.getDeviceFlow(client) deviceFlow, err := a.getDeviceFlow(client)
return err if err != nil {
return err
}
deviceFlow.SetLoginHint(hint)
flow = deviceFlow
return nil
} }
return err return err
} }
pkceFlow.SetLoginHint(hint)
flow = pkceFlow
return nil return nil
}) })
+2 -6
View File
@@ -97,9 +97,7 @@ func authenticateWithPKCEFlow(ctx context.Context, config *profilemanager.Config
return nil, fmt.Errorf("getting pkce authorization flow info failed with error: %v", err) return nil, fmt.Errorf("getting pkce authorization flow info failed with error: %v", err)
} }
if hint != "" { pkceFlowInfo.SetLoginHint(hint)
pkceFlowInfo.SetLoginHint(hint)
}
return pkceFlowInfo, nil return pkceFlowInfo, nil
} }
@@ -127,9 +125,7 @@ func authenticateWithDeviceCodeFlow(ctx context.Context, config *profilemanager.
} }
} }
if hint != "" { deviceFlowInfo.SetLoginHint(hint)
deviceFlowInfo.SetLoginHint(hint)
}
return deviceFlowInfo, nil return deviceFlowInfo, nil
} }
+1 -1
View File
@@ -323,7 +323,7 @@ func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName strin
const authInfoRequestTimeout = 30 * time.Second const authInfoRequestTimeout = 30 * time.Second
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, forceDeviceAuth bool) (*auth.TokenInfo, error) { func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, forceDeviceAuth bool) (*auth.TokenInfo, error) {
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth) oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth, "")
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to get OAuth flow: %v", err) return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
} }