From ee5e5b96affd23423212324314f4bcbdd4f2f8b5 Mon Sep 17 00:00:00 2001 From: what Date: Thu, 23 Jul 2026 16:24:05 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=20WithStreamErrors?= =?UTF-8?q?=20=E6=9F=A5=E8=AF=A2=E6=B5=81=E5=BC=8F=E8=B0=83=E7=94=A8?= =?UTF-8?q?=E6=89=A7=E8=A1=8C=E8=BF=87=E7=A8=8B=E4=B8=AD=E7=9A=84=E5=BC=82?= =?UTF-8?q?=E5=B8=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 流式输出/双向流的 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 看输出,现在都有真实断言。 --- README.md | 45 ++++++++++ client.go | 6 ++ client_test.go | 207 ++++++++++++++++++++++++++++++++++++++++++++++ example/main.go | 22 +++++ example/worker.py | 6 +- handlers_test.go | 125 ++++++++++++++++++++++++++++ pool_test.go | 100 +++++++++++++++++++++- session_test.go | 149 +++++++++++++++++++++++++++++++++ stream_error.go | 62 ++++++++++++++ 9 files changed, 717 insertions(+), 5 deletions(-) create mode 100644 client_test.go create mode 100644 handlers_test.go create mode 100644 session_test.go create mode 100644 stream_error.go diff --git a/README.md b/README.md index 25cb7d4..8cdff94 100644 --- a/README.md +++ b/README.md @@ -242,6 +242,51 @@ if ctx.Err() != nil { 完整可运行示例见 [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 亲和路由 默认情况下,每次 `Invoke` 通过轮询分配 worker 进程。当多次调用需要共享同一 Python 进程的状态时,可以使用 Session 或 StickyCtx 将调用固定到同一进程。 diff --git a/client.go b/client.go index 9fa549c..40047b3 100644 --- a/client.go +++ b/client.go @@ -221,6 +221,9 @@ func invokeStreamOut[R any](ctx context.Context, pool Pool, method string, rt re for { msg, err := readResult(ctx, conn, pool, write) 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 } if msg.Type == TypeChunk { @@ -395,6 +398,9 @@ func invokeStreamBoth[R any](ctx context.Context, pool Pool, method string, stre for { msg, err := readResult(ctx, conn, pool, write) 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 } if msg.Type == TypeChunk { diff --git a/client_test.go b/client_test.go new file mode 100644 index 0000000..5b5078e --- /dev/null +++ b/client_test.go @@ -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) + } + }) +} diff --git a/example/main.go b/example/main.go index 0c6b5fb..311b944 100644 --- a/example/main.go +++ b/example/main.go @@ -336,6 +336,28 @@ func demoTimeout(script string) { fmt.Print(" ", v) } fmt.Println() // 应该完整输出 1~9,不会被 500ms 默认超时打断 + + // ── 示例5:WithStreamErrors——流式 handler 执行过程中抛异常,channel 会静默 + // 提前关闭,Invoke 本身的 err 只描述"调用有没有发起成功",跟这个异常无关 + // (永远是 nil)。想知道流是不是因为 Python 侧异常提前结束,需要先用 + // WithStreamErrors 包一层 ctx(opt-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) { diff --git a/example/worker.py b/example/worker.py index a10af1a..79b5c24 100644 --- a/example/worker.py +++ b/example/worker.py @@ -65,10 +65,12 @@ def double_stream(numbers: Iterator[int]) -> Iterator[int]: @expose def stream_then_raise(n: int) -> Iterator[int]: - """流式输出中途抛异常,用于复现 _dispatch 里 end/error 消息错位的问题""" + """流式输出执行过程中抛异常(还没 yield 完就失败), + 用于复现 _dispatch 里 end/error 消息错位的问题""" for i in range(n): + if i == n - 1: + raise ValueError(f"boom: stream_then_raise 在第 {i} 个元素时失败") yield i - raise ValueError("boom: stream_then_raise 中途失败") # ── struct(dataclass / dict)类型 ─────────────────────────────────────────── diff --git a/handlers_test.go b/handlers_test.go new file mode 100644 index 0000000..26fd752 --- /dev/null +++ b/handlers_test.go @@ -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) + } + }) +} diff --git a/pool_test.go b/pool_test.go index 8a5a640..e6bfe69 100644 --- a/pool_test.go +++ b/pool_test.go @@ -7,6 +7,7 @@ import ( "os/exec" "path/filepath" "runtime" + "strings" "sync" "testing" "time" @@ -40,6 +41,26 @@ func newTestPool(t *testing.T, opts ...Option) 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)的池子被占满后, // 新的 Invoke 调用会排队阻塞等待连接释放,而不是立刻失败或被跳过; // 一旦有连接释放,新调用能正常拿到连接并执行。 @@ -162,8 +183,8 @@ func TestStreamErrorMidwayCorruptsNextCall(t *testing.T) { for v := range ch { got = append(got, v) } - if len(got) != 3 { - t.Fatalf("want 3 items before the raise, got %v", got) + if len(got) != 2 { + t.Fatalf("want 2 items before the raise, got %v", got) } // 连接理论上已经被当作"健康"放回池子;下一次调用应该拿到自己的正常结果, @@ -177,6 +198,76 @@ func TestStreamErrorMidwayCorruptsNextCall(t *testing.T) { } } +// TestStreamError 验证 WithStreamErrors:opt-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 nil,got %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 不会套用到流式输出/ // 双向流调用上:流式聊天这类响应可能持续很久,如果套用全局默认超时会在正常输出 // 过程中把它腰斩。没有显式传入 deadline 时,流式调用应该完全不受池子默认超时影响, @@ -215,7 +306,10 @@ func TestStreamInputCancelDoesNotLeakThread(t *testing.T) { pool := newTestPool(t, WithWorkers(1), WithMaxConns(2)) 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 { t.Fatalf("Invoke count_threads (baseline): %v", err) } diff --git a/session_test.go b/session_test.go new file mode 100644 index 0000000..08aa0a9 --- /dev/null +++ b/session_test.go @@ -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 落在同一个 worker,got worker %d / %d", workerA, workerC) + } + if workerB == workerA { + t.Fatalf("want sessB 落在独立的 worker,got 跟 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) + } + } +} diff --git a/stream_error.go b/stream_error.go new file mode 100644 index 0000000..411a568 --- /dev/null +++ b/stream_error.go @@ -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) +}