mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-03 03:59:07 +02:00
[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:
+11
-24
@@ -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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user