From 8fea0c3c116f97c05e19dda13679d5b59d99b9df Mon Sep 17 00:00:00 2001 From: m1saka Date: Mon, 20 Jul 2026 14:06:51 +0800 Subject: [PATCH] fix: propagate QR output failures --- auth.go | 28 +++++++++++++++++++++++++--- auth_test.go | 43 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 68 insertions(+), 3 deletions(-) diff --git a/auth.go b/auth.go index 2e8a94c..313b3df 100644 --- a/auth.go +++ b/auth.go @@ -6,6 +6,7 @@ import ( "encoding/json" "errors" "fmt" + "io" "net/http" "net/url" "os" @@ -251,6 +252,19 @@ type longPollData struct { 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 { decoded, err := decodeStringOrNumber(payload) 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 { return AuthData{}, fmt.Errorf("encode login QR code: %w", err) } - fmt.Fprintf(client.qrWriter, "请使用米家APP扫描下方二维码\n%s\n", loginData.LoginURL) - qrterminal.GenerateHalfBlock(loginData.LoginURL, qrterminal.L, client.qrWriter) + writer := &recordingWriter{w: 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 != "" { - 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) diff --git a/auth_test.go b/auth_test.go index f68b569..382e318 100644 --- a/auth_test.go +++ b/auth_test.go @@ -12,10 +12,19 @@ import ( "os" "path/filepath" "strings" + "sync/atomic" "testing" "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) { var result struct { 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) { var server *httptest.Server server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {