package auth import ( "bufio" "context" "crypto/rand" "encoding/hex" "errors" "fmt" "html" "io" "net" "net/http" "net/url" "os/exec" "runtime" "strings" "time" "gitea.dddbg.com/youbin/bdrclone/internal/baidu" ) const OOBRedirectURI = "oob" func Authorize(ctx context.Context, client *baidu.Client, redirectURI string, openBrowser bool) error { u, err := url.Parse(redirectURI) if err != nil { return fmt.Errorf("parse redirect_uri: %w", err) } if u.Scheme != "http" || u.Hostname() != "127.0.0.1" { return errors.New("automatic auth requires an http://127.0.0.1 redirect_uri") } state, err := randomState() if err != nil { return err } listener, err := net.Listen("tcp", u.Host) if err != nil { return fmt.Errorf("listen for OAuth callback on %s: %w", u.Host, err) } defer listener.Close() result := make(chan error, 1) mux := http.NewServeMux() mux.HandleFunc(u.Path, func(w http.ResponseWriter, r *http.Request) { if r.URL.Query().Get("state") != state { http.Error(w, "OAuth state mismatch", http.StatusBadRequest) select { case result <- errors.New("OAuth state mismatch"): default: } return } if code := r.URL.Query().Get("error"); code != "" { message := r.URL.Query().Get("error_description") if message == "" { message = code } http.Error(w, message, http.StatusBadRequest) select { case result <- errors.New(message): default: } return } code := r.URL.Query().Get("code") if code == "" { http.Error(w, "Authorization code is missing", http.StatusBadRequest) select { case result <- errors.New("authorization code is missing"): default: } return } _, exchangeErr := client.ExchangeCode(r.Context(), code, redirectURI) if exchangeErr != nil { http.Error(w, html.EscapeString(exchangeErr.Error()), http.StatusBadGateway) } else { w.Header().Set("Content-Type", "text/html; charset=utf-8") _, _ = w.Write([]byte("bdrclone

授权成功,可以关闭此页面。

")) } select { case result <- exchangeErr: default: } }) server := &http.Server{Handler: mux, ReadHeaderTimeout: 10 * time.Second} serverErr := make(chan error, 1) go func() { if err := server.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) { serverErr <- err } }() authorizeURL := client.AuthorizationURL(redirectURI, state) fmt.Printf("请在浏览器中授权:\n%s\n", authorizeURL) if openBrowser { _ = openURL(authorizeURL) } select { case err := <-result: shutdownCtx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() _ = server.Shutdown(shutdownCtx) return err case err := <-serverErr: return err case <-ctx.Done(): return ctx.Err() } } // AuthorizeOOB uses Baidu's documented out-of-band flow. Baidu displays the // authorization code in its own page instead of redirecting to a local server. func AuthorizeOOB(ctx context.Context, client *baidu.Client, openBrowser bool, input io.Reader, output io.Writer) error { state, err := randomState() if err != nil { return err } authorizeURL := client.AuthorizationURL(OOBRedirectURI, state) fmt.Fprintf(output, "请在浏览器中授权:\n%s\n\n授权后,将页面显示的授权码粘贴到这里:", authorizeURL) if openBrowser { _ = openURL(authorizeURL) } line, err := bufio.NewReader(input).ReadString('\n') if err != nil && !errors.Is(err, io.EOF) { return fmt.Errorf("read authorization code: %w", err) } if err := ctx.Err(); err != nil { return err } code, err := parseAuthorizationCode(line) if err != nil { return err } _, err = client.ExchangeCode(ctx, code, OOBRedirectURI) return err } func parseAuthorizationCode(input string) (string, error) { value := strings.TrimSpace(input) if value == "" { return "", errors.New("authorization code is empty") } if parsed, err := url.Parse(value); err == nil && parsed.Query().Get("code") != "" { value = parsed.Query().Get("code") } if strings.ContainsAny(value, " \t\r\n") { return "", errors.New("authorization code contains whitespace") } return value, nil } func randomState() (string, error) { b := make([]byte, 24) if _, err := rand.Read(b); err != nil { return "", fmt.Errorf("generate OAuth state: %w", err) } return hex.EncodeToString(b), nil } func openURL(target string) error { var command string var args []string switch runtime.GOOS { case "darwin": command, args = "open", []string{target} case "windows": command, args = "rundll32", []string{"url.dll,FileProtocolHandler", target} default: command, args = "xdg-open", []string{target} } return exec.Command(command, args...).Start() }