Skip to content

Commit 5d6b2e7

Browse files
committed
feat(orchestrator): drain sandboxes during shutdown
Introduce a shared draingate.Gate (counter plus notification channel) and use it in the sandbox factory and gRPC server to reject new sandbox starts while draining, wait for in-flight starts, and drain or force-stop live sandboxes before closers run. Forced shutdown preserves buffered close errors on context cancellation and avoids duplicate final-pass errors. Adds utils.WaitGroupWait to wait on a WaitGroup with context cancellation. The graceful drain phase is bounded by SHUTDOWN_DRAIN_TIMEOUT; when it expires the drain escalates to a forced sandbox shutdown. By default the drain waits forever, until sandboxes exit on their own or a force-stop API call empties the node.
1 parent de2f391 commit 5d6b2e7

13 files changed

Lines changed: 1270 additions & 21 deletions

File tree

‎packages/orchestrator/pkg/cfg/model.go‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -105,6 +105,10 @@ type Config struct {
105105
NBDPoolSize int `env:"NBD_POOL_SIZE" envDefault:"64"`
106106
Services []string `env:"ORCHESTRATOR_SERVICES" envDefault:"orchestrator"`
107107
PersistentVolumeMounts map[string]string `env:"PERSISTENT_VOLUME_MOUNTS"`
108+
// ShutdownDrainTimeout bounds the graceful drain phase during shutdown.
109+
// Zero (the default) drains forever: shutdown waits until sandboxes exit
110+
// on their own or a force-stop API call empties the node.
111+
ShutdownDrainTimeout time.Duration `env:"SHUTDOWN_DRAIN_TIMEOUT"`
108112
}
109113

110114
// AdditionalClickhouseEndpoints returns the non-blank entries from
Lines changed: 133 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,133 @@
1+
package draingate
2+
3+
import (
4+
"context"
5+
"errors"
6+
"fmt"
7+
"sync"
8+
)
9+
10+
var ErrDraining = errors.New("drain gate is draining")
11+
12+
type Gate struct {
13+
initOnce sync.Once
14+
drainOnce sync.Once
15+
16+
mu sync.Mutex
17+
done chan struct{}
18+
draining bool
19+
count int
20+
changed chan struct{}
21+
}
22+
23+
func New() *Gate {
24+
g := &Gate{}
25+
g.init()
26+
27+
return g
28+
}
29+
30+
func (g *Gate) init() {
31+
g.initOnce.Do(func() {
32+
g.done = make(chan struct{})
33+
g.changed = make(chan struct{})
34+
})
35+
}
36+
37+
func (g *Gate) StartDraining() bool {
38+
if g == nil {
39+
return false
40+
}
41+
42+
g.init()
43+
transitioned := false
44+
g.drainOnce.Do(func() {
45+
g.mu.Lock()
46+
defer g.mu.Unlock()
47+
48+
g.draining = true
49+
transitioned = true
50+
close(g.done)
51+
})
52+
53+
return transitioned
54+
}
55+
56+
func (g *Gate) Draining() bool {
57+
if g == nil {
58+
return false
59+
}
60+
61+
g.init()
62+
select {
63+
case <-g.done:
64+
return true
65+
default:
66+
return false
67+
}
68+
}
69+
70+
func (g *Gate) Done() <-chan struct{} {
71+
if g == nil {
72+
return nil
73+
}
74+
75+
g.init()
76+
77+
return g.done
78+
}
79+
80+
func (g *Gate) Enter() (func(), error) {
81+
if g == nil {
82+
return func() {}, nil
83+
}
84+
85+
g.init()
86+
g.mu.Lock()
87+
defer g.mu.Unlock()
88+
89+
if g.draining {
90+
return nil, ErrDraining
91+
}
92+
93+
g.count++
94+
g.notifyChangeLocked()
95+
96+
return sync.OnceFunc(func() {
97+
g.mu.Lock()
98+
defer g.mu.Unlock()
99+
100+
g.count--
101+
g.notifyChangeLocked()
102+
}), nil
103+
}
104+
105+
func (g *Gate) Wait(ctx context.Context) error {
106+
if g == nil {
107+
return nil
108+
}
109+
110+
g.init()
111+
for {
112+
g.mu.Lock()
113+
if g.count == 0 {
114+
g.mu.Unlock()
115+
116+
return nil
117+
}
118+
119+
changed := g.changed
120+
g.mu.Unlock()
121+
122+
select {
123+
case <-ctx.Done():
124+
return fmt.Errorf("%w", ctx.Err())
125+
case <-changed:
126+
}
127+
}
128+
}
129+
130+
func (g *Gate) notifyChangeLocked() {
131+
close(g.changed)
132+
g.changed = make(chan struct{})
133+
}
Lines changed: 207 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,207 @@
1+
package draingate
2+
3+
import (
4+
"context"
5+
"sync"
6+
"testing"
7+
"time"
8+
9+
"github.com/stretchr/testify/require"
10+
)
11+
12+
func TestEnterReleaseAndWait(t *testing.T) {
13+
t.Parallel()
14+
15+
g := New()
16+
release, err := g.Enter()
17+
require.NoError(t, err)
18+
19+
waitDone := make(chan error, 1)
20+
go func() {
21+
waitDone <- g.Wait(t.Context())
22+
}()
23+
24+
requireNotDone(t, waitDone)
25+
release()
26+
release()
27+
require.NoError(t, requireDone(t, waitDone))
28+
}
29+
30+
func TestRejectAfterDrain(t *testing.T) {
31+
t.Parallel()
32+
33+
g := New()
34+
require.True(t, g.StartDraining())
35+
require.False(t, g.StartDraining())
36+
require.True(t, g.Draining())
37+
38+
select {
39+
case <-g.Done():
40+
default:
41+
t.Fatal("Done channel was not closed")
42+
}
43+
44+
release, err := g.Enter()
45+
require.ErrorIs(t, err, ErrDraining)
46+
require.Nil(t, release)
47+
}
48+
49+
func TestWaitBlocksUntilAllReleases(t *testing.T) {
50+
t.Parallel()
51+
52+
g := New()
53+
release1, err := g.Enter()
54+
require.NoError(t, err)
55+
release2, err := g.Enter()
56+
require.NoError(t, err)
57+
58+
g.StartDraining()
59+
waitDone := make(chan error, 1)
60+
go func() {
61+
waitDone <- g.Wait(t.Context())
62+
}()
63+
64+
requireNotDone(t, waitDone)
65+
release1()
66+
requireNotDone(t, waitDone)
67+
release2()
68+
require.NoError(t, requireDone(t, waitDone))
69+
}
70+
71+
func TestWaitReturnsContextError(t *testing.T) {
72+
t.Parallel()
73+
74+
g := New()
75+
release, err := g.Enter()
76+
require.NoError(t, err)
77+
defer release()
78+
79+
ctx, cancel := context.WithCancel(t.Context())
80+
cancel()
81+
82+
require.ErrorIs(t, g.Wait(ctx), context.Canceled)
83+
}
84+
85+
func TestEnterDuringWaitWhileNotDrainingIsNotBlocked(t *testing.T) {
86+
t.Parallel()
87+
88+
g := New()
89+
release, err := g.Enter()
90+
require.NoError(t, err)
91+
92+
waitDone := make(chan error, 1)
93+
go func() {
94+
waitDone <- g.Wait(t.Context())
95+
}()
96+
requireNotDone(t, waitDone)
97+
98+
entered := make(chan error, 1)
99+
go func() {
100+
secondRelease, err := g.Enter()
101+
if err == nil {
102+
secondRelease()
103+
}
104+
entered <- err
105+
}()
106+
107+
require.NoError(t, requireDone(t, entered))
108+
release()
109+
require.NoError(t, requireDone(t, waitDone))
110+
}
111+
112+
func TestConcurrentStress(t *testing.T) {
113+
t.Parallel()
114+
115+
g := New()
116+
start := make(chan struct{})
117+
entered := make(chan func(), 100)
118+
rejected := make(chan error, 100)
119+
var wg sync.WaitGroup
120+
121+
for range 100 {
122+
wg.Go(func() {
123+
<-start
124+
release, err := g.Enter()
125+
if err != nil {
126+
rejected <- err
127+
128+
return
129+
}
130+
131+
entered <- release
132+
})
133+
}
134+
135+
close(start)
136+
require.Eventually(t, func() bool {
137+
return len(entered)+len(rejected) == 100
138+
}, time.Second, time.Millisecond)
139+
140+
g.StartDraining()
141+
wg.Wait()
142+
close(entered)
143+
close(rejected)
144+
145+
for err := range rejected {
146+
require.ErrorIs(t, err, ErrDraining)
147+
}
148+
149+
waitDone := make(chan error, 1)
150+
go func() {
151+
waitDone <- g.Wait(t.Context())
152+
}()
153+
154+
for release := range entered {
155+
release()
156+
}
157+
require.NoError(t, requireDone(t, waitDone))
158+
159+
_, err := g.Enter()
160+
require.ErrorIs(t, err, ErrDraining)
161+
}
162+
163+
func TestZeroValueAndNilGateAreSafe(t *testing.T) {
164+
t.Parallel()
165+
166+
var g Gate
167+
release, err := g.Enter()
168+
require.NoError(t, err)
169+
release()
170+
require.NoError(t, g.Wait(t.Context()))
171+
require.True(t, g.StartDraining())
172+
require.True(t, g.Draining())
173+
174+
var nilGate *Gate
175+
release, err = nilGate.Enter()
176+
require.NoError(t, err)
177+
release()
178+
require.NoError(t, nilGate.Wait(t.Context()))
179+
require.False(t, nilGate.StartDraining())
180+
require.False(t, nilGate.Draining())
181+
require.Nil(t, nilGate.Done())
182+
}
183+
184+
func requireDone[T any](t *testing.T, ch <-chan T) T {
185+
t.Helper()
186+
187+
select {
188+
case got := <-ch:
189+
return got
190+
case <-time.After(time.Second):
191+
t.Fatal("channel did not receive")
192+
193+
var zero T
194+
195+
return zero
196+
}
197+
}
198+
199+
func requireNotDone[T any](t *testing.T, ch <-chan T) {
200+
t.Helper()
201+
202+
select {
203+
case got := <-ch:
204+
t.Fatalf("channel received unexpectedly: %v", got)
205+
case <-time.After(25 * time.Millisecond):
206+
}
207+
}

0 commit comments

Comments
 (0)