From 604bd05e3b71446fcfe091df6c620474f7697388 Mon Sep 17 00:00:00 2001 From: NoahLan <6995syu@163.com> Date: Sat, 22 Aug 2026 21:23:22 +0800 Subject: [PATCH] feat: add non-blocking per-connection handler execution --- internal/request/request.go | 4 +- internal/request/request_test.go | 8 + internal/server/gnet_server.go | 31 +++- internal/server/handler_executor.go | 58 +++++++ internal/server/helpers.go | 7 +- internal/server/message_handler.go | 60 ++++++- internal/server/serial_server.go | 30 +++- internal/server/unified_event_handler.go | 17 +- internal/server/websocket_server.go | 20 ++- pkg/config/config.go | 42 ++++- pkg/executor/executor.go | 192 +++++++++++++++++++++++ pkg/executor/executor_test.go | 187 ++++++++++++++++++++++ pkg/nnet/server.go | 10 ++ 13 files changed, 641 insertions(+), 25 deletions(-) create mode 100644 internal/server/handler_executor.go create mode 100644 pkg/executor/executor.go create mode 100644 pkg/executor/executor_test.go diff --git a/internal/request/request.go b/internal/request/request.go index b57dd09..3a1f9aa 100644 --- a/internal/request/request.go +++ b/internal/request/request.go @@ -23,8 +23,10 @@ var _ RequestSetter = (*requestImpl)(nil) // New 创建新的请求对象 func New(raw []byte, protocol protocolpkg.Protocol) requestpkg.Request { + rawCopy := make([]byte, len(raw)) + copy(rawCopy, raw) return &requestImpl{ - raw: raw, + raw: rawCopy, protocol: protocol, } } diff --git a/internal/request/request_test.go b/internal/request/request_test.go index aba9ef8..fb999b0 100644 --- a/internal/request/request_test.go +++ b/internal/request/request_test.go @@ -18,6 +18,14 @@ func TestRequest(t *testing.T) { 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) { req := New([]byte("test"), nil).(*requestImpl) diff --git a/internal/server/gnet_server.go b/internal/server/gnet_server.go index c45b8d8..c13cff9 100644 --- a/internal/server/gnet_server.go +++ b/internal/server/gnet_server.go @@ -153,6 +153,8 @@ func newGnetServer(cfg *config.Config, protocol TransportProtocol) (*gnetServer, } } + ctx, cancel := context.WithCancel(context.Background()) + // 创建事件处理器 handlerTimeout := configHelper.HandlerTimeout() cloneHeader := configHelper.CloneHeader() @@ -160,9 +162,8 @@ func newGnetServer(cfg *config.Config, protocol TransportProtocol) (*gnetServer, protocolName := cfg.ApplicationProtocol bootCh := make(chan error, 1) bootOnce := &sync.Once{} - eventHandler := newUnifiedEventHandler(connManager, r, log, codecRegistry, codecResolverChain, defaultCodec, protocolManager, protocolName, enableProtocolEncode, cloneHeader, handlerTimeout, protocol, bootCh, bootOnce) - - ctx, cancel := context.WithCancel(context.Background()) + handlerExec := newHandlerExecutor(cfg, log) + eventHandler := newUnifiedEventHandler(connManager, r, log, codecRegistry, codecResolverChain, defaultCodec, protocolManager, protocolName, enableProtocolEncode, cloneHeader, handlerTimeout, protocol, bootCh, bootOnce, ctx, handlerExec) return &gnetServer{ config: cfg, @@ -234,6 +235,7 @@ func (s *gnetServer) Start() error { // 在goroutine中启动服务器 go func() { err := gnet.Run(s.eventHandler, addr, options...) + s.cancel() if err != nil { s.mu.Lock() s.started = false @@ -254,6 +256,11 @@ func (s *gnetServer) Start() error { s.stopped = true } 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() if !s.started { s.mu.Unlock() + _ = s.shutdownHandlerExecutor() return errors.ErrServerNotStarted } s.mu.Unlock() @@ -313,9 +321,26 @@ func (s *gnetServer) Stop() error { s.started = false s.mu.Unlock() + if err := s.shutdownHandlerExecutor(); err != nil && stopErr == nil { + stopErr = err + } + 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地址 func formatGnetAddr(addr, protocolStr string) string { prefix := protocolStr + "://" diff --git a/internal/server/handler_executor.go b/internal/server/handler_executor.go new file mode 100644 index 0000000..9c2d2b4 --- /dev/null +++ b/internal/server/handler_executor.go @@ -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 +} diff --git a/internal/server/helpers.go b/internal/server/helpers.go index 88caba2..b589d4d 100644 --- a/internal/server/helpers.go +++ b/internal/server/helpers.go @@ -99,9 +99,10 @@ func createContext( // 注意:response的header会在parseProtocolHeader之后设置 // 因为此时request的header已经解析完成 - // 如果有超时设置,创建带超时的context - // 注意:cancel函数需要被调用以避免context泄漏 - // 但由于handler执行时间可能很长,cancel会在context被GC时自动调用 + // If a timeout is configured, it starts when the request is admitted. In + // asynchronous mode the request is created before queueing, so queue wait + // time is included. Cancellation is cooperative; it cannot forcibly stop a + // handler that ignores ctx.Done(). var ctx context.Context var cancel context.CancelFunc if handlerTimeout > 0 { diff --git a/internal/server/message_handler.go b/internal/server/message_handler.go index 23d324f..2e5e0ff 100644 --- a/internal/server/message_handler.go +++ b/internal/server/message_handler.go @@ -47,9 +47,10 @@ func newMessageHandler( // handleMessage 处理单个消息(通用函数,TCP/UDP/串口共享) func (mh *messageHandler) handleMessage( ctx ctxpkg.Context, - message []byte, protocol protocolpkg.Protocol, ) error { + message := ctx.Request().Raw() + // 1. 协议解码 if protocol != nil { if err := parseProtocolHeader(ctx.Request(), message, protocol); err != nil { @@ -158,18 +159,67 @@ func (mh *messageHandler) handleMessageWithContext( conn ctxpkg.Connection, message []byte, protocol protocolpkg.Protocol, - codecRegistry codecpkg.Registry, 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 { defer cancel() } - // 处理消息 - if err := mh.handleMessage(ctx, message, protocol); err != nil { + mh.handlePreparedMessage(ctx, protocol) +} + +// 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) _ = 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 +} diff --git a/internal/server/serial_server.go b/internal/server/serial_server.go index bdcc820..fc5cf81 100644 --- a/internal/server/serial_server.go +++ b/internal/server/serial_server.go @@ -6,6 +6,7 @@ import ( "io" "net/http" "sync" + "time" "git.noahlan.cn/noahlan/nnet/v2/internal/connection" "git.noahlan.cn/noahlan/nnet/v2/internal/logger" @@ -39,6 +40,7 @@ type SerialServer struct { metrics metricspkg.Metrics healthChecker health.Checker msgHandler *messageHandler + handlerExecutor *handlerExecutor } // NewSerialServer 创建串口服务器 @@ -99,6 +101,7 @@ func NewSerialServer(cfg *config.Config) (*SerialServer, error) { ctx, cancel := context.WithCancel(context.Background()) cloneHeader := configHelper.CloneHeader() + handlerExec := newHandlerExecutor(cfg, log) return &SerialServer{ config: cfg, logger: log, @@ -115,6 +118,7 @@ func NewSerialServer(cfg *config.Config) (*SerialServer, error) { cancel: cancel, started: false, msgHandler: newMessageHandler(log, codecRegistry, codecResolverChain, defaultCodec, r, cloneHeader), + handlerExecutor: handlerExec, }, nil } @@ -205,6 +209,9 @@ func (s *SerialServer) Stop() error { defer s.mu.Unlock() if !s.started { + if s.handlerExecutor != nil { + _ = s.handlerExecutor.shutdown(context.Background()) + } 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.logger.Info("Serial server stopped") return nil @@ -311,14 +330,17 @@ func (s *SerialServer) handleConnection() { handlerTimeout := newServerConfigHelper(s.config).HandlerTimeout() for _, message := range messages { s.metrics.IncRequests() - s.msgHandler.handleMessageWithContext( - context.Background(), + if err := s.msgHandler.handleMessageWithExecutor( + s.ctx, + connID, ctxConn, message, protocol, - s.codecRegistry, handlerTimeout, - ) + s.handlerExecutor, + ); err != nil { + s.logger.Error("Failed to dispatch serial handler for connection %s: %v", connID, err) + } // 注意:metrics错误计数在handleMessageWithContext内部处理 // 如果需要更细粒度的错误处理,可以在handleMessageWithContext中回调 } diff --git a/internal/server/unified_event_handler.go b/internal/server/unified_event_handler.go index 4a40dfa..e795e9e 100644 --- a/internal/server/unified_event_handler.go +++ b/internal/server/unified_event_handler.go @@ -39,10 +39,12 @@ type unifiedEventHandler struct { bootCh chan error bootOnce *sync.Once msgHandler *messageHandler + handlerExecutor *handlerExecutor + parentCtx context.Context } // 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{ connManager: connManager, router: r, @@ -61,6 +63,8 @@ func newUnifiedEventHandler(connManager connection.ManagerInterface, r routerpkg bootCh: bootCh, bootOnce: bootOnce, 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) for _, message := range messages { - h.msgHandler.handleMessageWithContext( - context.Background(), + if err := h.msgHandler.handleMessageWithExecutor( + h.parentCtx, + connID, ctxConn, message, protocol, - h.codecRegistry, h.handlerTimeout, - ) + h.handlerExecutor, + ); err != nil { + h.logger.Error("Failed to dispatch handler for connection %s: %v", connID, err) + } } // 丢弃已处理的数据(仅面向连接的协议) diff --git a/internal/server/websocket_server.go b/internal/server/websocket_server.go index 24f6255..fa91fcb 100644 --- a/internal/server/websocket_server.go +++ b/internal/server/websocket_server.go @@ -171,6 +171,9 @@ func (s *WebSocketServer) Stop() error { s.mu.Lock() if !s.started { s.mu.Unlock() + if s.gnetServer != nil { + _ = s.gnetServer.shutdownHandlerExecutor() + } return pkgerrors.ErrServerNotStarted } 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 { if err := hook.OnStop(); err != nil { 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 { - s.msgHandler.handleMessageWithContext( + if err := s.msgHandler.handleMessageWithExecutor( s.ctx, + connID, ctxConn, message, protocol, - s.codecRegistry, s.handlerTimeout, - ) + s.gnetServer.eventHandler.handlerExecutor, + ); err != nil { + s.gnetServer.logger.Error("Failed to dispatch WebSocket handler for connection %s: %v", connID, err) + } } } } diff --git a/pkg/config/config.go b/pkg/config/config.go index 07c5c6f..a790200 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -4,6 +4,20 @@ import ( "time" "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 服务器配置 @@ -81,6 +95,23 @@ type Config struct { // Handler在事件循环中的最大执行时间,如果为0表示不限制。 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时使用) Serial *SerialConfig } @@ -183,7 +214,8 @@ func DefaultConfig() *Config { DefaultCodec: "binary", 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 { 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 } diff --git a/pkg/executor/executor.go b/pkg/executor/executor.go new file mode 100644 index 0000000..74e1b12 --- /dev/null +++ b/pkg/executor/executor.go @@ -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() +} diff --git a/pkg/executor/executor_test.go b/pkg/executor/executor_test.go new file mode 100644 index 0000000..bc005e1 --- /dev/null +++ b/pkg/executor/executor_test.go @@ -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) + } +} diff --git a/pkg/nnet/server.go b/pkg/nnet/server.go index 18fcca5..35966e8 100644 --- a/pkg/nnet/server.go +++ b/pkg/nnet/server.go @@ -145,6 +145,16 @@ type SessionConfig = config.SessionConfig // CodecConfig 编解码器配置类型别名 type CodecConfig = config.CodecConfig +// HandlerExecutionMode 配置处理器的执行位置 +type HandlerExecutionMode = config.HandlerExecutionMode + +const ( + // HandlerExecutionInline 在传输事件循环中同步执行处理器 + HandlerExecutionInline = config.HandlerExecutionInline + // HandlerExecutionPerConnection 在执行器中执行,并保证同一连接内顺序 + HandlerExecutionPerConnection = config.HandlerExecutionPerConnection +) + // NewTCPServer 创建TCP服务器 func NewTCPServer(cfg *Config) (Server, error) { return server.NewServer(cfg)