Skip to content

Commit 16b97bc

Browse files
committed
Add AfterTimeout to Deadline
1 parent af01cdd commit 16b97bc

2 files changed

Lines changed: 125 additions & 3 deletions

File tree

deadline/deadline.go

Lines changed: 43 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,8 @@ type Deadline struct {
3030
deadline time.Time
3131
state deadlineState
3232
pending uint8
33+
cbs map[int]func()
34+
nextCbID int
3335
}
3436

3537
// New creates new deadline timer.
@@ -47,11 +49,41 @@ func (d *Deadline) timeout() {
4749
return
4850
}
4951

50-
d.state = deadlineExceeded
51-
done := d.done
52+
d.fire()
5253
d.mu.Unlock()
54+
}
55+
56+
// AfterTimeout attaches a callback to the deadline.
57+
// The added callback will be triggered when the deadline is met.
58+
// The callback can be detached via the returned function.
59+
// This function mimics context.AfterFunc behavior.
60+
func (d *Deadline) AfterTimeout(cb func()) func() bool {
61+
d.mu.Lock()
62+
defer d.mu.Unlock()
63+
if d.state == deadlineExceeded {
64+
go cb()
5365

54-
close(done)
66+
return func() bool { return false }
67+
}
68+
d.nextCbID++
69+
usedID := d.nextCbID
70+
if d.cbs == nil {
71+
d.cbs = map[int]func(){}
72+
}
73+
d.cbs[d.nextCbID] = cb
74+
cancel := func() bool {
75+
d.mu.Lock()
76+
defer d.mu.Unlock()
77+
if _, has := d.cbs[usedID]; has {
78+
delete(d.cbs, usedID)
79+
80+
return true
81+
}
82+
83+
return false
84+
}
85+
86+
return cancel
5587
}
5688

5789
// Set new deadline. Zero value means no deadline.
@@ -89,8 +121,16 @@ func (d *Deadline) Set(setTo time.Time) {
89121
}
90122

91123
d.pending--
124+
d.fire()
125+
}
126+
127+
func (d *Deadline) fire() {
92128
d.state = deadlineExceeded
93129
close(d.done)
130+
for _, cb := range d.cbs {
131+
go cb()
132+
}
133+
clear(d.cbs)
94134
}
95135

96136
// Done receives deadline signal.

deadline/deadline_test.go

Lines changed: 82 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ package deadline
55

66
import (
77
"context"
8+
"sync/atomic"
89
"testing"
910
"time"
1011

@@ -93,6 +94,87 @@ func TestDeadline(t *testing.T) {
9394
expectedCalls := []byte{0}
9495
assert.Equal(t, expectedCalls, calls, "Wrong order of deadline signal")
9596
})
97+
98+
t.Run("DeadlineAfterTimeoutExceedFuture", func(t *testing.T) {
99+
now := time.Now()
100+
101+
d := New()
102+
var trigger atomic.Int32
103+
d.AfterTimeout(func() {
104+
trigger.Add(1)
105+
})
106+
d.Set(now.Add(10 * time.Millisecond))
107+
<-time.After(20 * time.Millisecond)
108+
d.Set(now.Add(10 * time.Millisecond))
109+
<-time.After(20 * time.Millisecond)
110+
111+
assert.Equal(t, int32(1), trigger.Load(), "Function triggered wrong number of time")
112+
})
113+
114+
t.Run("DeadlineAfterTimeoutExceedPast", func(t *testing.T) {
115+
now := time.Now()
116+
117+
d := New()
118+
var trigger atomic.Int32
119+
d.AfterTimeout(func() {
120+
trigger.Add(1)
121+
})
122+
d.Set(now.Add(10 * time.Millisecond))
123+
d.Set(now.Add(-10 * time.Millisecond))
124+
<-time.After(20 * time.Millisecond)
125+
126+
assert.Equal(t, int32(1), trigger.Load(), "Function triggered wrong number of time")
127+
})
128+
129+
t.Run("DeadlineAfterTimeoutAlreadyExceeded", func(t *testing.T) {
130+
now := time.Now()
131+
132+
d := New()
133+
var trigger atomic.Int32
134+
d.Set(now.Add(10 * time.Millisecond))
135+
d.Set(now.Add(-10 * time.Millisecond))
136+
<-time.After(20 * time.Millisecond)
137+
d.AfterTimeout(func() {
138+
trigger.Add(1)
139+
})
140+
141+
assert.Equal(t, int32(1), trigger.Load(), "Function triggered wrong number of time")
142+
})
143+
144+
t.Run("DeadlineAfterTimeoutDetach", func(t *testing.T) {
145+
now := time.Now()
146+
147+
d := New()
148+
var trigger atomic.Int32
149+
detach := d.AfterTimeout(func() {
150+
trigger.Add(1)
151+
})
152+
ret := detach()
153+
d.Set(now.Add(10 * time.Millisecond))
154+
<-time.After(20 * time.Millisecond)
155+
156+
assert.Equal(t, int32(0), trigger.Load(), "Function triggered wrong number of time")
157+
assert.Equal(t, true, ret, "The detach function returned wrong value")
158+
})
159+
160+
t.Run("DeadlineAfterTimeoutMultiCallbacks", func(t *testing.T) {
161+
now := time.Now()
162+
163+
d := New()
164+
var trigger atomic.Int32
165+
d.AfterTimeout(func() {
166+
trigger.Add(1)
167+
})
168+
d.AfterTimeout(func() {
169+
trigger.Add(1)
170+
})
171+
assert.Equal(t, int32(0), trigger.Load(), "Function triggered wrong number of time")
172+
173+
d.Set(now.Add(10 * time.Millisecond))
174+
<-time.After(20 * time.Millisecond)
175+
176+
assert.Equal(t, int32(2), trigger.Load(), "Function triggered wrong number of time")
177+
})
96178
}
97179

98180
func sendOnDone(ctx context.Context, done <-chan struct{}, dest chan byte, val byte) {

0 commit comments

Comments
 (0)