fix: stop QR output after write failure
This commit is contained in:
@@ -258,8 +258,11 @@ type recordingWriter struct {
|
||||
}
|
||||
|
||||
func (writer *recordingWriter) Write(payload []byte) (int, error) {
|
||||
if writer.err != nil {
|
||||
return 0, writer.err
|
||||
}
|
||||
written, err := writer.w.Write(payload)
|
||||
if err != nil && writer.err == nil {
|
||||
if err != nil {
|
||||
writer.err = err
|
||||
}
|
||||
return written, err
|
||||
|
||||
+53
-4
@@ -15,14 +15,51 @@ import (
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/mdp/qrterminal/v3"
|
||||
)
|
||||
|
||||
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
|
||||
}
|
||||
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) {
|
||||
@@ -308,6 +345,7 @@ func TestLoginQRCoreFlow(t *testing.T) {
|
||||
|
||||
func TestLoginReturnsQROutputErrorBeforeLongPoll(t *testing.T) {
|
||||
var longPollRequests atomic.Int32
|
||||
qrWriter := &failThenBlockWriter{failOnCall: 2, blockWrites: make(chan struct{})}
|
||||
var server *httptest.Server
|
||||
server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
switch request.URL.Path {
|
||||
@@ -324,14 +362,25 @@ func TestLoginReturnsQROutputErrorBeforeLongPoll(t *testing.T) {
|
||||
}))
|
||||
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 {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client.serviceLoginURL = server.URL + "/serviceLogin"
|
||||
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) {
|
||||
t.Fatalf("Login() error = %v, want %v", err, errQRWriterFailed)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user