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:
@@ -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 将调用固定到同一进程。
|
||||
|
||||
@@ -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 {
|
||||
|
||||
+207
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
+4
-2
@@ -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)类型 ───────────────────────────────────────────
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
+149
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user