From f106894100728ddd1bc2b7d9b668f8c2b181dd72 Mon Sep 17 00:00:00 2001 From: Thomas Legris Date: Thu, 3 Sep 2026 00:05:52 +0900 Subject: [PATCH] Add AfterTimeout to Deadline --- deadline/deadline.go | 46 ++++++++++++++++++++-- deadline/deadline_test.go | 83 +++++++++++++++++++++++++++++++++++++++ 2 files changed, 126 insertions(+), 3 deletions(-) diff --git a/deadline/deadline.go b/deadline/deadline.go index b0645bee..b81685b6 100644 --- a/deadline/deadline.go +++ b/deadline/deadline.go @@ -30,6 +30,8 @@ type Deadline struct { deadline time.Time state deadlineState pending uint8 + cbs map[int]func() + nextCbID int } // New creates new deadline timer. @@ -47,11 +49,41 @@ func (d *Deadline) timeout() { return } - d.state = deadlineExceeded - done := d.done + d.fire() d.mu.Unlock() +} + +// AfterTimeout attaches a callback to the deadline. +// The added callback will be triggered when the deadline is met. +// The callback can be detached via the returned function. +// This function mimics context.AfterFunc behavior. +func (d *Deadline) AfterTimeout(cb func()) func() bool { + d.mu.Lock() + defer d.mu.Unlock() + if d.state == deadlineExceeded { + go cb() - close(done) + return func() bool { return false } + } + d.nextCbID++ + usedID := d.nextCbID + if d.cbs == nil { + d.cbs = map[int]func(){} + } + d.cbs[d.nextCbID] = cb + cancel := func() bool { + d.mu.Lock() + defer d.mu.Unlock() + if _, has := d.cbs[usedID]; has { + delete(d.cbs, usedID) + + return true + } + + return false + } + + return cancel } // Set new deadline. Zero value means no deadline. @@ -89,8 +121,16 @@ func (d *Deadline) Set(setTo time.Time) { } d.pending-- + d.fire() +} + +func (d *Deadline) fire() { d.state = deadlineExceeded close(d.done) + for _, cb := range d.cbs { + go cb() + } + clear(d.cbs) } // Done receives deadline signal. diff --git a/deadline/deadline_test.go b/deadline/deadline_test.go index 280d7ecf..323ef9d1 100644 --- a/deadline/deadline_test.go +++ b/deadline/deadline_test.go @@ -5,6 +5,7 @@ package deadline import ( "context" + "sync/atomic" "testing" "time" @@ -93,6 +94,88 @@ func TestDeadline(t *testing.T) { expectedCalls := []byte{0} assert.Equal(t, expectedCalls, calls, "Wrong order of deadline signal") }) + + t.Run("DeadlineAfterTimeoutExceedFuture", func(t *testing.T) { + now := time.Now() + + d := New() + var trigger atomic.Int32 + d.AfterTimeout(func() { + trigger.Add(1) + }) + d.Set(now.Add(10 * time.Millisecond)) + <-time.After(20 * time.Millisecond) + d.Set(now.Add(10 * time.Millisecond)) + <-time.After(20 * time.Millisecond) + + assert.Equal(t, int32(1), trigger.Load(), "Function triggered wrong number of time") + }) + + t.Run("DeadlineAfterTimeoutExceedPast", func(t *testing.T) { + now := time.Now() + + d := New() + var trigger atomic.Int32 + d.AfterTimeout(func() { + trigger.Add(1) + }) + d.Set(now.Add(10 * time.Millisecond)) + d.Set(now.Add(-10 * time.Millisecond)) + <-time.After(20 * time.Millisecond) + + assert.Equal(t, int32(1), trigger.Load(), "Function triggered wrong number of time") + }) + + t.Run("DeadlineAfterTimeoutAlreadyExceeded", func(t *testing.T) { + now := time.Now() + + d := New() + var trigger atomic.Int32 + d.Set(now.Add(10 * time.Millisecond)) + d.Set(now.Add(-10 * time.Millisecond)) + <-time.After(20 * time.Millisecond) + d.AfterTimeout(func() { + trigger.Add(1) + }) + <-time.After(20 * time.Millisecond) + + assert.Equal(t, int32(1), trigger.Load(), "Function triggered wrong number of time") + }) + + t.Run("DeadlineAfterTimeoutDetach", func(t *testing.T) { + now := time.Now() + + d := New() + var trigger atomic.Int32 + detach := d.AfterTimeout(func() { + trigger.Add(1) + }) + ret := detach() + d.Set(now.Add(10 * time.Millisecond)) + <-time.After(20 * time.Millisecond) + + assert.Equal(t, int32(0), trigger.Load(), "Function triggered wrong number of time") + assert.Equal(t, true, ret, "The detach function returned wrong value") + }) + + t.Run("DeadlineAfterTimeoutMultiCallbacks", func(t *testing.T) { + now := time.Now() + + d := New() + var trigger atomic.Int32 + d.AfterTimeout(func() { + trigger.Add(1) + }) + d.AfterTimeout(func() { + trigger.Add(1) + }) + assert.Equal(t, int32(0), trigger.Load(), "Function triggered wrong number of time") + + d.Set(now.Add(10 * time.Millisecond)) + <-time.After(20 * time.Millisecond) + + assert.Equal(t, int32(2), trigger.Load(), "Function triggered wrong number of time") + }) } func sendOnDone(ctx context.Context, done <-chan struct{}, dest chan byte, val byte) {