fix(auth): improve no-browser OAuth login

This commit is contained in:
lc6464 2026-04-16 23:01:28 +08:00
parent ab019d3f18
commit ffd30d7db7
No known key found for this signature in database
GPG key ID: 53C61B42FEC71D6D
4 changed files with 139 additions and 102 deletions

View file

@ -17,7 +17,7 @@ import (
) )
const ( const (
supportedProvidersMsg = "supported providers: openai, anthropic, google-antigravity" supportedProvidersMsg = "supported providers: openai, anthropic, google-antigravity, antigravity"
defaultAnthropicModel = "claude-sonnet-4.6" defaultAnthropicModel = "claude-sonnet-4.6"
) )

View file

@ -20,7 +20,7 @@ func newLoginCommand() *cobra.Command {
} }
cmd.Flags().StringVarP( cmd.Flags().StringVarP(
&provider, "provider", "p", "", "Provider to login with (openai, anthropic, google-antigravity)", &provider, "provider", "p", "", "Provider to login with (openai, anthropic, google-antigravity, antigravity)",
) )
cmd.Flags().BoolVar(&useDeviceCode, "device-code", false, "Use device code flow (for headless environments)") cmd.Flags().BoolVar(&useDeviceCode, "device-code", false, "Use device code flow (for headless environments)")
cmd.Flags().BoolVar(&noBrowser, "no-browser", false, "Do not auto-open a browser during OAuth login") cmd.Flags().BoolVar(&noBrowser, "no-browser", false, "Do not auto-open a browser during OAuth login")

View file

@ -99,45 +99,33 @@ func LoginBrowserWithOptions(cfg OAuthProviderConfig, opts LoginBrowserOptions)
return nil, fmt.Errorf("generating state: %w", err) return nil, fmt.Errorf("generating state: %w", err)
} }
redirectURI := fmt.Sprintf("http://localhost:%d/auth/callback", cfg.Port) redirectURI := oauthCallbackRedirectURI(cfg.Port)
callbackPort := cfg.Port
var resultCh <-chan callbackResult
authURL := buildAuthorizeURL(cfg, pkce, state, redirectURI) if !opts.NoBrowser {
callbackResultCh := make(chan callbackResult, 1)
resultCh := make(chan callbackResult, 1) listener, actualPort, err := listenOAuthCallback(cfg.Port)
if err != nil {
mux := http.NewServeMux() return nil, fmt.Errorf("starting callback server on port %d: %w", cfg.Port, err)
mux.HandleFunc("/auth/callback", func(w http.ResponseWriter, r *http.Request) {
if r.URL.Query().Get("state") != state {
resultCh <- callbackResult{err: fmt.Errorf("state mismatch")}
http.Error(w, "State mismatch", http.StatusBadRequest)
return
} }
code := r.URL.Query().Get("code") redirectURI = oauthCallbackRedirectURI(actualPort)
if code == "" { callbackPort = actualPort
errMsg := r.URL.Query().Get("error") resultCh = callbackResultCh
resultCh <- callbackResult{err: fmt.Errorf("no code received: %s", errMsg)}
http.Error(w, "No authorization code received", http.StatusBadRequest)
return
}
w.Header().Set("Content-Type", "text/html") server := &http.Server{Handler: oauthCallbackHandler(state, callbackResultCh)}
fmt.Fprint(w, "<html><body><h2>Authentication successful!</h2><p>You can close this window.</p></body></html>") go func() {
resultCh <- callbackResult{code: code} _ = server.Serve(listener)
}) }()
defer func() {
listener, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", cfg.Port)) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
if err != nil { defer cancel()
return nil, fmt.Errorf("starting callback server on port %d: %w", cfg.Port, err) _ = server.Shutdown(ctx)
}()
} }
server := &http.Server{Handler: mux} authURL := buildAuthorizeURL(cfg, pkce, state, redirectURI)
go server.Serve(listener)
defer func() {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
server.Shutdown(ctx)
}()
fmt.Printf("Open this URL to authenticate:\n\n%s\n\n", authURL) fmt.Printf("Open this URL to authenticate:\n\n%s\n\n", authURL)
@ -149,7 +137,7 @@ func LoginBrowserWithOptions(cfg OAuthProviderConfig, opts LoginBrowserOptions)
fmt.Printf( fmt.Printf(
"Wait! If you are in a headless environment (like Coolify/VPS) and cannot reach localhost:%d,\n", "Wait! If you are in a headless environment (like Coolify/VPS) and cannot reach localhost:%d,\n",
cfg.Port, callbackPort,
) )
fmt.Println( fmt.Println(
"please complete the login in your local browser and then PASTE the final redirect URL (or just the code) here.", "please complete the login in your local browser and then PASTE the final redirect URL (or just the code) here.",
@ -157,11 +145,16 @@ func LoginBrowserWithOptions(cfg OAuthProviderConfig, opts LoginBrowserOptions)
fmt.Println("Waiting for authentication (browser or manual paste)...") fmt.Println("Waiting for authentication (browser or manual paste)...")
// Start manual input in a goroutine // Start manual input in a goroutine
manualCh := make(chan string) manualCh := make(chan string, 1)
manualDone := make(chan struct{})
defer close(manualDone)
go func() { go func() {
reader := bufio.NewReader(browserLoginInput) reader := bufio.NewReader(browserLoginInput)
input, _ := reader.ReadString('\n') input, _ := reader.ReadString('\n')
manualCh <- strings.TrimSpace(input) select {
case manualCh <- strings.TrimSpace(input):
case <-manualDone:
}
}() }()
select { select {
@ -191,6 +184,49 @@ func LoginBrowserWithOptions(cfg OAuthProviderConfig, opts LoginBrowserOptions)
} }
} }
func oauthCallbackRedirectURI(port int) string {
return fmt.Sprintf("http://localhost:%d/auth/callback", port)
}
func oauthCallbackHandler(state string, resultCh chan<- callbackResult) http.Handler {
mux := http.NewServeMux()
mux.HandleFunc("/auth/callback", func(w http.ResponseWriter, r *http.Request) {
if r.URL.Query().Get("state") != state {
resultCh <- callbackResult{err: fmt.Errorf("state mismatch")}
http.Error(w, "State mismatch", http.StatusBadRequest)
return
}
code := r.URL.Query().Get("code")
if code == "" {
errMsg := r.URL.Query().Get("error")
resultCh <- callbackResult{err: fmt.Errorf("no code received: %s", errMsg)}
http.Error(w, "No authorization code received", http.StatusBadRequest)
return
}
w.Header().Set("Content-Type", "text/html")
fmt.Fprint(w, "<html><body><h2>Authentication successful!</h2><p>You can close this window.</p></body></html>")
resultCh <- callbackResult{code: code}
})
return mux
}
func listenOAuthCallback(port int) (net.Listener, int, error) {
listener, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", port))
if err != nil {
return nil, 0, err
}
tcpAddr, ok := listener.Addr().(*net.TCPAddr)
if !ok {
_ = listener.Close()
return nil, 0, fmt.Errorf("unexpected listener address type %T", listener.Addr())
}
return listener, tcpAddr.Port, nil
}
type callbackResult struct { type callbackResult struct {
code string code string
err error err error

View file

@ -375,22 +375,16 @@ func TestParseDeviceCodeResponseInvalidInterval(t *testing.T) {
} }
} }
func TestLoginBrowserWithOptionsSkipsAutoOpenWhenDisabled(t *testing.T) { func TestLoginBrowserWithOptionsNoBrowserDoesNotRequireCallbackPort(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { server := newMockOAuthTokenServer()
if r.URL.Path != "/oauth/token" {
http.Error(w, "not found", http.StatusNotFound)
return
}
resp := map[string]any{
"access_token": "mock-access-token",
"refresh_token": "mock-refresh-token",
"expires_in": 3600,
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer server.Close() defer server.Close()
reservedListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("net.Listen() error: %v", err)
}
defer reservedListener.Close()
reservedPort := reservedListener.Addr().(*net.TCPAddr).Port
origOpenBrowserFunc := openBrowserFunc origOpenBrowserFunc := openBrowserFunc
origBrowserLoginInput := browserLoginInput origBrowserLoginInput := browserLoginInput
t.Cleanup(func() { t.Cleanup(func() {
@ -409,7 +403,7 @@ func TestLoginBrowserWithOptionsSkipsAutoOpenWhenDisabled(t *testing.T) {
Issuer: server.URL, Issuer: server.URL,
ClientID: "test-client", ClientID: "test-client",
Scopes: "openid", Scopes: "openid",
Port: freeLocalPort(t), Port: reservedPort,
} }
cred, err := LoginBrowserWithOptions(cfg, LoginBrowserOptions{NoBrowser: true}) cred, err := LoginBrowserWithOptions(cfg, LoginBrowserOptions{NoBrowser: true})
@ -426,7 +420,62 @@ func TestLoginBrowserWithOptionsSkipsAutoOpenWhenDisabled(t *testing.T) {
} }
func TestLoginBrowserWithOptionsAutoOpensByDefault(t *testing.T) { func TestLoginBrowserWithOptionsAutoOpensByDefault(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { server := newMockOAuthTokenServer()
defer server.Close()
origOpenBrowserFunc := openBrowserFunc
origBrowserLoginInput := browserLoginInput
t.Cleanup(func() {
openBrowserFunc = origOpenBrowserFunc
browserLoginInput = origBrowserLoginInput
})
var (
openCalls int
browserURL string
)
openBrowserFunc = func(url string) error {
openCalls++
browserURL = url
return nil
}
browserLoginInput = strings.NewReader("manual-code\n")
cfg := OAuthProviderConfig{
Issuer: server.URL,
ClientID: "test-client",
Scopes: "openid",
Port: 0,
}
_, err := LoginBrowserWithOptions(cfg, LoginBrowserOptions{})
if err != nil {
t.Fatalf("LoginBrowserWithOptions() error: %v", err)
}
if openCalls != 1 {
t.Fatalf("openBrowserFunc call count = %d, want 1", openCalls)
}
parsedBrowserURL, err := url.Parse(browserURL)
if err != nil {
t.Fatalf("url.Parse(browserURL) error: %v", err)
}
redirectURI, err := url.Parse(parsedBrowserURL.Query().Get("redirect_uri"))
if err != nil {
t.Fatalf("url.Parse(redirectURI) error: %v", err)
}
if redirectURI.Port() == "" {
t.Fatal("redirectURI port is empty")
}
if redirectURI.Port() == "0" {
t.Fatalf("redirectURI port = %q, want dynamically assigned port", redirectURI.Port())
}
}
func newMockOAuthTokenServer() *httptest.Server {
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/oauth/token" { if r.URL.Path != "/oauth/token" {
http.Error(w, "not found", http.StatusNotFound) http.Error(w, "not found", http.StatusNotFound)
return return
@ -439,52 +488,4 @@ func TestLoginBrowserWithOptionsAutoOpensByDefault(t *testing.T) {
} }
_ = json.NewEncoder(w).Encode(resp) _ = json.NewEncoder(w).Encode(resp)
})) }))
defer server.Close()
origOpenBrowserFunc := openBrowserFunc
origBrowserLoginInput := browserLoginInput
t.Cleanup(func() {
openBrowserFunc = origOpenBrowserFunc
browserLoginInput = origBrowserLoginInput
})
var openCalls int
openBrowserFunc = func(string) error {
openCalls++
return nil
}
browserLoginInput = strings.NewReader("manual-code\n")
cfg := OAuthProviderConfig{
Issuer: server.URL,
ClientID: "test-client",
Scopes: "openid",
Port: freeLocalPort(t),
}
_, err := LoginBrowserWithOptions(cfg, LoginBrowserOptions{})
if err != nil {
t.Fatalf("LoginBrowserWithOptions() error: %v", err)
}
if openCalls != 1 {
t.Fatalf("openBrowserFunc call count = %d, want 1", openCalls)
}
}
func freeLocalPort(t *testing.T) int {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("net.Listen() error: %v", err)
}
defer listener.Close()
addr, ok := listener.Addr().(*net.TCPAddr)
if !ok {
t.Fatalf("listener addr type = %T, want *net.TCPAddr", listener.Addr())
}
return addr.Port
} }