fix: propagate QR output failures

This commit is contained in:
2026-07-20 14:06:51 +08:00
parent 43c688af6f
commit 8fea0c3c11
2 changed files with 68 additions and 3 deletions
+25 -3
View File
@@ -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)
+43
View File
@@ -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) {