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) {
|
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
@@ -15,14 +15,51 @@ 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) {
|
||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user