返回 DeepSeek-Reasonix
stream_lifecycle_test.go
根目录 / internal / extension / providerext / stream_lifecycle_test.go
1 package providerext
2
3 import (
4 "context"
5 "sync"
6 "sync/atomic"
7 "testing"
8
9 "reasonix/internal/extension/protocol"
10 "reasonix/internal/provider"
11 )
12
13 func TestFinishedStreamUnregistersDrainCancel(t *testing.T) {
14 r := testResolver(t, baseCatalog(), nil)
15 var unregistered atomic.Int32
16 stream := &extensionStream{
17 done: make(chan struct{}),
18 unregisterDrainCancel: func() {
19 unregistered.Add(1)
20 },
21 }
22 r.mu.Lock()
23 r.streams["finished"] = stream
24 r.finishLocked("finished", stream, provider.Chunk{})
25 r.finishLocked("finished", stream, provider.Chunk{})
26 r.mu.Unlock()
27 if got := unregistered.Load(); got != 1 {
28 t.Fatalf("drain cancel unregister count = %d, want 1", got)
29 }
30 }
31
32 func TestDrainCancelInstallUnregistersWhenStreamAlreadyFinished(t *testing.T) {
33 r := testResolver(t, baseCatalog(), nil)
34 stream := &extensionStream{done: make(chan struct{})}
35 r.mu.Lock()
36 r.streams["finished-before-install"] = stream
37 r.finishLocked("finished-before-install", stream, provider.Chunk{})
38 r.mu.Unlock()
39
40 var unregistered atomic.Int32
41 r.installDrainCancel("finished-before-install", stream, func() {
42 unregistered.Add(1)
43 })
44 if got := unregistered.Load(); got != 1 {
45 t.Fatalf("late drain cancel unregister count = %d, want 1", got)
46 }
47 if stream.unregisterDrainCancel != nil {
48 t.Fatal("completed stream retained a drain cancel unregister callback")
49 }
50 }
51
52 func TestConcurrentStreamsReuseProviderHandleWithoutMutation(t *testing.T) {
53 fc := newFakeClient("demo", demoDescriptor())
54 r := testResolver(t, baseCatalog(), nil, fc)
55 p, err := r.Resolve(provider.Selection{Ref: "plugin/demo/fake/x"})
56 if err != nil {
57 t.Fatalf("Resolve: %v", err)
58 }
59
60 type result struct {
61 out <-chan provider.Chunk
62 err error
63 }
64 const streamCount = 8
65 results := make(chan result, streamCount)
66 var wg sync.WaitGroup
67 for range streamCount {
68 wg.Go(func() {
69 out, streamErr := p.Stream(context.Background(), provider.Request{
70 Messages: []provider.Message{{Role: provider.RoleUser}},
71 })
72 results <- result{out: out, err: streamErr}
73 })
74 }
75 wg.Wait()
76 close(results)
77
78 var outputs []<-chan provider.Chunk
79 for item := range results {
80 if item.err != nil {
81 t.Fatalf("Stream: %v", item.err)
82 }
83 outputs = append(outputs, item.out)
84 }
85 fc.mu.Lock()
86 opened := append([]protocol.StreamOpenParams(nil), fc.opened...)
87 fc.mu.Unlock()
88 if len(opened) != streamCount {
89 t.Fatalf("opened streams = %d, want %d", len(opened), streamCount)
90 }
91 for _, params := range opened {
92 r.RouteStreamEnd(protocol.StreamEndParams{StreamID: params.StreamID, LastSeq: 0})
93 }
94 for _, out := range outputs {
95 if chunks := collectChunks(t, out); len(chunks) != 0 {
96 t.Fatalf("clean empty stream delivered %d chunks", len(chunks))
97 }
98 }
99 }
100
100 lines GO