fix(auth): improve no-browser OAuth login
This commit is contained in:
parent
ab019d3f18
commit
ffd30d7db7
4 changed files with 139 additions and 102 deletions
|
|
@ -17,7 +17,7 @@ import (
|
|||
)
|
||||
|
||||
const (
|
||||
supportedProvidersMsg = "supported providers: openai, anthropic, google-antigravity"
|
||||
supportedProvidersMsg = "supported providers: openai, anthropic, google-antigravity, antigravity"
|
||||
defaultAnthropicModel = "claude-sonnet-4.6"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -20,7 +20,7 @@ func newLoginCommand() *cobra.Command {
|
|||
}
|
||||
|
||||
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(&noBrowser, "no-browser", false, "Do not auto-open a browser during OAuth login")
|
||||
|
|
|
|||
|
|
@ -99,45 +99,33 @@ func LoginBrowserWithOptions(cfg OAuthProviderConfig, opts LoginBrowserOptions)
|
|||
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)
|
||||
|
||||
resultCh := make(chan callbackResult, 1)
|
||||
|
||||
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
|
||||
if !opts.NoBrowser {
|
||||
callbackResultCh := make(chan callbackResult, 1)
|
||||
listener, actualPort, err := listenOAuthCallback(cfg.Port)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("starting callback server on port %d: %w", cfg.Port, err)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
redirectURI = oauthCallbackRedirectURI(actualPort)
|
||||
callbackPort = actualPort
|
||||
resultCh = callbackResultCh
|
||||
|
||||
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}
|
||||
})
|
||||
|
||||
listener, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", cfg.Port))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("starting callback server on port %d: %w", cfg.Port, err)
|
||||
server := &http.Server{Handler: oauthCallbackHandler(state, callbackResultCh)}
|
||||
go func() {
|
||||
_ = server.Serve(listener)
|
||||
}()
|
||||
defer func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
_ = server.Shutdown(ctx)
|
||||
}()
|
||||
}
|
||||
|
||||
server := &http.Server{Handler: mux}
|
||||
go server.Serve(listener)
|
||||
defer func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
server.Shutdown(ctx)
|
||||
}()
|
||||
authURL := buildAuthorizeURL(cfg, pkce, state, redirectURI)
|
||||
|
||||
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(
|
||||
"Wait! If you are in a headless environment (like Coolify/VPS) and cannot reach localhost:%d,\n",
|
||||
cfg.Port,
|
||||
callbackPort,
|
||||
)
|
||||
fmt.Println(
|
||||
"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)...")
|
||||
|
||||
// Start manual input in a goroutine
|
||||
manualCh := make(chan string)
|
||||
manualCh := make(chan string, 1)
|
||||
manualDone := make(chan struct{})
|
||||
defer close(manualDone)
|
||||
go func() {
|
||||
reader := bufio.NewReader(browserLoginInput)
|
||||
input, _ := reader.ReadString('\n')
|
||||
manualCh <- strings.TrimSpace(input)
|
||||
select {
|
||||
case manualCh <- strings.TrimSpace(input):
|
||||
case <-manualDone:
|
||||
}
|
||||
}()
|
||||
|
||||
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 {
|
||||
code string
|
||||
err error
|
||||
|
|
|
|||
|
|
@ -375,22 +375,16 @@ func TestParseDeviceCodeResponseInvalidInterval(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestLoginBrowserWithOptionsSkipsAutoOpenWhenDisabled(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
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)
|
||||
}))
|
||||
func TestLoginBrowserWithOptionsNoBrowserDoesNotRequireCallbackPort(t *testing.T) {
|
||||
server := newMockOAuthTokenServer()
|
||||
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
|
||||
origBrowserLoginInput := browserLoginInput
|
||||
t.Cleanup(func() {
|
||||
|
|
@ -409,7 +403,7 @@ func TestLoginBrowserWithOptionsSkipsAutoOpenWhenDisabled(t *testing.T) {
|
|||
Issuer: server.URL,
|
||||
ClientID: "test-client",
|
||||
Scopes: "openid",
|
||||
Port: freeLocalPort(t),
|
||||
Port: reservedPort,
|
||||
}
|
||||
|
||||
cred, err := LoginBrowserWithOptions(cfg, LoginBrowserOptions{NoBrowser: true})
|
||||
|
|
@ -426,7 +420,62 @@ func TestLoginBrowserWithOptionsSkipsAutoOpenWhenDisabled(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" {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
return
|
||||
|
|
@ -439,52 +488,4 @@ func TestLoginBrowserWithOptionsAutoOpensByDefault(t *testing.T) {
|
|||
}
|
||||
_ = 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
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue