fix: stop QR output after write failure

This commit is contained in:
2026-07-21 11:06:44 +08:00
parent 503f7cd2b4
commit 965e035052
2 changed files with 58 additions and 6 deletions
+4 -1
View File
@@ -258,8 +258,11 @@ type recordingWriter struct {
} }
func (writer *recordingWriter) Write(payload []byte) (int, error) { func (writer *recordingWriter) Write(payload []byte) (int, error) {
if writer.err != nil {
return 0, writer.err
}
written, err := writer.w.Write(payload) written, err := writer.w.Write(payload)
if err != nil && writer.err == nil { if err != nil {
writer.err = err writer.err = err
} }
return written, err return written, err
+53 -4
View File
@@ -15,15 +15,52 @@ import (
"sync/atomic" "sync/atomic"
"testing" "testing"
"time" "time"
"github.com/mdp/qrterminal/v3"
) )
var errQRWriterFailed = errors.New("QR writer failed") var errQRWriterFailed = errors.New("QR writer failed")
type failingQRWriter struct{} type failThenBlockWriter struct {
calls atomic.Int32
failOnCall int32
blockWrites chan struct{}
}
func (failingQRWriter) Write([]byte) (int, error) { func (writer *failThenBlockWriter) Write(payload []byte) (int, error) {
call := writer.calls.Add(1)
if call == writer.failOnCall {
return 0, errQRWriterFailed return 0, errQRWriterFailed
} }
if call > writer.failOnCall {
<-writer.blockWrites
}
return len(payload), nil
}
func TestRecordingWriterStopsAfterFirstError(t *testing.T) {
underlying := &failThenBlockWriter{failOnCall: 1, blockWrites: make(chan struct{})}
writer := &recordingWriter{w: underlying}
done := make(chan struct{})
go func() {
qrterminal.GenerateHalfBlock("https://qr.example/login", qrterminal.L, writer)
close(done)
}()
select {
case <-done:
case <-time.After(time.Second):
close(underlying.blockWrites)
<-done
t.Fatal("GenerateHalfBlock blocked after the first write failure")
}
if calls := underlying.calls.Load(); calls != 1 {
t.Fatalf("underlying Write calls = %d, want 1", calls)
}
if !errors.Is(writer.err, errQRWriterFailed) {
t.Fatalf("recorded error = %v, want %v", writer.err, errQRWriterFailed)
}
}
func TestParseServiceResponse(t *testing.T) { func TestParseServiceResponse(t *testing.T) {
var result struct { var result struct {
@@ -308,6 +345,7 @@ func TestLoginQRCoreFlow(t *testing.T) {
func TestLoginReturnsQROutputErrorBeforeLongPoll(t *testing.T) { func TestLoginReturnsQROutputErrorBeforeLongPoll(t *testing.T) {
var longPollRequests atomic.Int32 var longPollRequests atomic.Int32
qrWriter := &failThenBlockWriter{failOnCall: 2, blockWrites: make(chan struct{})}
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) {
switch request.URL.Path { switch request.URL.Path {
@@ -324,14 +362,25 @@ func TestLoginReturnsQROutputErrorBeforeLongPoll(t *testing.T) {
})) }))
defer server.Close() defer server.Close()
client, err := NewClient(t.TempDir(), WithHTTPClient(server.Client()), WithQRWriter(failingQRWriter{})) client, err := NewClient(t.TempDir(), WithHTTPClient(server.Client()), WithQRWriter(qrWriter))
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
client.serviceLoginURL = server.URL + "/serviceLogin" client.serviceLoginURL = server.URL + "/serviceLogin"
client.loginURL = server.URL + "/loginUrl" client.loginURL = server.URL + "/loginUrl"
_, err = client.Login(context.Background()) done := make(chan error, 1)
go func() {
_, loginErr := client.Login(context.Background())
done <- loginErr
}()
select {
case err = <-done:
case <-time.After(time.Second):
close(qrWriter.blockWrites)
<-done
t.Fatal("Login() blocked after QR output failed")
}
if !errors.Is(err, errQRWriterFailed) { if !errors.Is(err, errQRWriterFailed) {
t.Fatalf("Login() error = %v, want %v", err, errQRWriterFailed) t.Fatalf("Login() error = %v, want %v", err, errQRWriterFailed)
} }