package notify import ( "fmt" "net" "sync" ) const defaultInboundDispatchSource = "_notify.default_inbound_source" // Inbound dispatch runs one serial worker per source so messages from the same // connection keep their relative order. The pending queue is bounded both per // source and in total: a slow handler can never let one connection grow the // queue without limit, and a saturated connection can never starve the others. // Callers block until there is room, which pushes backpressure to the transport // reader instead of growing memory; CloseAndWait unblocks every caller. const ( defaultInboundDispatchQueueLimit = 4096 defaultInboundDispatchQueueBytes = 64 << 20 defaultInboundSourceQueueLimit = 512 defaultInboundSourceQueueBytes = 16 << 20 ) type inboundDispatchItem struct { size int fn func() } type inboundDispatcher struct { mu sync.Mutex closed bool closeCh chan struct{} roomCh chan struct{} queued int queuedBytes int maxItems int maxBytes int sourceItems int sourceBytes int workers map[string]*inboundDispatchWorker wg sync.WaitGroup } type inboundDispatchWorker struct { queue []inboundDispatchItem running bool queued int queuedBytes int } func newInboundDispatcher() *inboundDispatcher { return newInboundDispatcherWithCaps( defaultInboundDispatchQueueLimit, defaultInboundDispatchQueueBytes, defaultInboundSourceQueueLimit, defaultInboundSourceQueueBytes, ) } // newInboundDispatcherWithLimits sizes a dispatcher for a single source: the // per-source caps equal the global caps. func newInboundDispatcherWithLimits(maxItems int, maxBytes int) *inboundDispatcher { return newInboundDispatcherWithCaps(maxItems, maxBytes, maxItems, maxBytes) } func newInboundDispatcherWithCaps(maxItems int, maxBytes int, sourceItems int, sourceBytes int) *inboundDispatcher { if maxItems <= 0 { maxItems = defaultInboundDispatchQueueLimit } if maxBytes <= 0 { maxBytes = defaultInboundDispatchQueueBytes } if sourceItems <= 0 || sourceItems > maxItems { sourceItems = maxItems } if sourceBytes <= 0 || sourceBytes > maxBytes { sourceBytes = maxBytes } return &inboundDispatcher{ closeCh: make(chan struct{}), roomCh: make(chan struct{}, 1), maxItems: maxItems, maxBytes: maxBytes, sourceItems: sourceItems, sourceBytes: sourceBytes, workers: make(map[string]*inboundDispatchWorker), } } // Dispatch queues fn for the given source without byte accounting. Prefer // DispatchSized when the queued payload size is known. Like DispatchSized it // blocks while the queue is at its limit. func (d *inboundDispatcher) Dispatch(source string, fn func()) bool { return d.DispatchSized(source, 0, fn) } // DispatchSized queues fn for the given source. It blocks while the source or // the dispatcher is at its item or byte limit, and returns false once the // dispatcher is closed. A parked caller is released by CloseAndWait, so the // owner of the reader must keep CloseAndWait reachable (a concurrent closer, or // the reader's own stop path once the wait unblocks). func (d *inboundDispatcher) DispatchSized(source string, size int, fn func()) bool { if d == nil || fn == nil { return false } if source == "" { source = defaultInboundDispatchSource } if size < 0 { size = 0 } for { d.mu.Lock() if d.closed { d.mu.Unlock() return false } worker := d.workers[source] if worker == nil { worker = &inboundDispatchWorker{} d.workers[source] = worker } if d.roomLocked(worker, size) { worker.queue = append(worker.queue, inboundDispatchItem{size: size, fn: fn}) worker.queued++ worker.queuedBytes += size d.queued++ d.queuedBytes += size if worker.running { d.mu.Unlock() return true } worker.running = true d.wg.Add(1) d.mu.Unlock() go d.run(source, worker) return true } d.mu.Unlock() // roomCh is signalled whenever a queued item is consumed; closeCh // releases the caller during shutdown. select { case <-d.roomCh: case <-d.closeCh: return false } } } func (d *inboundDispatcher) roomLocked(worker *inboundDispatchWorker, size int) bool { if d.maxItems > 0 && d.queued >= d.maxItems { return false } if d.sourceItems > 0 && worker.queued >= d.sourceItems { return false } // Always admit at least one item per scope so an oversized payload cannot // deadlock the reader behind an empty queue. if d.maxBytes > 0 && d.queued > 0 && d.queuedBytes+size > d.maxBytes { return false } if d.sourceBytes > 0 && worker.queued > 0 && worker.queuedBytes+size > d.sourceBytes { return false } return true } func (d *inboundDispatcher) signalRoomLocked() { select { case d.roomCh <- struct{}{}: default: } } func (d *inboundDispatcher) run(source string, worker *inboundDispatchWorker) { defer d.wg.Done() for { d.mu.Lock() if len(worker.queue) == 0 { worker.running = false if current := d.workers[source]; current == worker { delete(d.workers, source) } d.signalRoomLocked() d.mu.Unlock() return } item := worker.queue[0] worker.queue[0] = inboundDispatchItem{} worker.queue = worker.queue[1:] worker.queued-- worker.queuedBytes -= item.size d.queued-- d.queuedBytes -= item.size d.signalRoomLocked() d.mu.Unlock() item.fn() } } func (d *inboundDispatcher) CloseAndWait() { if d == nil { return } d.mu.Lock() if !d.closed { d.closed = true close(d.closeCh) } d.signalRoomLocked() d.mu.Unlock() d.wg.Wait() } func clientInboundDispatchSource() string { return "client" } func serverInboundDispatchSource(source interface{}) string { switch data := source.(type) { case serverInboundSource: return serverInboundDispatchSourceKey(data) case *serverInboundSource: if data == nil { return defaultInboundDispatchSource } return serverInboundDispatchSourceKey(*data) case net.Conn: return fmt.Sprintf("conn:%p", data) case string: if data == "" { return defaultInboundDispatchSource } return "peer:" + data default: return defaultInboundDispatchSource } } func serverInboundDispatchSourceKey(source serverInboundSource) string { if source.Conn != nil { return fmt.Sprintf("conn:%p:%d", source.Conn, source.TransportGeneration) } if source.Logical != nil { return fmt.Sprintf("logical:%s:%d", source.Logical.ID(), source.TransportGeneration) } if source.Source != "" { return fmt.Sprintf("peer:%s:%d", source.Source, source.TransportGeneration) } if source.RemoteAddr != nil { return fmt.Sprintf("addr:%s:%d", source.RemoteAddr.String(), source.TransportGeneration) } return defaultInboundDispatchSource }