Skip to content

Commit a44e35d

Browse files
committed
inference: fix embedding runtime lifecycle race
1 parent a7d9e80 commit a44e35d

7 files changed

Lines changed: 55 additions & 27 deletions

File tree

.agents/skills/tidb-test-guidelines/references/domain-case-map.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010

1111
### Tests
1212
- `pkg/domain/db_test.go` - Tests domain session.
13-
- `pkg/domain/domain_test.go` - Tests domain info.
13+
- `pkg/domain/domain_test.go` - Tests domain info and concurrent embedding runtime lifecycle.
1414
- `pkg/domain/domain_utils_test.go` - Tests error code.
1515
- `pkg/domain/domainctx_test.go` - Tests domain ctx.
1616
- `pkg/domain/extract_test.go` - Tests extract plan without history view.

pkg/domain/BUILD.bazel

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -156,7 +156,7 @@ go_test(
156156
],
157157
embed = [":domain"],
158158
flaky = True,
159-
shard_count = 30,
159+
shard_count = 31,
160160
deps = [
161161
"//pkg/config",
162162
"//pkg/ddl",

pkg/domain/domain.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -229,7 +229,7 @@ type Domain struct {
229229
minJobIDRefresher *systable.MinJobIDRefresher
230230

231231
instancePlanCache sessionctx.InstancePlanCache // the instance level plan cache
232-
embedFn *inference.EmbedFn
232+
embedFn atomic.Pointer[inference.EmbedFn]
233233

234234
statsOwner owner.Manager
235235

pkg/domain/domain_test.go

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,29 @@ import (
4747
"go.etcd.io/etcd/tests/v3/integration"
4848
)
4949

50+
func TestEmbeddingRuntimeConcurrentLifecycle(t *testing.T) {
51+
do := &Domain{}
52+
for range 32 {
53+
do.initInferenceProviders()
54+
start := make(chan struct{})
55+
readDone := make(chan struct{})
56+
closeDone := make(chan struct{})
57+
go func() {
58+
<-start
59+
_ = do.GetEmbedFn()
60+
close(readDone)
61+
}()
62+
go func() {
63+
<-start
64+
do.closeInferenceProviders()
65+
close(closeDone)
66+
}()
67+
close(start)
68+
<-readDone
69+
<-closeDone
70+
}
71+
}
72+
5073
func TestInfo(t *testing.T) {
5174
t.Skip("TestInfo will hang currently, it should be fixed later")
5275

pkg/domain/inference.go

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -17,22 +17,23 @@ package domain
1717
import "github.com/pingcap/tidb/pkg/inference"
1818

1919
func (do *Domain) initInferenceProviders() {
20-
do.closeInferenceProviders()
21-
do.embedFn = inference.NewEmbedFn()
20+
oldEmbedFn := do.embedFn.Swap(inference.NewEmbedFn())
21+
if oldEmbedFn != nil {
22+
oldEmbedFn.Close()
23+
}
2224
}
2325

2426
func (do *Domain) closeInferenceProviders() {
25-
if do.embedFn == nil {
26-
return
27+
embedFn := do.embedFn.Swap(nil)
28+
if embedFn != nil {
29+
embedFn.Close()
2730
}
28-
do.embedFn.Close()
29-
do.embedFn = nil
3031
}
3132

3233
// GetEmbedFn returns the embedding function managed by this Domain.
3334
func (do *Domain) GetEmbedFn() *inference.EmbedFn {
3435
if do == nil {
3536
return nil
3637
}
37-
return do.embedFn
38+
return do.embedFn.Load()
3839
}

pkg/expression/integration_test/integration_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4530,7 +4530,7 @@ func TestEmbedTextFunction(t *testing.T) {
45304530
err := tk.QueryToErr(`select embed_text('mock/json', '[1, 3, 4]')`)
45314531
require.ErrorContains(t, err, "EMBED_TEXT is only supported in starter deployment mode")
45324532
if !enableStarterDeployModeForEmbeddingTest(t) {
4533-
return
4533+
t.Skip("EMBED_TEXT functional tests require nextgen kernel in starter deployment mode")
45344534
}
45354535

45364536
err = tk.QueryToErr("select embed_text('text-embedding-3', 'hello world')")

pkg/inference/sqlembed_test.go

Lines changed: 20 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,8 @@ import (
2727
"github.com/stretchr/testify/require"
2828
)
2929

30+
const embeddingTestTimeout = 10 * time.Second
31+
3032
type staticEmbedder struct {
3133
embeddings [][]float32
3234
err error
@@ -197,13 +199,8 @@ func TestEmbedFnSharedCallCancellation(t *testing.T) {
197199
result2 <- result{embedding: embedding, err: err}
198200
}()
199201
require.Eventually(t, func() bool {
200-
embedFn.mu.Lock()
201-
defer embedFn.mu.Unlock()
202-
for _, call := range embedFn.inFlight {
203-
return call.waiters == 2
204-
}
205-
return false
206-
}, time.Second, 10*time.Millisecond)
202+
return hasSingleInFlightCallWithWaiters(embedFn, 2)
203+
}, embeddingTestTimeout, 10*time.Millisecond)
207204

208205
firstCause := errors.New("first caller canceled")
209206
cancel1(firstCause)
@@ -245,21 +242,16 @@ func TestEmbedFnCancelsProviderAfterAllCallersCancel(t *testing.T) {
245242
err2 <- err
246243
}()
247244
require.Eventually(t, func() bool {
248-
embedFn.mu.Lock()
249-
defer embedFn.mu.Unlock()
250-
for _, call := range embedFn.inFlight {
251-
return call.waiters == 2
252-
}
253-
return false
254-
}, time.Second, 10*time.Millisecond)
245+
return hasSingleInFlightCallWithWaiters(embedFn, 2)
246+
}, embeddingTestTimeout, 10*time.Millisecond)
255247

256248
cancel1()
257249
cancel2()
258250
require.ErrorIs(t, receiveFromChannel(t, err1, "first caller cancellation"), context.Canceled)
259251
require.ErrorIs(t, receiveFromChannel(t, err2, "second caller cancellation"), context.Canceled)
260252
select {
261253
case <-provider.canceled:
262-
case <-time.After(time.Second):
254+
case <-time.After(embeddingTestTimeout):
263255
t.Fatal("provider request was not canceled after all callers canceled")
264256
}
265257
}
@@ -278,12 +270,24 @@ func waitForChannel(t *testing.T, ch <-chan struct{}, description string) {
278270
receiveFromChannel(t, ch, description)
279271
}
280272

273+
func hasSingleInFlightCallWithWaiters(embedFn *EmbedFn, waiters int) bool {
274+
embedFn.mu.Lock()
275+
defer embedFn.mu.Unlock()
276+
if len(embedFn.inFlight) != 1 {
277+
return false
278+
}
279+
for _, call := range embedFn.inFlight {
280+
return call.waiters == waiters
281+
}
282+
return false
283+
}
284+
281285
func receiveFromChannel[T any](t *testing.T, ch <-chan T, description string) T {
282286
t.Helper()
283287
select {
284288
case value := <-ch:
285289
return value
286-
case <-time.After(time.Second):
290+
case <-time.After(embeddingTestTimeout):
287291
t.Fatalf("timed out waiting for %s", description)
288292
var zero T
289293
return zero

0 commit comments

Comments
 (0)