diff --git a/await.go b/await.go index 6b7b4e69b..a6b616502 100644 --- a/await.go +++ b/await.go @@ -138,7 +138,10 @@ func (a *PublicationAwaiter) pollLoop(ctx context.Context, readCheckpoint func(c cp, cpErr = readCheckpoint(ctx) switch { case errors.Is(cpErr, os.ErrNotExist): - return false, nil + // The log has not published a checkpoint yet. This is not an error, + // but we still need to fall through and broadcast so that waiters + // get a chance to notice if their context has been cancelled. + cp, cpSize, cpErr = nil, 0, nil case cpErr != nil: cpSize = 0 default: diff --git a/await_test.go b/await_test.go index c077a99df..bec0eb51d 100644 --- a/await_test.go +++ b/await_test.go @@ -19,12 +19,14 @@ import ( "context" "crypto/sha256" "fmt" + "os" "sync" "sync/atomic" "time" "errors" "testing" + "testing/synctest" "github.com/transparency-dev/formats/log" "golang.org/x/mod/sumdb/note" @@ -294,6 +296,49 @@ func TestAwait_contextCancel(t *testing.T) { wg.Wait() } +// TestAwait_noCheckpointRespectsContext checks that Await returns once its +// context expires, even if the log has not yet published a checkpoint. +// +// See https://github.com/transparency-dev/tessera/issues/1192. +func TestAwait_noCheckpointRespectsContext(t *testing.T) { + t.Parallel() + + // Within the synctest bubble, time only advances once every goroutine is + // durably blocked, so Await is guaranteed to be waiting on the awaiter + // before its context deadline passes. + synctest.Test(t, func(t *testing.T) { + // The awaiter is long-lived, so it has its own context which outlives the + // Await call below. Cancelling it at the end of the test lets the poll loop + // broadcast one final time, releasing any goroutine still stuck in Await. + awaiterCtx, cancelAwaiter := context.WithCancel(t.Context()) + defer cancelAwaiter() + + // Simulate a log which has not yet published a checkpoint. + readCheckpoint := func(_ context.Context) ([]byte, error) { + return nil, os.ErrNotExist + } + awaiter := NewPublicationAwaiter(awaiterCtx, readCheckpoint, 10*time.Millisecond) + + awaitCtx, cancelAwait := context.WithTimeout(t.Context(), 50*time.Millisecond) + defer cancelAwait() + + errC := make(chan error, 1) + go func() { + _, _, err := awaiter.Await(awaitCtx, func() (Index, error) { return Index{Index: 0}, nil }) + errC <- err + }() + + select { + case err := <-errC: + if !errors.Is(err, context.DeadlineExceeded) { + t.Errorf("Await: got err %v, want %v", err, context.DeadlineExceeded) + } + case <-time.After(1 * time.Second): + t.Error("Await did not return after its context expired") + } + }) +} + func BenchmarkAwait(b *testing.B) { cpFormat := "origin/\n%d\nhash\n\nsig" cpSize := atomic.Uint64{}