175 lines
4.6 KiB
Go
175 lines
4.6 KiB
Go
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()
|
||
}
|