fix: propagate QR output failures
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
"os"
|
"os"
|
||||||
@@ -251,6 +252,19 @@ type longPollData struct {
|
|||||||
|
|
||||||
type stringOrNumber string
|
type stringOrNumber string
|
||||||
|
|
||||||
|
type recordingWriter struct {
|
||||||
|
w io.Writer
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (writer *recordingWriter) Write(payload []byte) (int, error) {
|
||||||
|
written, err := writer.w.Write(payload)
|
||||||
|
if err != nil && writer.err == nil {
|
||||||
|
writer.err = err
|
||||||
|
}
|
||||||
|
return written, err
|
||||||
|
}
|
||||||
|
|
||||||
func (value *stringOrNumber) UnmarshalJSON(payload []byte) error {
|
func (value *stringOrNumber) UnmarshalJSON(payload []byte) error {
|
||||||
decoded, err := decodeStringOrNumber(payload)
|
decoded, err := decodeStringOrNumber(payload)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -282,10 +296,18 @@ func (client *Client) Login(ctx context.Context) (AuthData, error) {
|
|||||||
if _, err := qr.Encode(loginData.LoginURL, qr.L); err != nil {
|
if _, err := qr.Encode(loginData.LoginURL, qr.L); err != nil {
|
||||||
return AuthData{}, fmt.Errorf("encode login QR code: %w", err)
|
return AuthData{}, fmt.Errorf("encode login QR code: %w", err)
|
||||||
}
|
}
|
||||||
fmt.Fprintf(client.qrWriter, "请使用米家APP扫描下方二维码\n%s\n", loginData.LoginURL)
|
writer := &recordingWriter{w: client.qrWriter}
|
||||||
qrterminal.GenerateHalfBlock(loginData.LoginURL, qrterminal.L, client.qrWriter)
|
if _, err := fmt.Fprintf(writer, "请使用米家APP扫描下方二维码\n%s\n", loginData.LoginURL); err != nil {
|
||||||
|
return AuthData{}, fmt.Errorf("write QR login output: %w", err)
|
||||||
|
}
|
||||||
|
qrterminal.GenerateHalfBlock(loginData.LoginURL, qrterminal.L, writer)
|
||||||
|
if writer.err != nil {
|
||||||
|
return AuthData{}, fmt.Errorf("write QR login output: %w", writer.err)
|
||||||
|
}
|
||||||
if loginData.QR != "" {
|
if loginData.QR != "" {
|
||||||
fmt.Fprintf(client.qrWriter, "二维码图片: %s\n", loginData.QR)
|
if _, err := fmt.Fprintf(writer, "二维码图片: %s\n", loginData.QR); err != nil {
|
||||||
|
return AuthData{}, fmt.Errorf("write QR login output: %w", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return client.completeQRLogin(ctx, loginData)
|
return client.completeQRLogin(ctx, loginData)
|
||||||
|
|||||||
@@ -12,10 +12,19 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var errQRWriterFailed = errors.New("QR writer failed")
|
||||||
|
|
||||||
|
type failingQRWriter struct{}
|
||||||
|
|
||||||
|
func (failingQRWriter) Write([]byte) (int, error) {
|
||||||
|
return 0, errQRWriterFailed
|
||||||
|
}
|
||||||
|
|
||||||
func TestParseServiceResponse(t *testing.T) {
|
func TestParseServiceResponse(t *testing.T) {
|
||||||
var result struct {
|
var result struct {
|
||||||
Code int `json:"code"`
|
Code int `json:"code"`
|
||||||
@@ -244,6 +253,40 @@ func TestLoginQRCoreFlow(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLoginReturnsQROutputErrorBeforeLongPoll(t *testing.T) {
|
||||||
|
var longPollRequests atomic.Int32
|
||||||
|
var server *httptest.Server
|
||||||
|
server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||||
|
switch request.URL.Path {
|
||||||
|
case "/serviceLogin":
|
||||||
|
_, _ = io.WriteString(writer, `&&&START&&&{"code":70016,"location":"`+server.URL+`/prepare"}`)
|
||||||
|
case "/loginUrl":
|
||||||
|
_, _ = io.WriteString(writer, `&&&START&&&{"code":0,"loginUrl":"https://qr.example/login","qr":"https://qr.example/image","lp":"`+server.URL+`/lp"}`)
|
||||||
|
case "/lp":
|
||||||
|
longPollRequests.Add(1)
|
||||||
|
_, _ = io.WriteString(writer, `&&&START&&&{"code":70016}`)
|
||||||
|
default:
|
||||||
|
http.NotFound(writer, request)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
client, err := NewClient(t.TempDir(), WithHTTPClient(server.Client()), WithQRWriter(failingQRWriter{}))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
client.serviceLoginURL = server.URL + "/serviceLogin"
|
||||||
|
client.loginURL = server.URL + "/loginUrl"
|
||||||
|
|
||||||
|
_, err = client.Login(context.Background())
|
||||||
|
if !errors.Is(err, errQRWriterFailed) {
|
||||||
|
t.Fatalf("Login() error = %v, want %v", err, errQRWriterFailed)
|
||||||
|
}
|
||||||
|
if requests := longPollRequests.Load(); requests != 0 {
|
||||||
|
t.Fatalf("long-poll requests = %d, want 0", requests)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestQRLoginRejectsIncompleteCallbackWithoutSaving(t *testing.T) {
|
func TestQRLoginRejectsIncompleteCallbackWithoutSaving(t *testing.T) {
|
||||||
var server *httptest.Server
|
var server *httptest.Server
|
||||||
server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||||
|
|||||||
Reference in New Issue
Block a user