diff --git a/auth.go b/auth.go index fe80196..5d8e3ab 100644 --- a/auth.go +++ b/auth.go @@ -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 diff --git a/auth_test.go b/auth_test.go index 9055f48..a6c3204 100644 --- a/auth_test.go +++ b/auth_test.go @@ -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) { - return 0, errQRWriterFailed +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) }