Files
bdrclone/internal/auth/auth.go
2026-08-13 23:02:49 +08:00

175 lines
4.6 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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("<!doctype html><meta charset=utf-8><title>bdrclone</title><p>授权成功,可以关闭此页面。</p>"))
}
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()
}