feat: add non-blocking per-connection handler execution

main v2.1.0
NoahLan 20 hours ago
parent e805b244d7
commit 604bd05e3b

@ -23,8 +23,10 @@ var _ RequestSetter = (*requestImpl)(nil)
// New 创建新的请求对象 // New 创建新的请求对象
func New(raw []byte, protocol protocolpkg.Protocol) requestpkg.Request { func New(raw []byte, protocol protocolpkg.Protocol) requestpkg.Request {
rawCopy := make([]byte, len(raw))
copy(rawCopy, raw)
return &requestImpl{ return &requestImpl{
raw: raw, raw: rawCopy,
protocol: protocol, protocol: protocol,
} }
} }

@ -18,6 +18,14 @@ func TestRequest(t *testing.T) {
assert.Equal(t, raw, req.Raw(), "Expected raw data to match") assert.Equal(t, raw, req.Raw(), "Expected raw data to match")
} }
func TestRequestCopiesRawData(t *testing.T) {
raw := []byte("test data")
req := New(raw, nil)
raw[0] = 'X'
assert.Equal(t, []byte("test data"), req.Raw())
}
func TestRequestBody(t *testing.T) { func TestRequestBody(t *testing.T) {
req := New([]byte("test"), nil).(*requestImpl) req := New([]byte("test"), nil).(*requestImpl)

@ -153,6 +153,8 @@ func newGnetServer(cfg *config.Config, protocol TransportProtocol) (*gnetServer,
} }
} }
ctx, cancel := context.WithCancel(context.Background())
// 创建事件处理器 // 创建事件处理器
handlerTimeout := configHelper.HandlerTimeout() handlerTimeout := configHelper.HandlerTimeout()
cloneHeader := configHelper.CloneHeader() cloneHeader := configHelper.CloneHeader()
@ -160,9 +162,8 @@ func newGnetServer(cfg *config.Config, protocol TransportProtocol) (*gnetServer,
protocolName := cfg.ApplicationProtocol protocolName := cfg.ApplicationProtocol
bootCh := make(chan error, 1) bootCh := make(chan error, 1)
bootOnce := &sync.Once{} bootOnce := &sync.Once{}
eventHandler := newUnifiedEventHandler(connManager, r, log, codecRegistry, codecResolverChain, defaultCodec, protocolManager, protocolName, enableProtocolEncode, cloneHeader, handlerTimeout, protocol, bootCh, bootOnce) handlerExec := newHandlerExecutor(cfg, log)
eventHandler := newUnifiedEventHandler(connManager, r, log, codecRegistry, codecResolverChain, defaultCodec, protocolManager, protocolName, enableProtocolEncode, cloneHeader, handlerTimeout, protocol, bootCh, bootOnce, ctx, handlerExec)
ctx, cancel := context.WithCancel(context.Background())
return &gnetServer{ return &gnetServer{
config: cfg, config: cfg,
@ -234,6 +235,7 @@ func (s *gnetServer) Start() error {
// 在goroutine中启动服务器 // 在goroutine中启动服务器
go func() { go func() {
err := gnet.Run(s.eventHandler, addr, options...) err := gnet.Run(s.eventHandler, addr, options...)
s.cancel()
if err != nil { if err != nil {
s.mu.Lock() s.mu.Lock()
s.started = false s.started = false
@ -254,6 +256,11 @@ func (s *gnetServer) Start() error {
s.stopped = true s.stopped = true
} }
s.stopMu.Unlock() s.stopMu.Unlock()
go func() {
if shutdownErr := s.shutdownHandlerExecutor(); shutdownErr != nil {
s.logger.Warn("%s handler executor shutdown error: %v", s.protocol.String(), shutdownErr)
}
}()
}() }()
// 等待服务器启动 // 等待服务器启动
@ -277,6 +284,7 @@ func (s *gnetServer) Stop() error {
s.mu.Lock() s.mu.Lock()
if !s.started { if !s.started {
s.mu.Unlock() s.mu.Unlock()
_ = s.shutdownHandlerExecutor()
return errors.ErrServerNotStarted return errors.ErrServerNotStarted
} }
s.mu.Unlock() s.mu.Unlock()
@ -313,9 +321,26 @@ func (s *gnetServer) Stop() error {
s.started = false s.started = false
s.mu.Unlock() s.mu.Unlock()
if err := s.shutdownHandlerExecutor(); err != nil && stopErr == nil {
stopErr = err
}
return stopErr return stopErr
} }
func (s *gnetServer) shutdownHandlerExecutor() error {
if s == nil || s.eventHandler == nil || s.eventHandler.handlerExecutor == nil {
return nil
}
timeout := s.config.ShutdownTimeout
if timeout <= 0 {
timeout = 30 * time.Second
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
return s.eventHandler.handlerExecutor.shutdown(ctx)
}
// formatGnetAddr 格式化gnet地址 // formatGnetAddr 格式化gnet地址
func formatGnetAddr(addr, protocolStr string) string { func formatGnetAddr(addr, protocolStr string) string {
prefix := protocolStr + "://" prefix := protocolStr + "://"

@ -0,0 +1,58 @@
package server
import (
"context"
"sync"
"git.noahlan.cn/noahlan/nnet/v2/internal/logger"
"git.noahlan.cn/noahlan/nnet/v2/pkg/config"
"git.noahlan.cn/noahlan/nnet/v2/pkg/executor"
)
// handlerExecutor owns the optional asynchronous handler executor. An
// injected executor is deliberately not owned by the server.
type handlerExecutor struct {
exec executor.Executor
owned bool
shutdownOnce sync.Once
shutdownErr error
}
func newHandlerExecutor(cfg *config.Config, log logger.Logger) *handlerExecutor {
if cfg == nil || cfg.HandlerExecutionMode != config.HandlerExecutionPerConnection {
return nil
}
if cfg.HandlerExecutor != nil {
return &handlerExecutor{exec: cfg.HandlerExecutor}
}
return &handlerExecutor{
exec: executor.NewPerKey(executor.Config{
Workers: cfg.HandlerWorkers,
QueueSize: cfg.HandlerQueueSize,
OnPanic: func(value interface{}) {
log.Error("Handler panic recovered: %v", value)
},
}),
owned: true,
}
}
func (h *handlerExecutor) submit(key string, task executor.Task) error {
if h == nil || h.exec == nil {
return nil
}
return h.exec.Submit(key, task)
}
func (h *handlerExecutor) shutdown(ctx context.Context) error {
if h == nil || !h.owned || h.exec == nil {
return nil
}
h.shutdownOnce.Do(func() {
h.shutdownErr = h.exec.Shutdown(ctx)
})
return h.shutdownErr
}

@ -99,9 +99,10 @@ func createContext(
// 注意response的header会在parseProtocolHeader之后设置 // 注意response的header会在parseProtocolHeader之后设置
// 因为此时request的header已经解析完成 // 因为此时request的header已经解析完成
// 如果有超时设置创建带超时的context // If a timeout is configured, it starts when the request is admitted. In
// 注意cancel函数需要被调用以避免context泄漏 // asynchronous mode the request is created before queueing, so queue wait
// 但由于handler执行时间可能很长cancel会在context被GC时自动调用 // time is included. Cancellation is cooperative; it cannot forcibly stop a
// handler that ignores ctx.Done().
var ctx context.Context var ctx context.Context
var cancel context.CancelFunc var cancel context.CancelFunc
if handlerTimeout > 0 { if handlerTimeout > 0 {

@ -47,9 +47,10 @@ func newMessageHandler(
// handleMessage 处理单个消息通用函数TCP/UDP/串口共享) // handleMessage 处理单个消息通用函数TCP/UDP/串口共享)
func (mh *messageHandler) handleMessage( func (mh *messageHandler) handleMessage(
ctx ctxpkg.Context, ctx ctxpkg.Context,
message []byte,
protocol protocolpkg.Protocol, protocol protocolpkg.Protocol,
) error { ) error {
message := ctx.Request().Raw()
// 1. 协议解码 // 1. 协议解码
if protocol != nil { if protocol != nil {
if err := parseProtocolHeader(ctx.Request(), message, protocol); err != nil { if err := parseProtocolHeader(ctx.Request(), message, protocol); err != nil {
@ -158,18 +159,67 @@ func (mh *messageHandler) handleMessageWithContext(
conn ctxpkg.Connection, conn ctxpkg.Connection,
message []byte, message []byte,
protocol protocolpkg.Protocol, protocol protocolpkg.Protocol,
codecRegistry codecpkg.Registry,
handlerTimeout time.Duration, handlerTimeout time.Duration,
) { ) {
ctx, cancel := createContext(parentCtx, conn, message, protocol, codecRegistry, mh.codecResolverChain, mh.defaultCodec, mh.cloneHeader, handlerTimeout) ctx, cancel := mh.newContext(parentCtx, conn, message, protocol, handlerTimeout)
if cancel != nil { if cancel != nil {
defer cancel() defer cancel()
} }
// 处理消息 mh.handlePreparedMessage(ctx, protocol)
if err := mh.handleMessage(ctx, message, protocol); err != nil { }
// handleMessageWithExecutor creates the request context before enqueueing so
// HandlerTimeout includes queue wait time as well as handler execution time.
func (mh *messageHandler) handleMessageWithExecutor(
parentCtx context.Context,
connID string,
conn ctxpkg.Connection,
message []byte,
protocol protocolpkg.Protocol,
handlerTimeout time.Duration,
dispatcher *handlerExecutor,
) error {
if dispatcher == nil {
mh.handleMessageWithContext(parentCtx, conn, message, protocol, handlerTimeout)
return nil
}
messageCopy := cloneMessage(message)
ctx, cancel := mh.newContext(parentCtx, conn, messageCopy, protocol, handlerTimeout)
if err := dispatcher.submit(connID, func() {
defer cancel()
mh.handlePreparedMessage(ctx, protocol)
}); err != nil {
cancel()
return err
}
return nil
}
func (mh *messageHandler) newContext(
parentCtx context.Context,
conn ctxpkg.Connection,
message []byte,
protocol protocolpkg.Protocol,
handlerTimeout time.Duration,
) (ctxpkg.Context, context.CancelFunc) {
return createContext(parentCtx, conn, message, protocol, mh.codecRegistry, mh.codecResolverChain, mh.defaultCodec, mh.cloneHeader, handlerTimeout)
}
func (mh *messageHandler) handlePreparedMessage(ctx ctxpkg.Context, protocol protocolpkg.Protocol) {
if err := mh.handleMessage(ctx, protocol); err != nil {
// 根据错误类型决定是否发送错误响应 // 根据错误类型决定是否发送错误响应
errorMsg := fmt.Sprintf("Error: %v\n", err) errorMsg := fmt.Sprintf("Error: %v\n", err)
_ = ctx.Response().WriteBytes([]byte(errorMsg)) _ = ctx.Response().WriteBytes([]byte(errorMsg))
} }
} }
func cloneMessage(message []byte) []byte {
if len(message) == 0 {
return nil
}
copyMessage := make([]byte, len(message))
copy(copyMessage, message)
return copyMessage
}

@ -6,6 +6,7 @@ import (
"io" "io"
"net/http" "net/http"
"sync" "sync"
"time"
"git.noahlan.cn/noahlan/nnet/v2/internal/connection" "git.noahlan.cn/noahlan/nnet/v2/internal/connection"
"git.noahlan.cn/noahlan/nnet/v2/internal/logger" "git.noahlan.cn/noahlan/nnet/v2/internal/logger"
@ -39,6 +40,7 @@ type SerialServer struct {
metrics metricspkg.Metrics metrics metricspkg.Metrics
healthChecker health.Checker healthChecker health.Checker
msgHandler *messageHandler msgHandler *messageHandler
handlerExecutor *handlerExecutor
} }
// NewSerialServer 创建串口服务器 // NewSerialServer 创建串口服务器
@ -99,6 +101,7 @@ func NewSerialServer(cfg *config.Config) (*SerialServer, error) {
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
cloneHeader := configHelper.CloneHeader() cloneHeader := configHelper.CloneHeader()
handlerExec := newHandlerExecutor(cfg, log)
return &SerialServer{ return &SerialServer{
config: cfg, config: cfg,
logger: log, logger: log,
@ -115,6 +118,7 @@ func NewSerialServer(cfg *config.Config) (*SerialServer, error) {
cancel: cancel, cancel: cancel,
started: false, started: false,
msgHandler: newMessageHandler(log, codecRegistry, codecResolverChain, defaultCodec, r, cloneHeader), msgHandler: newMessageHandler(log, codecRegistry, codecResolverChain, defaultCodec, r, cloneHeader),
handlerExecutor: handlerExec,
}, nil }, nil
} }
@ -205,6 +209,9 @@ func (s *SerialServer) Stop() error {
defer s.mu.Unlock() defer s.mu.Unlock()
if !s.started { if !s.started {
if s.handlerExecutor != nil {
_ = s.handlerExecutor.shutdown(context.Background())
}
return errors.ErrServerNotStarted return errors.ErrServerNotStarted
} }
@ -218,6 +225,18 @@ func (s *SerialServer) Stop() error {
} }
} }
shutdownTimeout := s.config.ShutdownTimeout
if shutdownTimeout <= 0 {
shutdownTimeout = 30 * time.Second
}
if s.handlerExecutor != nil {
ctx, cancel := context.WithTimeout(context.Background(), shutdownTimeout)
defer cancel()
if err := s.handlerExecutor.shutdown(ctx); err != nil {
s.logger.Warn("Serial handler executor shutdown error: %v", err)
}
}
s.started = false s.started = false
s.logger.Info("Serial server stopped") s.logger.Info("Serial server stopped")
return nil return nil
@ -311,14 +330,17 @@ func (s *SerialServer) handleConnection() {
handlerTimeout := newServerConfigHelper(s.config).HandlerTimeout() handlerTimeout := newServerConfigHelper(s.config).HandlerTimeout()
for _, message := range messages { for _, message := range messages {
s.metrics.IncRequests() s.metrics.IncRequests()
s.msgHandler.handleMessageWithContext( if err := s.msgHandler.handleMessageWithExecutor(
context.Background(), s.ctx,
connID,
ctxConn, ctxConn,
message, message,
protocol, protocol,
s.codecRegistry,
handlerTimeout, handlerTimeout,
) s.handlerExecutor,
); err != nil {
s.logger.Error("Failed to dispatch serial handler for connection %s: %v", connID, err)
}
// 注意metrics错误计数在handleMessageWithContext内部处理 // 注意metrics错误计数在handleMessageWithContext内部处理
// 如果需要更细粒度的错误处理可以在handleMessageWithContext中回调 // 如果需要更细粒度的错误处理可以在handleMessageWithContext中回调
} }

@ -39,10 +39,12 @@ type unifiedEventHandler struct {
bootCh chan error bootCh chan error
bootOnce *sync.Once bootOnce *sync.Once
msgHandler *messageHandler msgHandler *messageHandler
handlerExecutor *handlerExecutor
parentCtx context.Context
} }
// newUnifiedEventHandler 创建统一的事件处理器 // newUnifiedEventHandler 创建统一的事件处理器
func newUnifiedEventHandler(connManager connection.ManagerInterface, r routerpkg.Router, logger logger.Logger, codecRegistry codecpkg.Registry, codecResolverChain *codecpkg.ResolverChain, defaultCodec string, protocolManager protocolpkg.Manager, protocolName string, enableProtocolEncode bool, cloneHeader bool, handlerTimeout time.Duration, protocol TransportProtocol, bootCh chan error, bootOnce *sync.Once) *unifiedEventHandler { func newUnifiedEventHandler(connManager connection.ManagerInterface, r routerpkg.Router, logger logger.Logger, codecRegistry codecpkg.Registry, codecResolverChain *codecpkg.ResolverChain, defaultCodec string, protocolManager protocolpkg.Manager, protocolName string, enableProtocolEncode bool, cloneHeader bool, handlerTimeout time.Duration, protocol TransportProtocol, bootCh chan error, bootOnce *sync.Once, parentCtx context.Context, handlerExec *handlerExecutor) *unifiedEventHandler {
return &unifiedEventHandler{ return &unifiedEventHandler{
connManager: connManager, connManager: connManager,
router: r, router: r,
@ -61,6 +63,8 @@ func newUnifiedEventHandler(connManager connection.ManagerInterface, r routerpkg
bootCh: bootCh, bootCh: bootCh,
bootOnce: bootOnce, bootOnce: bootOnce,
msgHandler: newMessageHandler(logger, codecRegistry, codecResolverChain, defaultCodec, r, cloneHeader), msgHandler: newMessageHandler(logger, codecRegistry, codecResolverChain, defaultCodec, r, cloneHeader),
handlerExecutor: handlerExec,
parentCtx: parentCtx,
} }
} }
@ -311,14 +315,17 @@ func (h *unifiedEventHandler) handleTraffic(c gnet.Conn, udpPacket []byte) gnet.
// 处理每个完整的消息 // 处理每个完整的消息
ctxConn := toContextConnection(conn) ctxConn := toContextConnection(conn)
for _, message := range messages { for _, message := range messages {
h.msgHandler.handleMessageWithContext( if err := h.msgHandler.handleMessageWithExecutor(
context.Background(), h.parentCtx,
connID,
ctxConn, ctxConn,
message, message,
protocol, protocol,
h.codecRegistry,
h.handlerTimeout, h.handlerTimeout,
) h.handlerExecutor,
); err != nil {
h.logger.Error("Failed to dispatch handler for connection %s: %v", connID, err)
}
} }
// 丢弃已处理的数据(仅面向连接的协议) // 丢弃已处理的数据(仅面向连接的协议)

@ -171,6 +171,9 @@ func (s *WebSocketServer) Stop() error {
s.mu.Lock() s.mu.Lock()
if !s.started { if !s.started {
s.mu.Unlock() s.mu.Unlock()
if s.gnetServer != nil {
_ = s.gnetServer.shutdownHandlerExecutor()
}
return pkgerrors.ErrServerNotStarted return pkgerrors.ErrServerNotStarted
} }
httpServer := s.httpServer httpServer := s.httpServer
@ -219,6 +222,14 @@ waitConnections:
} }
} }
if s.gnetServer.eventHandler != nil && s.gnetServer.eventHandler.handlerExecutor != nil {
executorCtx, executorCancel := context.WithTimeout(context.Background(), shutdownTimeout)
if err := s.gnetServer.eventHandler.handlerExecutor.shutdown(executorCtx); err != nil {
s.gnetServer.logger.Warn("WebSocket handler executor shutdown error: %v", err)
}
executorCancel()
}
for _, hook := range s.serverLifecycleHooks { for _, hook := range s.serverLifecycleHooks {
if err := hook.OnStop(); err != nil { if err := hook.OnStop(); err != nil {
s.gnetServer.logger.Error("OnStop hook error: %v", err) s.gnetServer.logger.Error("OnStop hook error: %v", err)
@ -293,14 +304,17 @@ func (s *WebSocketServer) handleWebSocket(w http.ResponseWriter, r *http.Request
} }
for _, message := range messages { for _, message := range messages {
s.msgHandler.handleMessageWithContext( if err := s.msgHandler.handleMessageWithExecutor(
s.ctx, s.ctx,
connID,
ctxConn, ctxConn,
message, message,
protocol, protocol,
s.codecRegistry,
s.handlerTimeout, s.handlerTimeout,
) s.gnetServer.eventHandler.handlerExecutor,
); err != nil {
s.gnetServer.logger.Error("Failed to dispatch WebSocket handler for connection %s: %v", connID, err)
}
} }
} }
} }

@ -4,6 +4,20 @@ import (
"time" "time"
"git.noahlan.cn/noahlan/nnet/v2/pkg/errors" "git.noahlan.cn/noahlan/nnet/v2/pkg/errors"
"git.noahlan.cn/noahlan/nnet/v2/pkg/executor"
)
// HandlerExecutionMode controls where registered handlers run.
//
// Inline preserves the historical behavior and runs handlers in the transport
// event loop. PerConnection dispatches handlers to an executor and serializes
// tasks for each connection while allowing different connections to run in
// parallel.
type HandlerExecutionMode string
const (
HandlerExecutionInline HandlerExecutionMode = "inline"
HandlerExecutionPerConnection HandlerExecutionMode = "per_connection"
) )
// Config 服务器配置 // Config 服务器配置
@ -81,6 +95,23 @@ type Config struct {
// Handler在事件循环中的最大执行时间如果为0表示不限制。 // Handler在事件循环中的最大执行时间如果为0表示不限制。
HandlerTimeout time.Duration HandlerTimeout time.Duration
// HandlerExecutionMode controls handler dispatch. Empty is equivalent to
// HandlerExecutionInline.
HandlerExecutionMode HandlerExecutionMode
// HandlerWorkers is the number of workers used by the internally-created
// handler executor. Zero uses the executor default.
HandlerWorkers int
// HandlerQueueSize is the maximum number of waiting tasks per connection.
// Zero uses the executor default.
HandlerQueueSize int
// HandlerExecutor optionally supplies an executor for
// HandlerExecutionPerConnection. The server does not shut down an injected
// executor; its owner remains responsible for its lifecycle.
HandlerExecutor executor.Executor
// 串口配置仅当TransportProtocol为serial时使用 // 串口配置仅当TransportProtocol为serial时使用
Serial *SerialConfig Serial *SerialConfig
} }
@ -184,6 +215,7 @@ func DefaultConfig() *Config {
EnableProtocolEncode: false, EnableProtocolEncode: false,
}, },
HandlerTimeout: 30 * time.Second, HandlerTimeout: 30 * time.Second,
HandlerExecutionMode: HandlerExecutionInline,
} }
} }
@ -206,5 +238,13 @@ func (c *Config) Validate() error {
if c.MaxConnections <= 0 { if c.MaxConnections <= 0 {
c.MaxConnections = 10000 c.MaxConnections = 10000
} }
if c.HandlerExecutionMode == "" {
c.HandlerExecutionMode = HandlerExecutionInline
}
switch c.HandlerExecutionMode {
case HandlerExecutionInline, HandlerExecutionPerConnection:
default:
return errors.Newf("unsupported handler execution mode: %s", c.HandlerExecutionMode)
}
return nil return nil
} }

@ -0,0 +1,192 @@
package executor
import (
"context"
"errors"
"runtime"
"sync"
)
var (
ErrClosed = errors.New("executor is closed")
ErrQueueFull = errors.New("executor queue is full")
ErrNilTask = errors.New("executor task is nil")
)
type Task func()
type Executor interface {
Submit(key string, task Task) error
Shutdown(ctx context.Context) error
}
type Config struct {
Workers int
QueueSize int
OnPanic func(value interface{})
}
type PerKey struct {
mu sync.Mutex
cond *sync.Cond
queues map[string]*taskQueue
ready []*taskQueue
workers int
queueSize int
onPanic func(value interface{})
pending int
closed bool
done chan struct{}
wg sync.WaitGroup
}
type taskQueue struct {
key string
tasks []Task
pending int
running bool
ready bool
}
func NewPerKey(cfg Config) *PerKey {
workers := cfg.Workers
if workers <= 0 {
workers = runtime.GOMAXPROCS(0)
}
if workers <= 0 {
workers = 1
}
queueSize := cfg.QueueSize
if queueSize <= 0 {
queueSize = 1024
}
e := &PerKey{
queues: make(map[string]*taskQueue),
workers: workers,
queueSize: queueSize,
onPanic: cfg.OnPanic,
done: make(chan struct{}),
}
e.cond = sync.NewCond(&e.mu)
e.wg.Add(workers)
for i := 0; i < workers; i++ {
go e.worker()
}
go func() {
e.wg.Wait()
close(e.done)
}()
return e
}
func (e *PerKey) Submit(key string, task Task) error {
if task == nil {
return ErrNilTask
}
e.mu.Lock()
defer e.mu.Unlock()
if e.closed {
return ErrClosed
}
queue := e.queues[key]
if queue == nil {
queue = &taskQueue{key: key}
e.queues[key] = queue
}
if len(queue.tasks) >= e.queueSize {
return ErrQueueFull
}
queue.tasks = append(queue.tasks, task)
queue.pending++
e.pending++
if !queue.running && !queue.ready {
queue.ready = true
e.ready = append(e.ready, queue)
e.cond.Signal()
}
return nil
}
func (e *PerKey) Shutdown(ctx context.Context) error {
if ctx == nil {
ctx = context.Background()
}
e.mu.Lock()
if !e.closed {
e.closed = true
e.cond.Broadcast()
}
e.mu.Unlock()
select {
case <-e.done:
return nil
case <-ctx.Done():
return ctx.Err()
}
}
func (e *PerKey) worker() {
defer e.wg.Done()
for {
task, queue, ok := e.nextTask()
if !ok {
return
}
e.run(task)
e.mu.Lock()
queue.running = false
queue.pending--
e.pending--
if len(queue.tasks) > 0 {
queue.ready = true
e.ready = append(e.ready, queue)
e.cond.Signal()
} else {
delete(e.queues, queue.key)
}
e.cond.Broadcast()
e.mu.Unlock()
}
}
func (e *PerKey) nextTask() (Task, *taskQueue, bool) {
e.mu.Lock()
defer e.mu.Unlock()
for len(e.ready) == 0 {
if e.closed && e.pending == 0 {
return nil, nil, false
}
e.cond.Wait()
}
queue := e.ready[0]
e.ready[0] = nil
e.ready = e.ready[1:]
queue.ready = false
task := queue.tasks[0]
queue.tasks[0] = nil
queue.tasks = queue.tasks[1:]
queue.running = true
return task, queue, true
}
func (e *PerKey) run(task Task) {
defer func() {
if value := recover(); value != nil && e.onPanic != nil {
e.onPanic(value)
}
}()
task()
}

@ -0,0 +1,187 @@
package executor
import (
"context"
"errors"
"sync"
"testing"
"time"
)
func TestPerKeyPreservesOrderAndRunsDifferentKeysConcurrently(t *testing.T) {
e := NewPerKey(Config{Workers: 2, QueueSize: 8})
defer func() {
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := e.Shutdown(ctx); err != nil {
t.Fatalf("shutdown error: %v", err)
}
}()
var mu sync.Mutex
var order []int
firstStarted := make(chan struct{})
releaseFirst := make(chan struct{})
secondDone := make(chan struct{})
if err := e.Submit("a", func() {
close(firstStarted)
<-releaseFirst
mu.Lock()
order = append(order, 1)
mu.Unlock()
}); err != nil {
t.Fatalf("submit first task: %v", err)
}
if err := e.Submit("a", func() {
mu.Lock()
order = append(order, 2)
mu.Unlock()
}); err != nil {
t.Fatalf("submit second task: %v", err)
}
if err := e.Submit("b", func() {
close(secondDone)
}); err != nil {
t.Fatalf("submit different-key task: %v", err)
}
select {
case <-firstStarted:
case <-time.After(time.Second):
t.Fatal("first task did not start")
}
select {
case <-secondDone:
case <-time.After(time.Second):
t.Fatal("different-key task did not run concurrently")
}
close(releaseFirst)
deadline := time.After(time.Second)
for {
mu.Lock()
done := len(order) == 2
mu.Unlock()
if done {
break
}
select {
case <-deadline:
t.Fatal("same-key tasks did not drain")
default:
time.Sleep(time.Millisecond)
}
}
mu.Lock()
defer mu.Unlock()
if order[0] != 1 || order[1] != 2 {
t.Fatalf("order = %v, want [1 2]", order)
}
}
func TestPerKeyRejectsFullQueueWithoutBlocking(t *testing.T) {
e := NewPerKey(Config{Workers: 1, QueueSize: 1})
defer func() {
closeTask := make(chan struct{})
_ = e.Submit("a", func() { close(closeTask) })
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = e.Shutdown(ctx)
}()
started := make(chan struct{})
release := make(chan struct{})
if err := e.Submit("a", func() {
close(started)
<-release
}); err != nil {
t.Fatalf("submit running task: %v", err)
}
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("running task did not start")
}
if err := e.Submit("a", func() {}); err != nil {
t.Fatalf("submit queued task: %v", err)
}
done := make(chan error, 1)
go func() {
done <- e.Submit("a", func() {})
}()
select {
case err := <-done:
if !errors.Is(err, ErrQueueFull) {
t.Fatalf("submit error = %v, want ErrQueueFull", err)
}
case <-time.After(100 * time.Millisecond):
t.Fatal("queue-full submit blocked")
}
close(release)
}
func TestPerKeyShutdownDrainsAcceptedTasks(t *testing.T) {
e := NewPerKey(Config{Workers: 1, QueueSize: 4})
var countMu sync.Mutex
count := 0
for i := 0; i < 3; i++ {
if err := e.Submit("a", func() {
countMu.Lock()
count++
countMu.Unlock()
}); err != nil {
t.Fatalf("submit task %d: %v", i, err)
}
}
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := e.Shutdown(ctx); err != nil {
t.Fatalf("shutdown error: %v", err)
}
if err := e.Submit("a", func() {}); !errors.Is(err, ErrClosed) {
t.Fatalf("submit after shutdown error = %v, want ErrClosed", err)
}
countMu.Lock()
defer countMu.Unlock()
if count != 3 {
t.Fatalf("completed tasks = %d, want 3", count)
}
}
func TestPerKeyRecoversTaskPanic(t *testing.T) {
panicSeen := make(chan interface{}, 1)
e := NewPerKey(Config{
Workers: 1,
QueueSize: 2,
OnPanic: func(value interface{}) {
panicSeen <- value
},
})
if err := e.Submit("a", func() { panic("boom") }); err != nil {
t.Fatalf("submit panic task: %v", err)
}
if err := e.Submit("a", func() {}); err != nil {
t.Fatalf("submit follow-up task: %v", err)
}
select {
case value := <-panicSeen:
if value != "boom" {
t.Fatalf("panic value = %v, want boom", value)
}
case <-time.After(time.Second):
t.Fatal("panic callback was not called")
}
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := e.Shutdown(ctx); err != nil {
t.Fatalf("shutdown error: %v", err)
}
}

@ -145,6 +145,16 @@ type SessionConfig = config.SessionConfig
// CodecConfig 编解码器配置类型别名 // CodecConfig 编解码器配置类型别名
type CodecConfig = config.CodecConfig type CodecConfig = config.CodecConfig
// HandlerExecutionMode 配置处理器的执行位置
type HandlerExecutionMode = config.HandlerExecutionMode
const (
// HandlerExecutionInline 在传输事件循环中同步执行处理器
HandlerExecutionInline = config.HandlerExecutionInline
// HandlerExecutionPerConnection 在执行器中执行,并保证同一连接内顺序
HandlerExecutionPerConnection = config.HandlerExecutionPerConnection
)
// NewTCPServer 创建TCP服务器 // NewTCPServer 创建TCP服务器
func NewTCPServer(cfg *Config) (Server, error) { func NewTCPServer(cfg *Config) (Server, error) {
return server.NewServer(cfg) return server.NewServer(cfg)

Loading…
Cancel
Save