feat: 添加 WithStreamErrors 查询流式调用执行过程中的异常

流式输出/双向流的 handler(Python 生成器)如果执行过程中抛异常,
Invoke[chan T] 本身的 err 只描述"调用有没有发起成功",跟这个异常
无关(永远是 nil),channel 只会静默提前关闭,调用方原本完全无法
感知。新增 WithStreamErrors(ctx) 返回一个包过的 ctx 和一个查询函数
streamErr,opt-in 之后可以查到具体错误。

错误记录挂在 WithStreamErrors 返回的 ctx 的对象图里(context.WithValue),
不是全局表——调用方不再引用 ctx/channel 时会被 GC 自然回收,不需要
任何显式清理逻辑,也不依赖 ctx.Done(),即使用 context.Background()
也能正常释放;ctx 之后被别的 context.With*(包括 StickyCtx)再包一层
也不影响查询。

同时补充完整的自动化测试覆盖 example/main.go 里演示过的所有功能:
四种调用模式 × int/struct/slice/[]byte 的组合(client_test.go)、
WithHandlers/call_go 全双工(handlers_test.go)、NewSession 隔离性
与 StickyCtx 路由(session_test.go),之前这些只能靠人肉跑 go run
看输出,现在都有真实断言。
This commit is contained in:
2026-07-23 16:24:05 +08:00
parent 6ccec66a1e
commit ee5e5b96af
9 changed files with 717 additions and 5 deletions
+45
View File
@@ -242,6 +242,51 @@ if ctx.Err() != nil {
完整可运行示例见 [example/main.go](example/main.go) 中的 `demoTimeout`(默认超时 / 显式 deadline 优先级 / 流式超时静默关闭 channel)和 `demoBlocking`(池子占满后阻塞排队)。 完整可运行示例见 [example/main.go](example/main.go) 中的 `demoTimeout`(默认超时 / 显式 deadline 优先级 / 流式超时静默关闭 channel)和 `demoBlocking`(池子占满后阻塞排队)。
### 查询流式调用执行过程中的错误(WithStreamErrors
流式输出/双向流的 handler(Python 生成器)如果**执行过程中**抛异常(不是等全部 yield 完才失败,是还没跑完就失败),默认情况下调用方完全无法感知——`Invoke[chan T]` 返回的那个 `err` 只描述"调用有没有成功发起"(参数序列化、抢连接、写 `call` 消息),跟 Python 函数体最终是正常结束还是执行过程中抛异常没有任何关系,因为 handler 真正执行是在 `Invoke` 返回之后的后台 goroutine 里才发生的。Python 侧异常只会让 `ch` 静默提前关闭,现象上跟"正常读完"一模一样,`ctx.Err()` 也查不出来(不是 ctx 取消导致的)。
想知道是不是这种情况,用 `WithStreamErrors` 包一层 ctx(opt-in),它会额外返回一个查询函数,之后不需要再传 ctx:
```go
ctx, streamErr := gobridge.WithStreamErrors(ctx) // opt-in,不包这一层就是上面说的默认行为
ch, err := gobridge.Invoke[chan int](ctx, pool, "range_gen", 1, 10)
if err != nil {
// 建立阶段失败,比如池子占满/参数错误
}
for v := range ch {
fmt.Println(v)
}
if err := streamErr(ch); err != nil {
// channel 是因为 Python 侧执行过程中抛异常提前关闭的,不是正常 yield 完
}
```
错误记录挂在 `WithStreamErrors` 返回的那个 `ctx` 的对象图里(通过 `context.WithValue` 携带一个可变的存储,`streamErr` 直接持有这块存储,不需要再查 ctx),不是全局表——调用方不再引用这个 `ctx`(和对应的 `ch`)时,整条链会被 GC 自然回收,**不需要任何显式清理逻辑,也不依赖 `ctx.Done()`**,即使用 `context.Background()` 也能正常释放。`WithStreamErrors` 之后即使 `ctx` 又被别的 `context.With*`(包括库里自己的 `StickyCtx`)再包一层,也不影响——`context.Value()` 的查找是逐层往上找的,不会因为后续再包装而丢失(`streamErr` 是从最初那次 `WithStreamErrors` 调用里直接拿到的闭包,跟后续怎么包 ctx 完全无关)。
不调用 `WithStreamErrors` 是完全零成本的默认行为,现有的流式调用代码不需要做任何改动。完整可运行示例见 `example/main.go``demoTimeout` 的示例5。
### 查询流式调用中途的错误(StreamError
流式输出/双向流的 handler(Python 生成器)如果**执行过程中**抛异常(不是等全部 yield 完才失败,是还没跑完就失败),默认情况下调用方完全无法感知——`Invoke[chan T]` 返回的那个 `err` 只描述"调用有没有成功发起"(参数序列化、抢连接、写 `call` 消息),跟 Python 函数体最终是正常结束还是中途抛异常没有任何关系,因为 handler 真正执行是在 `Invoke` 返回之后的后台 goroutine 里才发生的。Python 侧异常只会让 `ch` 静默提前关闭,现象上跟"正常读完"一模一样,`ctx.Err()` 也查不出来(不是 ctx 取消导致的)。
想知道是不是这种情况,用 `StreamError` 查——不需要任何额外配置,直接传入 `Invoke[chan T]` 返回的那个 channel
```go
ch, err := gobridge.Invoke[chan int](ctx, pool, "range_gen", 1, 10)
if err != nil {
// 建立阶段失败,比如池子占满/参数错误
}
for v := range ch {
fmt.Println(v)
}
if streamErr := gobridge.StreamError(ch); streamErr != nil {
// channel 是因为 Python 侧执行过程中抛异常提前关闭的,不是正常 yield 完
}
```
`StreamError` 内部用调用方传给 `Invoke` 的那个 `ctx` 挂了 `context.AfterFunc``ctx` 结束(取消/超时/调用方自己 `cancel()`)时会自动清理对应记录,不需要调用方主动查询才释放——前提是 ctx 本身会结束(`context.Background()` 永不 `Done()`,如果还从来不查 `StreamError`,记录会一直留着,这也是本文档反复强调"别用没有 deadline 的 ctx"的场景之一)。
## Session 亲和路由 ## Session 亲和路由
默认情况下,每次 `Invoke` 通过轮询分配 worker 进程。当多次调用需要共享同一 Python 进程的状态时,可以使用 Session 或 StickyCtx 将调用固定到同一进程。 默认情况下,每次 `Invoke` 通过轮询分配 worker 进程。当多次调用需要共享同一 Python 进程的状态时,可以使用 Session 或 StickyCtx 将调用固定到同一进程。
+6
View File
@@ -221,6 +221,9 @@ func invokeStreamOut[R any](ctx context.Context, pool Pool, method string, rt re
for { for {
msg, err := readResult(ctx, conn, pool, write) msg, err := readResult(ctx, conn, pool, write)
if err != nil || msg.Type == TypeEnd || msg.Type == TypeError { if err != nil || msg.Type == TypeEnd || msg.Type == TypeError {
if msg.Type == TypeError {
recordStreamError(ctx, ch.Interface(), fmt.Errorf("remote error: %s", msg.Error))
}
return return
} }
if msg.Type == TypeChunk { if msg.Type == TypeChunk {
@@ -395,6 +398,9 @@ func invokeStreamBoth[R any](ctx context.Context, pool Pool, method string, stre
for { for {
msg, err := readResult(ctx, conn, pool, write) msg, err := readResult(ctx, conn, pool, write)
if err != nil || msg.Type == TypeEnd || msg.Type == TypeError { if err != nil || msg.Type == TypeEnd || msg.Type == TypeError {
if msg.Type == TypeError {
recordStreamError(ctx, outCh.Interface(), fmt.Errorf("remote error: %s", msg.Error))
}
return return
} }
if msg.Type == TypeChunk { if msg.Type == TypeChunk {
+207
View File
@@ -0,0 +1,207 @@
package gobridge
import (
"context"
"reflect"
"testing"
)
// testUser 对应 example/worker.py 里的 User dataclass。
type testUser struct {
ID int `json:"id"`
Name string `json:"name"`
Score float64 `json:"score"`
Level string `json:"level,omitempty"`
}
// TestInvokeModes 覆盖 example/main.go 的 demoPool 演示过的四种调用模式
// (普通调用、流式输出、流式输入、双向流)在 int / struct / slice / []byte
// 各种类型组合下的正确性,对应 README「四种调用模式」小节。
func TestInvokeModes(t *testing.T) {
pool := newTestPool(t, WithWorkers(2), WithMaxConns(4))
ctx := context.Background()
t.Run("普通调用_int", func(t *testing.T) {
sum, err := Invoke[int](ctx, pool, "add", 3, 4)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
if sum != 7 {
t.Fatalf("want 7, got %d", sum)
}
})
t.Run("流式输出_int", func(t *testing.T) {
ch, err := Invoke[chan int](ctx, pool, "range_gen", 1, 6)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
var got []int
for v := range ch {
got = append(got, v)
}
want := []int{1, 2, 3, 4, 5}
if !reflect.DeepEqual(got, want) {
t.Fatalf("want %v, got %v", want, got)
}
})
t.Run("流式输入_int", func(t *testing.T) {
inputCh := make(chan int, 10)
go func() {
for i := 1; i <= 5; i++ {
inputCh <- i
}
close(inputCh)
}()
total, err := Invoke[int](ctx, pool, "sum_stream", inputCh)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
if total != 15 {
t.Fatalf("want 15, got %d", total)
}
})
t.Run("双向流_int", func(t *testing.T) {
inputCh := make(chan int, 10)
go func() {
for i := 1; i <= 5; i++ {
inputCh <- i
}
close(inputCh)
}()
outCh, err := Invoke[chan int](ctx, pool, "double_stream", inputCh)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
var got []int
for v := range outCh {
got = append(got, v)
}
want := []int{1, 4, 9, 16, 25}
if !reflect.DeepEqual(got, want) {
t.Fatalf("want %v, got %v", want, got)
}
})
t.Run("普通调用_struct", func(t *testing.T) {
user, err := Invoke[testUser](ctx, pool, "get_user", 42)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
want := testUser{ID: 42, Name: "user_42", Score: 63}
if user != want {
t.Fatalf("want %+v, got %+v", want, user)
}
})
users := []testUser{
{ID: 1, Name: "alice", Score: 5.0},
{ID: 2, Name: "bob", Score: 8.0},
{ID: 3, Name: "carol", Score: 12.0},
}
t.Run("slice输入_返回标量", func(t *testing.T) {
scoreSum, err := Invoke[float64](ctx, pool, "total_score", users)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
if scoreSum != 25 {
t.Fatalf("want 25, got %v", scoreSum)
}
})
t.Run("slice输入输出", func(t *testing.T) {
enriched, err := Invoke[[]testUser](ctx, pool, "enrich_users", users)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
want := []testUser{
{ID: 1, Name: "alice", Score: 5.0, Level: "silver"},
{ID: 2, Name: "bob", Score: 8.0, Level: "silver"},
{ID: 3, Name: "carol", Score: 12.0, Level: "gold"},
}
if !reflect.DeepEqual(enriched, want) {
t.Fatalf("want %+v, got %+v", want, enriched)
}
})
t.Run("流式输出_struct", func(t *testing.T) {
userCh, err := Invoke[chan testUser](ctx, pool, "gen_users", 3)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
var got []testUser
for u := range userCh {
got = append(got, u)
}
want := []testUser{
{ID: 1, Name: "user_1", Score: 3},
{ID: 2, Name: "user_2", Score: 6},
{ID: 3, Name: "user_3", Score: 9},
}
if !reflect.DeepEqual(got, want) {
t.Fatalf("want %+v, got %+v", want, got)
}
})
t.Run("双向流_struct", func(t *testing.T) {
inCh := make(chan testUser, len(users))
go func() {
for _, u := range users {
inCh <- u
}
close(inCh)
}()
procCh, err := Invoke[chan testUser](ctx, pool, "process_users", inCh)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
var got []testUser
for u := range procCh {
got = append(got, u)
}
want := []testUser{
{ID: 1, Name: "ALICE", Score: 10},
{ID: 2, Name: "BOB", Score: 16},
{ID: 3, Name: "CAROL", Score: 24},
}
if !reflect.DeepEqual(got, want) {
t.Fatalf("want %+v, got %+v", want, got)
}
})
t.Run("bytes_输入输出", func(t *testing.T) {
rev, err := Invoke[[]byte](ctx, pool, "bytes_reverse", []byte("hello"))
if err != nil {
t.Fatalf("Invoke: %v", err)
}
if string(rev) != "olleh" {
t.Fatalf("want olleh, got %s", rev)
}
cat, err := Invoke[[]byte](ctx, pool, "bytes_concat", []byte("foo"), []byte("bar"))
if err != nil {
t.Fatalf("Invoke: %v", err)
}
if string(cat) != "foobar" {
t.Fatalf("want foobar, got %s", cat)
}
})
t.Run("bytes_流式输出", func(t *testing.T) {
bCh, err := Invoke[chan []byte](ctx, pool, "bytes_chunks", []byte("abcdefgh"), 3)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
var got []string
for chunk := range bCh {
got = append(got, string(chunk))
}
want := []string{"abc", "def", "gh"}
if !reflect.DeepEqual(got, want) {
t.Fatalf("want %v, got %v", want, got)
}
})
}
+22
View File
@@ -336,6 +336,28 @@ func demoTimeout(script string) {
fmt.Print(" ", v) fmt.Print(" ", v)
} }
fmt.Println() // 应该完整输出 1~9,不会被 500ms 默认超时打断 fmt.Println() // 应该完整输出 1~9,不会被 500ms 默认超时打断
// ── 示例5WithStreamErrors——流式 handler 执行过程中抛异常,channel 会静默
// 提前关闭,Invoke 本身的 err 只描述"调用有没有发起成功",跟这个异常无关
// (永远是 nil)。想知道流是不是因为 Python 侧异常提前结束,需要先用
// WithStreamErrors 包一层 ctxopt-in),拿到的 streamErr 函数不用再传 ctx。
//
// 错误记录挂在 WithStreamErrors 返回的这个 ctx 的对象图里,调用方不再引用
// ctx5/ch3 时会被 GC 自然回收,不需要任何显式清理,即使用 context.Background()
// 也一样能正常释放(不依赖 ctx.Done())。
ctx5, streamErr := gobridge.WithStreamErrors(context.Background())
ch3, err := gobridge.Invoke[chan int](ctx5, pool, "stream_then_raise", 3)
if err != nil {
log.Fatal(err) // 这里的 err 只可能是"发起调用失败",不会是 stream_then_raise 里的异常
}
fmt.Print("stream_then_raise(3)(执行过程中会抛异常)=")
for v := range ch3 {
fmt.Print(" ", v)
}
fmt.Println()
if err := streamErr(ch3); err != nil {
fmt.Println("streamErr 查到执行过程中的异常:", err)
}
} }
func demoBlocking(script string) { func demoBlocking(script string) {
+4 -2
View File
@@ -65,10 +65,12 @@ def double_stream(numbers: Iterator[int]) -> Iterator[int]:
@expose @expose
def stream_then_raise(n: int) -> Iterator[int]: def stream_then_raise(n: int) -> Iterator[int]:
"""流式输出中途抛异常,用于复现 _dispatch 里 end/error 消息错位的问题""" """流式输出执行过程中抛异常(还没 yield 完就失败),
用于复现 _dispatch 里 end/error 消息错位的问题"""
for i in range(n): for i in range(n):
if i == n - 1:
raise ValueError(f"boom: stream_then_raise 在第 {i} 个元素时失败")
yield i yield i
raise ValueError("boom: stream_then_raise 中途失败")
# ── structdataclass / dict)类型 ─────────────────────────────────────────── # ── structdataclass / dict)类型 ───────────────────────────────────────────
+125
View File
@@ -0,0 +1,125 @@
package gobridge
import (
"context"
"fmt"
"reflect"
"sync"
"testing"
)
// testGoService 对应 example/main.go 里的 goService,实现 Handler 接口,
// 公开方法自动暴露给 Python 通过 call_go() 调用。用于覆盖 example/main.go
// 的 demoServer 演示过的 WithHandlers/call_go 全双工功能。
type testGoService struct {
pool Pool // 供 EnrichName 内部再调 Python
mu sync.Mutex
logLines []string
}
func (s *testGoService) Multiply(ctx context.Context, a, b int) (int, error) {
return a * b, nil
}
func (s *testGoService) Log(msg string) {
s.mu.Lock()
defer s.mu.Unlock()
s.logLines = append(s.logLines, msg)
}
func (s *testGoService) takeLogLines() []string {
s.mu.Lock()
defer s.mu.Unlock()
lines := s.logLines
s.logLines = nil
return lines
}
// EnrichName 内部通过 Invoke 调用 Python 的 to_upper,构成 Go→Python→Go→Python 四层链路。
func (s *testGoService) EnrichName(ctx context.Context, name string) (string, error) {
upper, err := Invoke[string](ctx, s.pool, "to_upper", name)
if err != nil {
return "", err
}
return "Hello, " + upper + "!", nil
}
func (s *testGoService) MakeUser(ctx context.Context, uid int) (testUser, error) {
return testUser{ID: uid, Name: fmt.Sprintf("user_%d", uid), Score: float64(uid) * 1.5}, nil
}
func newTestServerPool(t *testing.T, svc *testGoService) Pool {
t.Helper()
pool := newTestPool(t, WithWorkers(1), WithHandlers(svc))
svc.pool = pool
return pool
}
// TestHandlers 覆盖 example/main.go 的 demoServer 演示过的 WithHandlers/call_go
// 全双工功能:Python 调用 Go 方法、流式输出过程中回调 Go、Go→Python→Go→Python
// 四层链路、call_go[T] 自动构造 dataclass。
func TestHandlers(t *testing.T) {
svc := &testGoService{}
pool := newTestServerPool(t, svc)
ctx := context.Background()
t.Run("Python调用Go方法", func(t *testing.T) {
result, err := Invoke[int](ctx, pool, "compute_with_go_mul", 6, 7)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
if result != 42 {
t.Fatalf("want 42, got %d", result)
}
})
t.Run("流式输出过程中回调Go", func(t *testing.T) {
ch, err := Invoke[chan int](ctx, pool, "squared_with_log", 4)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
var got []int
for v := range ch {
got = append(got, v)
}
want := []int{1, 4, 9, 16}
if !reflect.DeepEqual(got, want) {
t.Fatalf("want %v, got %v", want, got)
}
wantLogs := []string{
"yielding 1² = 1",
"yielding 2² = 4",
"yielding 3² = 9",
"yielding 4² = 16",
}
if gotLogs := svc.takeLogLines(); !reflect.DeepEqual(gotLogs, wantLogs) {
t.Fatalf("want log lines %v, got %v", wantLogs, gotLogs)
}
})
t.Run("四层全双工链路", func(t *testing.T) {
// full_chain("world") → call_go[str]("EnrichName","world")
// → Invoke[string](ctx, serv, "to_upper", "world") → "WORLD"
// ← "Hello, WORLD!"
greeting, err := Invoke[string](ctx, pool, "full_chain", "world")
if err != nil {
t.Fatalf("Invoke: %v", err)
}
if greeting != "Hello, WORLD!" {
t.Fatalf("want %q, got %q", "Hello, WORLD!", greeting)
}
})
t.Run("call_go自动构造dataclass", func(t *testing.T) {
enriched, err := Invoke[testUser](ctx, pool, "get_user_via_go", 12)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
want := testUser{ID: 12, Name: "user_12", Score: 18, Level: "gold"}
if enriched != want {
t.Fatalf("want %+v, got %+v", want, enriched)
}
})
}
+97 -3
View File
@@ -7,6 +7,7 @@ import (
"os/exec" "os/exec"
"path/filepath" "path/filepath"
"runtime" "runtime"
"strings"
"sync" "sync"
"testing" "testing"
"time" "time"
@@ -40,6 +41,26 @@ func newTestPool(t *testing.T, opts ...Option) Pool {
return pool return pool
} }
// stableThreadCount 轮询 count_threads(),直到连续两次读数一致再返回,
// 用于绕开 pool 刚创建时 Python 侧连接处理线程还没必然全部起稳的竞态。
func stableThreadCount(t *testing.T, ctx context.Context, pool Pool) (int, error) {
t.Helper()
prev := -1
deadline := time.Now().Add(3 * time.Second)
for time.Now().Before(deadline) {
n, err := Invoke[int](ctx, pool, "count_threads")
if err != nil {
return 0, err
}
if n == prev {
return n, nil
}
prev = n
time.Sleep(50 * time.Millisecond)
}
return prev, nil
}
// TestPoolExhaustionBlocks 验证:workers=2、maxConns=2(总容量 4)的池子被占满后, // TestPoolExhaustionBlocks 验证:workers=2、maxConns=2(总容量 4)的池子被占满后,
// 新的 Invoke 调用会排队阻塞等待连接释放,而不是立刻失败或被跳过; // 新的 Invoke 调用会排队阻塞等待连接释放,而不是立刻失败或被跳过;
// 一旦有连接释放,新调用能正常拿到连接并执行。 // 一旦有连接释放,新调用能正常拿到连接并执行。
@@ -162,8 +183,8 @@ func TestStreamErrorMidwayCorruptsNextCall(t *testing.T) {
for v := range ch { for v := range ch {
got = append(got, v) got = append(got, v)
} }
if len(got) != 3 { if len(got) != 2 {
t.Fatalf("want 3 items before the raise, got %v", got) t.Fatalf("want 2 items before the raise, got %v", got)
} }
// 连接理论上已经被当作"健康"放回池子;下一次调用应该拿到自己的正常结果, // 连接理论上已经被当作"健康"放回池子;下一次调用应该拿到自己的正常结果,
@@ -177,6 +198,76 @@ func TestStreamErrorMidwayCorruptsNextCall(t *testing.T) {
} }
} }
// TestStreamError 验证 WithStreamErrorsopt-in 之后,流式调用执行过程中因为 Python
// 侧真实抛出的异常提前结束时,能通过它返回的 streamErr 函数查到具体错误(不需要再传
// ctx);正常跑完的流式调用查询结果是 nil。
//
// 这个方案没有全局登记表——错误记录挂在 WithStreamErrors 返回的那个 ctx 的对象图里,
// 调用方不再引用这个 ctx(和对应的 ch)时,整条链会被 GC 自然回收,不需要任何显式
// 清理逻辑,也不依赖 ctx.Done()context.Background() 一样能正常释放。
func TestStreamError(t *testing.T) {
pool := newTestPool(t, WithWorkers(1), WithMaxConns(2))
t.Run("opt-in 后能查到 Python 侧真实抛出的异常", func(t *testing.T) {
ctx, streamErr := WithStreamErrors(context.Background())
// stream_then_raise 是 example/worker.py 里真实的 Python 函数:
// 在还没 yield 完 n 个数之前,执行过程中就 raise ValueError("boom: ...")
// 不是模拟出来的错误,也不是等全部 yield 完才失败。
ch, err := Invoke[chan int](ctx, pool, "stream_then_raise", 3)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
var got []int
for v := range ch {
got = append(got, v)
}
if len(got) != 2 {
t.Fatalf("want 2 items before the raise, got %v", got)
}
err = streamErr(ch)
if err == nil {
t.Fatal("want non-nil error after Python 侧执行过程中抛异常")
}
if !strings.Contains(err.Error(), "boom") {
t.Fatalf("want error containing %q, got %v", "boom", err)
}
t.Logf("streamErr 查到的错误: %v", err)
})
t.Run("正常跑完不会有错误", func(t *testing.T) {
ctx, streamErr := WithStreamErrors(context.Background())
ch, err := Invoke[chan int](ctx, pool, "range_gen", 1, 4)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
var got []int
for v := range ch {
got = append(got, v)
}
if len(got) != 3 {
t.Fatalf("want 3 items, got %v", got)
}
if err := streamErr(ch); err != nil {
t.Fatalf("want nilgot %v", err)
}
})
t.Run("再包一层 ctx(比如加超时)不会丢失 WithStreamErrors 记录的信息", func(t *testing.T) {
ctx, streamErr := WithStreamErrors(context.Background())
ctx, cancel := context.WithTimeout(ctx, 5*time.Second) // WithStreamErrors 之后又包了一层
defer cancel()
ch, err := Invoke[chan int](ctx, pool, "stream_then_raise", 3)
if err != nil {
t.Fatalf("Invoke: %v", err)
}
for range ch {
}
if err := streamErr(ch); err == nil {
t.Fatal("再包一层 WithTimeout 之后应该还能查到 WithStreamErrors 记录的错误")
}
})
}
// TestStreamOutIgnoresDefaultTimeout 验证 WithDefaultTimeout 不会套用到流式输出/ // TestStreamOutIgnoresDefaultTimeout 验证 WithDefaultTimeout 不会套用到流式输出/
// 双向流调用上:流式聊天这类响应可能持续很久,如果套用全局默认超时会在正常输出 // 双向流调用上:流式聊天这类响应可能持续很久,如果套用全局默认超时会在正常输出
// 过程中把它腰斩。没有显式传入 deadline 时,流式调用应该完全不受池子默认超时影响, // 过程中把它腰斩。没有显式传入 deadline 时,流式调用应该完全不受池子默认超时影响,
@@ -215,7 +306,10 @@ func TestStreamInputCancelDoesNotLeakThread(t *testing.T) {
pool := newTestPool(t, WithWorkers(1), WithMaxConns(2)) pool := newTestPool(t, WithWorkers(1), WithMaxConns(2))
ctx := context.Background() ctx := context.Background()
baseline, err := Invoke[int](ctx, pool, "count_threads") // 刚创建的 pool,Python 那边给每条预建连接起的 _handle_conn/_reader 线程
// 还没必然全部起稳(server.accept() 是异步接受的),第一次量出来的线程数
// 可能偏低导致误判。轮询到连续两次读数一致再当作稳定的 baseline。
baseline, err := stableThreadCount(t, ctx, pool)
if err != nil { if err != nil {
t.Fatalf("Invoke count_threads (baseline): %v", err) t.Fatalf("Invoke count_threads (baseline): %v", err)
} }
+149
View File
@@ -0,0 +1,149 @@
package gobridge
import (
"context"
"fmt"
"reflect"
"testing"
)
// sessionResult 对应 example/worker.py 里 session_result 返回的
// {"value": int, "steps": [int, ...]}。
type sessionResult struct {
Value int `json:"value"`
Steps []int `json:"steps"`
}
// parseWorkerID 从 "[worker N] ..." 格式的字符串里解析出 worker 编号,
// example/worker.py 里 session_init/global_increment/global_get 都是这个格式。
func parseWorkerID(t *testing.T, msg string) int {
t.Helper()
var id int
if _, err := fmt.Sscanf(msg, "[worker %d]", &id); err != nil {
t.Fatalf("parse worker id from %q: %v", msg, err)
}
return id
}
// TestNewSessionIsolation 覆盖 example/main.go 的 demoSession 演示过的
// NewSession 功能:不同 session 之间状态互相隔离;同一个 worker 进程内的
// session 共享进程级全局变量,不同 worker 之间完全独立。
func TestNewSessionIsolation(t *testing.T) {
pool := newTestPool(t, WithWorkers(2))
ctx := context.Background()
// NewSession 内部按 pool 共享的原子计数器轮询:第一次 idx=1%2=1,
// 第二次 idx=2%2=0,第三次 idx=3%2=1——sessA 和 sessC 落在同一个 worker
// sessB 落在另一个(用全新 pool,不掺杂其它调用,保证这个顺序确定)。
sessA := NewSession(pool)
sessB := NewSession(pool)
sessC := NewSession(pool)
if _, err := Invoke[string](ctx, sessA, "session_init", "A", 100); err != nil {
t.Fatalf("sessA init: %v", err)
}
if _, err := Invoke[string](ctx, sessB, "session_init", "B", 200); err != nil {
t.Fatalf("sessB init: %v", err)
}
if _, err := Invoke[string](ctx, sessC, "session_init", "C", 300); err != nil {
t.Fatalf("sessC init: %v", err)
}
// 各自 step,验证会话状态互不干扰
if v, err := Invoke[int](ctx, sessA, "session_step", "A", 10); err != nil || v != 110 {
t.Fatalf("sessA step(+10): v=%d err=%v, want 110", v, err)
}
if v, err := Invoke[int](ctx, sessA, "session_step", "A", 5); err != nil || v != 115 {
t.Fatalf("sessA step(+5): v=%d err=%v, want 115", v, err)
}
if v, err := Invoke[int](ctx, sessB, "session_step", "B", 50); err != nil || v != 250 {
t.Fatalf("sessB step(+50): v=%d err=%v, want 250", v, err)
}
if v, err := Invoke[int](ctx, sessC, "session_step", "C", 99); err != nil || v != 399 {
t.Fatalf("sessC step(+99): v=%d err=%v, want 399(与 sessA 同 worker 但状态独立)", v, err)
}
rA, err := Invoke[sessionResult](ctx, sessA, "session_result", "A")
if err != nil {
t.Fatalf("sessA result: %v", err)
}
if want := (sessionResult{Value: 115, Steps: []int{10, 5}}); !reflect.DeepEqual(rA, want) {
t.Fatalf("sessA result: want %+v, got %+v", want, rA)
}
rB, err := Invoke[sessionResult](ctx, sessB, "session_result", "B")
if err != nil {
t.Fatalf("sessB result: %v", err)
}
if want := (sessionResult{Value: 250, Steps: []int{50}}); !reflect.DeepEqual(rB, want) {
t.Fatalf("sessB result: want %+v, got %+v", want, rB)
}
rC, err := Invoke[sessionResult](ctx, sessC, "session_result", "C")
if err != nil {
t.Fatalf("sessC result: %v", err)
}
if want := (sessionResult{Value: 399, Steps: []int{99}}); !reflect.DeepEqual(rC, want) {
t.Fatalf("sessC result: want %+v, got %+v", want, rC)
}
// 进程级全局变量:sessA 和 sessC 同一个 worker 进程,应该共享 _global_counter
// sessB 是独立进程,从 0 开始,不受影响。
rawA, err := Invoke[string](ctx, sessA, "global_increment", 10)
if err != nil {
t.Fatalf("sessA global_increment: %v", err)
}
workerA := parseWorkerID(t, rawA)
rawC, err := Invoke[string](ctx, sessC, "global_increment", 5)
if err != nil {
t.Fatalf("sessC global_increment: %v", err)
}
workerC := parseWorkerID(t, rawC)
rawB, err := Invoke[string](ctx, sessB, "global_increment", 99)
if err != nil {
t.Fatalf("sessB global_increment: %v", err)
}
workerB := parseWorkerID(t, rawB)
if workerA != workerC {
t.Fatalf("want sessA/sessC 落在同一个 workergot worker %d / %d", workerA, workerC)
}
if workerB == workerA {
t.Fatalf("want sessB 落在独立的 workergot 跟 sessA/sessC 一样都是 worker %d", workerB)
}
wantA := fmt.Sprintf("[worker %d] counter = 15", workerA) // sessA(+10) + sessC(+5) 共享
if rawA2, err := Invoke[string](ctx, sessA, "global_get"); err != nil || rawA2 != wantA {
t.Fatalf("sessA global_get: want %q, got %q (err=%v)", wantA, rawA2, err)
}
wantB := fmt.Sprintf("[worker %d] counter = 99", workerB) // 独立进程,只有自己的 +99
if rawB != wantB {
t.Fatalf("sessB global_increment: want %q, got %q", wantB, rawB)
}
}
// TestStickyCtx 覆盖 example/main.go 的 demoSession 演示过的 StickyCtx 功能:
// 相同的亲和 key 无论调用多少次,都应该稳定路由到同一个 worker 进程。
func TestStickyCtx(t *testing.T) {
pool := newTestPool(t, WithWorkers(2))
ctx := StickyCtx(context.Background(), "sticky-key")
var firstWorker int
for i := 0; i < 4; i++ {
msg, err := Invoke[string](ctx, pool, "session_init", fmt.Sprintf("aff-%d", i), i)
if err != nil {
t.Fatalf("Invoke #%d: %v", i, err)
}
worker := parseWorkerID(t, msg)
if i == 0 {
firstWorker = worker
continue
}
if worker != firstWorker {
t.Fatalf("StickyCtx 路由不稳定:第 1 次落在 worker %d,第 %d 次落在 worker %d", firstWorker, i+1, worker)
}
}
}
+62
View File
@@ -0,0 +1,62 @@
package gobridge
import (
"context"
"sync"
)
// streamErrBox 是挂在 ctx.Value 树上的一块可变存储,key 是具体的 channel 值本身。
// 生命周期完全跟着调用方持有的 ctx/channel 走,不需要任何显式清理:
// 调用方不再引用它们时,整棵对象图(ctx → streamErrBox → errs → error
// 会被 GC 自然回收——不依赖 ctx.Done()context.Background() 一样能正常释放。
type streamErrBox struct {
mu sync.Mutex
errs map[any]error
}
func (b *streamErrBox) get(ch any) error {
b.mu.Lock()
defer b.mu.Unlock()
return b.errs[ch]
}
func (b *streamErrBox) set(ch any, err error) {
b.mu.Lock()
defer b.mu.Unlock()
b.errs[ch] = err
}
type streamErrBoxKey struct{}
// WithStreamErrors 返回一个包过的 ctx,以及一个用于查询流式调用错误的函数 streamErr。
// 之后用这个 ctx(或它的子 ctx,即使被别的 context.With* 再包一层也不受影响)发起的
// 流式输出/双向流调用,如果 Python handler 执行过程中抛异常,可以用 streamErr(ch)
// 查到具体错误,不需要再传一次 ctx。
//
// 不调用这个函数是默认行为:流式调用中途出错只会让 channel 静默提前关闭,
// 调用方拿不到任何错误信息。这是因为 Invoke[chan T] 在建立调用阶段就已经成功
// 返回了 (ch, nil),后续 Python 生成器内部的异常没有天然的返回值可以携带。
//
// ctx, streamErr := gobridge.WithStreamErrors(ctx)
// ch, err := gobridge.Invoke[chan int](ctx, pool, "range_gen", 1, 10)
// for v := range ch {
// fmt.Println(v)
// }
// if err := streamErr(ch); err != nil {
// // channel 是因为 Python 侧异常提前关闭的,而不是正常 yield 完
// }
func WithStreamErrors(parent context.Context) (ctx context.Context, streamErr func(ch any) error) {
box := &streamErrBox{errs: make(map[any]error)}
ctx = context.WithValue(parent, streamErrBoxKey{}, box)
return ctx, box.get
}
// recordStreamError 供 invokeStreamOut/invokeStreamBoth 在读到 TypeError 时调用。
// 如果调用方没有用 WithStreamErrors 包过 ctx,直接是个空操作。
func recordStreamError(ctx context.Context, ch any, err error) {
box, ok := ctx.Value(streamErrBoxKey{}).(*streamErrBox)
if !ok {
return
}
box.set(ch, err)
}