Initial release of bdrclone
This commit is contained in:
174
internal/auth/auth.go
Normal file
174
internal/auth/auth.go
Normal file
@@ -0,0 +1,174 @@
|
||||
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()
|
||||
}
|
||||
25
internal/auth/auth_test.go
Normal file
25
internal/auth/auth_test.go
Normal file
@@ -0,0 +1,25 @@
|
||||
package auth
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseAuthorizationCode(t *testing.T) {
|
||||
tests := map[string]string{
|
||||
"plain-code\n": "plain-code",
|
||||
"http://openapi.baidu.com/success?code=a%2Bb": "a+b",
|
||||
"https://example.test/?state=x&code=xyz": "xyz",
|
||||
}
|
||||
for input, want := range tests {
|
||||
got, err := parseAuthorizationCode(input)
|
||||
if err != nil {
|
||||
t.Fatalf("parseAuthorizationCode(%q): %v", input, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Errorf("parseAuthorizationCode(%q) = %q, want %q", input, got, want)
|
||||
}
|
||||
}
|
||||
for _, input := range []string{"", " \n", "two words"} {
|
||||
if _, err := parseAuthorizationCode(input); err == nil {
|
||||
t.Errorf("parseAuthorizationCode(%q) unexpectedly succeeded", input)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user