Files
gobridge/client.go
T
what ee5e5b96af 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
看输出,现在都有真实断言。
2026-07-23 16:24:05 +08:00

418 lines
11 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Package gobridge 提供 Go 与 Python 之间的双向通信桥接,
// 支持普通调用、流式输出、流式输入和双向流四种模式。
package gobridge
import (
"context"
"encoding/json"
"fmt"
"log"
"net"
"reflect"
"sync"
)
// applyDefaultTimeout 在 ctx 未设置 deadline 时,套用 pool 配置的默认超时(WithDefaultTimeout)。
// 调用方显式设置的 deadline 优先级更高,不会被覆盖;未配置默认超时时返回原 ctx 和 no-op cancel。
func applyDefaultTimeout(ctx context.Context, pool Pool) (context.Context, context.CancelFunc) {
if _, ok := ctx.Deadline(); ok {
return ctx, func() {}
}
d := pool.defaultTimeout()
if d <= 0 {
return ctx, func() {}
}
return context.WithTimeout(ctx, d)
}
// Invoke 调用 Python 暴露的函数,支持四种模式:
//
// 普通调用: Invoke[int](ctx, pool, "Add", 3, 4)
// 流式输出: Invoke[chan int](ctx, pool, "RangeGen", 1, 10) // Python yield → Go channel
// 流式输入: Invoke[int](ctx, pool, "SumStream", inputChan) // Go channel → Python Iterator
// 双向流: Invoke[chan int](ctx, pool, "Transform", inputChan) // 两端均为流
//
// ctx 取消时会立即中断与 Python 的通信并返回 ctx.Err()。
// 对于流式输出/双向流,ctx 取消会关闭返回的 channel。
func Invoke[R any](ctx context.Context, pool Pool, method string, args ...any) (R, error) {
rt := reflect.TypeFor[R]()
// 查找 chan 类型的输入参数
streamArgIdx := -1
var streamCh reflect.Value
for i, arg := range args {
if arg != nil {
rv := reflect.ValueOf(arg)
if rv.Kind() == reflect.Chan {
streamArgIdx = i
streamCh = rv
break
}
}
}
switch {
case rt.Kind() == reflect.Chan && streamArgIdx >= 0:
return invokeStreamBoth[R](ctx, pool, method, streamArgIdx, streamCh, rt, args...)
case rt.Kind() == reflect.Chan:
return invokeStreamOut[R](ctx, pool, method, rt, args...)
case streamArgIdx >= 0:
return invokeStreamIn[R](ctx, pool, method, streamArgIdx, streamCh, args...)
default:
return invokeRegular[R](ctx, pool, method, args...)
}
}
// watchCtx 启动一个 goroutine 监听 ctx
// - ctx 取消时先发送 cancel 消息(Python 侧收到后注入 InterruptedError
// - 再关闭连接,解除阻塞中的读写操作
//
// write 是调用方提供的互斥写函数,保证与其他写操作不并发。
// 返回 stop 函数,必须在 conn 归还连接池前调用,可安全多次调用。
func watchCtx(ctx context.Context, conn net.Conn, id uint64, write func(Message)) (stop func()) {
done := make(chan struct{})
var once sync.Once
go func() {
select {
case <-ctx.Done():
write(Message{ID: id, Type: TypeCancel})
conn.Close()
case <-done:
}
}()
return func() { once.Do(func() { close(done) }) }
}
// chanRecv 从 ch 接收一个值,同时监听 ctx.Done()。
// 返回 (值, channel是否open, ctx是否已取消)。
func chanRecv(ctx context.Context, ch reflect.Value) (reflect.Value, bool, bool) {
chosen, val, ok := reflect.Select([]reflect.SelectCase{
{Dir: reflect.SelectRecv, Chan: ch},
{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(ctx.Done())},
})
if chosen == 1 {
return reflect.Value{}, false, true
}
return val, ok, false
}
// contextErr 在 io 错误时优先返回 ctx 的错误原因
func contextErr(ctx context.Context, err error) error {
if e := ctx.Err(); e != nil {
return e
}
return err
}
// readResult 读取下一条非 callback 消息,期间内联处理所有 Python→Go 回调。
// 保证 py→go→py→go→... 全链路复用同一条连接,不产生额外线程。
// write 是调用方提供的互斥写函数,与 watchCtx 共享同一把锁,避免并发写。
func readResult(ctx context.Context, conn net.Conn, pool Pool, write func(Message)) (Message, error) {
for {
msg, err := readMsg(conn)
if err != nil {
return Message{}, err
}
if msg.Type != TypeCallback {
return msg, nil
}
result, errStr := pool.callbackDispatch(ctx, msg)
var resp Message
if errStr != "" {
log.Printf("gobridge: handler %s error: %s", msg.Method, errStr)
resp = Message{ID: msg.ID, Type: TypeError, Error: errStr}
} else {
data, _ := json.Marshal(result)
resp = Message{ID: msg.ID, Type: TypeCallbackResult, Data: data}
}
write(resp)
}
}
func invokeRegular[R any](ctx context.Context, pool Pool, method string, args ...any) (R, error) {
var zero R
ctx, cancel := applyDefaultTimeout(ctx, pool)
defer cancel()
argsJSON, err := json.Marshal(args)
if err != nil {
return zero, fmt.Errorf("marshal args: %w", err)
}
conn, w, err := pool.acquire(ctx)
if err != nil {
return zero, err
}
var mu sync.Mutex
write := func(msg Message) { mu.Lock(); writeMsg(conn, msg); mu.Unlock() } //nolint
id := pool.nextReqID()
stop := watchCtx(ctx, conn, id, write)
defer stop()
mu.Lock()
err = writeMsg(conn, Message{ID: id, Type: TypeCall, Method: method, Args: argsJSON})
mu.Unlock()
if err != nil {
w.release(conn, false)
return zero, contextErr(ctx, fmt.Errorf("write call: %w", err))
}
resp, err := readResult(ctx, conn, pool, write)
if err != nil {
w.release(conn, false)
return zero, contextErr(ctx, fmt.Errorf("read response: %w", err))
}
stop()
w.release(conn, true)
if resp.Type == TypeError {
return zero, fmt.Errorf("remote error: %s", resp.Error)
}
var result R
if err := json.Unmarshal(resp.Data, &result); err != nil {
return zero, fmt.Errorf("unmarshal result: %w", err)
}
return result, nil
}
// invokeStreamOut 不套用 WithDefaultTimeout:返回的 channel 生命周期由调用方通过
// range 消费决定,可能持续很久(比如流式聊天回复),套一个全局默认超时会在正常
// 流式输出过程中把它腰斩。想要超时保护的话,调用方必须显式传入带 deadline 的 ctx。
func invokeStreamOut[R any](ctx context.Context, pool Pool, method string, rt reflect.Type, args ...any) (R, error) {
var zero R
argsJSON, err := json.Marshal(args)
if err != nil {
return zero, fmt.Errorf("marshal args: %w", err)
}
conn, w, err := pool.acquire(ctx)
if err != nil {
return zero, err
}
id := pool.nextReqID()
if err := writeMsg(conn, Message{
ID: id,
Type: TypeCall,
Method: method,
Args: argsJSON,
}); err != nil {
w.release(conn, false)
return zero, contextErr(ctx, fmt.Errorf("write call: %w", err))
}
ch := reflect.MakeChan(rt, 64)
go func() {
var mu sync.Mutex
write := func(msg Message) { mu.Lock(); writeMsg(conn, msg); mu.Unlock() } //nolint
stop := watchCtx(ctx, conn, id, write)
defer func() {
stop()
ch.Close()
w.release(conn, ctx.Err() == nil)
}()
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 {
val := reflect.New(rt.Elem())
if err := json.Unmarshal(msg.Data, val.Interface()); err != nil {
return
}
ch.Send(val.Elem())
}
}
}()
return ch.Interface().(R), nil
}
func invokeStreamIn[R any](ctx context.Context, pool Pool, method string, streamArgIdx int, streamCh reflect.Value, args ...any) (R, error) {
var zero R
ctx, cancel := applyDefaultTimeout(ctx, pool)
defer cancel()
jsonArgs := make([]any, len(args))
copy(jsonArgs, args)
jsonArgs[streamArgIdx] = nil
argsJSON, err := json.Marshal(jsonArgs)
if err != nil {
return zero, fmt.Errorf("marshal args: %w", err)
}
conn, w, err := pool.acquire(ctx)
if err != nil {
return zero, err
}
var mu sync.Mutex
writeErr := func(msg Message) error {
mu.Lock()
defer mu.Unlock()
return writeMsg(conn, msg)
}
write := func(msg Message) { writeErr(msg) } //nolint
id := pool.nextReqID()
stop := watchCtx(ctx, conn, id, write)
defer stop()
if err := writeErr(Message{
ID: id,
Type: TypeCall,
Method: method,
Args: argsJSON,
StreamInput: true,
StreamArgIdx: streamArgIdx,
}); err != nil {
w.release(conn, false)
return zero, contextErr(ctx, fmt.Errorf("write call: %w", err))
}
for {
val, ok, cancelled := chanRecv(ctx, streamCh)
if cancelled {
w.release(conn, false)
return zero, ctx.Err()
}
if !ok {
break
}
chunkData, err := json.Marshal(val.Interface())
if err != nil {
w.release(conn, false)
return zero, fmt.Errorf("marshal chunk: %w", err)
}
if err := writeErr(Message{ID: id, Type: TypeChunk, Data: chunkData}); err != nil {
w.release(conn, false)
return zero, contextErr(ctx, fmt.Errorf("write chunk: %w", err))
}
}
if err := writeErr(Message{ID: id, Type: TypeEnd}); err != nil {
w.release(conn, false)
return zero, contextErr(ctx, fmt.Errorf("write end: %w", err))
}
resp, err := readResult(ctx, conn, pool, write)
if err != nil {
w.release(conn, false)
return zero, contextErr(ctx, fmt.Errorf("read response: %w", err))
}
stop()
w.release(conn, true)
if resp.Type == TypeError {
return zero, fmt.Errorf("remote error: %s", resp.Error)
}
var result R
if err := json.Unmarshal(resp.Data, &result); err != nil {
return zero, fmt.Errorf("unmarshal result: %w", err)
}
return result, nil
}
// invokeStreamBoth 同样不套用 WithDefaultTimeout,理由同 invokeStreamOut——
// 双向流的生命周期由输入/输出两端共同决定,可能持续很久,需要超时保护时调用方
// 必须显式传入带 deadline 的 ctx。
func invokeStreamBoth[R any](ctx context.Context, pool Pool, method string, streamArgIdx int, streamCh reflect.Value, rt reflect.Type, args ...any) (R, error) {
var zero R
jsonArgs := make([]any, len(args))
copy(jsonArgs, args)
jsonArgs[streamArgIdx] = nil
argsJSON, err := json.Marshal(jsonArgs)
if err != nil {
return zero, fmt.Errorf("marshal args: %w", err)
}
conn, w, err := pool.acquire(ctx)
if err != nil {
return zero, err
}
id := pool.nextReqID()
if err := writeMsg(conn, Message{
ID: id,
Type: TypeCall,
Method: method,
Args: argsJSON,
StreamInput: true,
StreamArgIdx: streamArgIdx,
}); err != nil {
w.release(conn, false)
return zero, contextErr(ctx, fmt.Errorf("write call: %w", err))
}
outCh := reflect.MakeChan(rt, 64)
var mu sync.Mutex
write := func(msg Message) { mu.Lock(); writeMsg(conn, msg); mu.Unlock() } //nolint
// 写入 goroutine:输入 channel → Python chunks
go func() {
for {
val, ok, cancelled := chanRecv(ctx, streamCh)
if cancelled || !ok {
break
}
data, err := json.Marshal(val.Interface())
if err != nil {
break
}
mu.Lock()
err = writeMsg(conn, Message{ID: id, Type: TypeChunk, Data: data})
mu.Unlock()
if err != nil {
break
}
}
write(Message{ID: id, Type: TypeEnd})
}()
// 读取 goroutinePython chunks → 输出 channel,内联处理 callback
go func() {
stop := watchCtx(ctx, conn, id, write)
defer func() {
stop()
outCh.Close()
w.release(conn, ctx.Err() == nil)
}()
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 {
val := reflect.New(rt.Elem())
if err := json.Unmarshal(msg.Data, val.Interface()); err != nil {
return
}
outCh.Send(val.Elem())
}
}
}()
return outCh.Interface().(R), nil
}