package proto_socket import ( "errors" "fmt" "reflect" "sync" "sync/atomic" "time" "google.golang.org/protobuf/proto" "git.toki-labs.com/toki/proto-socket/go/packets" ) var ErrNotConnected = errors.New("not connected") const MaxNonce int32 = 1<<31 - 1 type ParserMap map[string]func([]byte) (proto.Message, error) type Transport interface { WritePacket(base *packets.PacketBase) error Close() error } type pendingRequest struct { expectedTypeName string ch chan proto.Message errCh chan error } type queuedPacket struct { base *packets.PacketBase done chan error } type inboundItem struct { typeName string data []byte incomingNonce int32 responseNonce int32 } type Communicator struct { mu sync.RWMutex nonce atomic.Int32 frameSeq atomic.Int64 isAlive atomic.Bool drainOnClose atomic.Bool gateway InboundGateway parserMap ParserMap handlers map[string][]func(proto.Message) reqHandlers map[string]func(proto.Message, int32) pendingRequests map[int32]*pendingRequest writeQueue chan queuedPacket inboundQueue chan inboundItem closed chan struct{} receiveDone chan struct{} closeOnce sync.Once transport Transport writeErrHandler func(error) frameErrHandler func(error) } func TypeNameOf(m proto.Message) string { return string(proto.MessageName(m)) } func NewCommunicator(transport Transport, parserMap ParserMap) *Communicator { c := &Communicator{} c.Initialize(transport, parserMap) return c } func (c *Communicator) Initialize(transport Transport, parserMap ParserMap) { c.transport = transport c.parserMap = make(ParserMap, len(parserMap)+1) for k, v := range parserMap { c.parserMap[k] = v } c.parserMap[TypeNameOf(&packets.HeartBeat{})] = func(b []byte) (proto.Message, error) { m := &packets.HeartBeat{} return m, proto.Unmarshal(b, m) } c.handlers = make(map[string][]func(proto.Message)) c.reqHandlers = make(map[string]func(proto.Message, int32)) c.pendingRequests = make(map[int32]*pendingRequest) c.writeQueue = make(chan queuedPacket, 64) c.inboundQueue = make(chan inboundItem, 64) c.closed = make(chan struct{}) c.receiveDone = make(chan struct{}) c.isAlive.Store(true) go c.writeLoop() go c.receiveLoop() } func (c *Communicator) IsAlive() bool { return c.isAlive.Load() } func (c *Communicator) SetWriteErrorHandler(fn func(error)) { c.mu.Lock() defer c.mu.Unlock() c.writeErrHandler = fn } // SetFrameErrorHandler registers a handler invoked when an inbound gateway fails // to decode a raw frame. Transports use it to apply their parse-error disconnect // semantics for the gateway path, mirroring SetWriteErrorHandler for the write // path. Without a gateway, OnReceivedFrame surfaces decode errors via its return // value instead and this handler is unused. func (c *Communicator) SetFrameErrorHandler(fn func(error)) { c.mu.Lock() defer c.mu.Unlock() c.frameErrHandler = fn } func (c *Communicator) onFrameError(err error) { c.mu.RLock() handler := c.frameErrHandler c.mu.RUnlock() if handler != nil { handler(err) } } func (c *Communicator) nextNonce() int32 { for { current := c.nonce.Load() var next int32 if current >= MaxNonce { next = 1 } else { next = current + 1 } if c.nonce.CompareAndSwap(current, next) { return next } } } func (c *Communicator) shutdown() { c.isAlive.Store(false) c.closeOnce.Do(func() { close(c.closed) }) } func (c *Communicator) ForceShutdown() { c.drainOnClose.Store(false) c.shutdown() c.closeGateway() if c.receiveDone != nil { <-c.receiveDone } c.cancelPendingRequests() } func (c *Communicator) Close() error { c.drainOnClose.Store(true) c.shutdown() c.closeGateway() if c.receiveDone != nil { <-c.receiveDone } c.cancelPendingRequests() return nil } func (c *Communicator) cancelPendingRequests() { c.mu.Lock() defer c.mu.Unlock() for k, pending := range c.pendingRequests { pending.errCh <- ErrNotConnected delete(c.pendingRequests, k) } } func (c *Communicator) writeLoop() { for { select { case item := <-c.writeQueue: err := c.transport.WritePacket(item.base) item.done <- err if err != nil { c.mu.RLock() handler := c.writeErrHandler c.mu.RUnlock() if handler != nil { handler(err) } } case <-c.closed: for { select { case item := <-c.writeQueue: err := c.transport.WritePacket(item.base) item.done <- err case <-c.receiveDone: for { select { case item := <-c.writeQueue: err := c.transport.WritePacket(item.base) item.done <- err default: return } } } } } } } func (c *Communicator) QueuePacket(base *packets.PacketBase) error { select { case <-c.receiveDone: return ErrNotConnected default: } done := make(chan error, 1) select { case c.writeQueue <- queuedPacket{base: base, done: done}: case <-c.receiveDone: return ErrNotConnected } select { case err := <-done: return err case <-c.receiveDone: return ErrNotConnected } } func (c *Communicator) Send(m proto.Message) error { if !c.IsAlive() { return ErrNotConnected } data, err := proto.Marshal(m) if err != nil { return err } return c.QueuePacket(&packets.PacketBase{ TypeName: TypeNameOf(m), Nonce: c.nextNonce(), Data: data, }) } func (c *Communicator) SendRequest(req proto.Message, resType proto.Message, timeout time.Duration) (proto.Message, error) { if !c.IsAlive() { return nil, ErrNotConnected } if timeout <= 0 { timeout = 30 * time.Second } requestNonce := c.nextNonce() pending := &pendingRequest{ expectedTypeName: TypeNameOf(resType), ch: make(chan proto.Message, 1), errCh: make(chan error, 1), } c.mu.Lock() c.pendingRequests[requestNonce] = pending c.mu.Unlock() data, err := proto.Marshal(req) if err != nil { c.removePending(requestNonce) return nil, err } err = c.QueuePacket(&packets.PacketBase{ TypeName: TypeNameOf(req), Nonce: requestNonce, Data: data, }) if err != nil { c.removePending(requestNonce) return nil, err } timer := time.NewTimer(timeout) defer timer.Stop() select { case res := <-pending.ch: return res, nil case err := <-pending.errCh: return nil, err case <-timer.C: c.removePending(requestNonce) return nil, fmt.Errorf("request timeout for nonce %d", requestNonce) case <-c.closed: c.removePending(requestNonce) return nil, ErrNotConnected } } func (c *Communicator) AddListener(typeName string, fn func(proto.Message)) { c.mu.Lock() defer c.mu.Unlock() if _, ok := c.reqHandlers[typeName]; ok { panic(fmt.Sprintf("type %s is already registered with AddRequestListener", typeName)) } c.handlers[typeName] = append(c.handlers[typeName], fn) } func (c *Communicator) RemoveListeners(typeName string) { c.mu.Lock() defer c.mu.Unlock() delete(c.handlers, typeName) } func (c *Communicator) AddRequestListener(typeName string, fn func(proto.Message, int32)) { c.mu.Lock() defer c.mu.Unlock() if len(c.handlers[typeName]) > 0 { panic(fmt.Sprintf("type %s is already registered with AddListener", typeName)) } if _, ok := c.reqHandlers[typeName]; ok { panic(fmt.Sprintf("type %s is already registered with AddRequestListener", typeName)) } c.reqHandlers[typeName] = fn } func (c *Communicator) OnReceivedData(typeName string, data []byte, incomingNonce, responseNonce int32) { if responseNonce > 0 { c.handleResponse(typeName, data, responseNonce) return } c.EnqueueInbound(typeName, data, incomingNonce, responseNonce) } // EnableInboundGateway attaches a goroutine worker-pool gateway in front of the // receive coordinator. Raw frames submitted via OnReceivedFrame are decoded by // the pool and reordered by an internal seq before reaching OnReceivedData, so // the coordinator keeps FIFO dispatch and sole ownership of stateful handling. // // The gateway performs pure PacketBase envelope decode only; it never touches // the pending-request map, listeners, or the write queue. Opt-in: without it // OnReceivedFrame decodes inline on the calling goroutine. func (c *Communicator) EnableInboundGateway(workers, queueSize int) { gateway := NewWorkerGateway(workers, queueSize, func(f DecodedFrame) { c.OnReceivedData(f.TypeName, f.Data, f.IncomingNonce, f.ResponseNonce) }, func(err error) { // Surface gateway decode errors to the transport's frame error // handler so the parse-error disconnect semantics match the inline // path. The gateway never tears down the transport itself. c.onFrameError(err) }, ) c.AttachInboundGateway(gateway) } // AttachInboundGateway installs a custom InboundGateway in front of the receive // coordinator. The gateway's sink must forward decoded frames to OnReceivedData // to preserve coordinator ownership. Use EnableInboundGateway for the default // worker-pool gateway. func (c *Communicator) AttachInboundGateway(gateway InboundGateway) { c.mu.Lock() c.gateway = gateway c.mu.Unlock() } // OnReceivedFrame ingests a raw PacketBase frame from a transport read loop. // // When a gateway is attached the frame is submitted with an internal seq for // off-coordinator decode + reorder and nil is returned; a frame that fails to // decode is reported asynchronously, in seq order, to the frame error handler // registered via SetFrameErrorHandler so the transport can apply the same // parse-error disconnect semantics as the inline path. The gateway never tears // down the transport itself. // // Without a gateway the frame is decoded inline on the calling goroutine and // forwarded to OnReceivedData; a decode error is returned so the transport can // apply its existing parse-error disconnect semantics. func (c *Communicator) OnReceivedFrame(raw []byte) error { if !c.IsAlive() { return nil } c.mu.RLock() gateway := c.gateway c.mu.RUnlock() if gateway != nil { gateway.Submit(InboundFrame{Seq: c.frameSeq.Add(1), Bytes: raw}) return nil } base, err := decodePacketBase(raw) if err != nil { return err } c.OnReceivedData(base.GetTypeName(), base.GetData(), base.GetNonce(), base.GetResponseNonce()) return nil } func (c *Communicator) closeGateway() { c.mu.Lock() gateway := c.gateway c.gateway = nil c.mu.Unlock() if gateway != nil { gateway.Close() } } func (c *Communicator) EnqueueInbound(typeName string, data []byte, incomingNonce, responseNonce int32) { if !c.IsAlive() { return } item := inboundItem{ typeName: typeName, data: data, incomingNonce: incomingNonce, responseNonce: responseNonce, } select { case c.inboundQueue <- item: case <-c.closed: } } func (c *Communicator) receiveLoop() { defer close(c.receiveDone) for { select { case item := <-c.inboundQueue: c.dispatchInbound(item) case <-c.closed: if c.drainOnClose.Load() { // drain remaining in queue for { select { case item := <-c.inboundQueue: c.dispatchInbound(item) default: return } } } else { return } } } } func (c *Communicator) dispatchInbound(item inboundItem) { c.mu.RLock() reqHandler := c.reqHandlers[item.typeName] listeners := append([]func(proto.Message){}, c.handlers[item.typeName]...) c.mu.RUnlock() if reqHandler != nil { msg, err := c.parse(item.typeName, item.data) if err != nil { return } reqHandler(msg, item.incomingNonce) return } if len(listeners) == 0 { return } msg, err := c.parse(item.typeName, item.data) if err != nil { return } for _, listener := range listeners { listener(msg) } } func (c *Communicator) handleResponse(typeName string, data []byte, responseNonce int32) { pending := c.removePending(responseNonce) if pending == nil { return } if typeName != pending.expectedTypeName { pending.errCh <- fmt.Errorf("response type mismatch for nonce %d: expected %s, got %s", responseNonce, pending.expectedTypeName, typeName) return } msg, err := c.parse(typeName, data) if err != nil { pending.errCh <- err return } pending.ch <- msg } func (c *Communicator) parse(typeName string, data []byte) (proto.Message, error) { c.mu.RLock() parser := c.parserMap[typeName] c.mu.RUnlock() if parser == nil { return nil, fmt.Errorf("protobuf parser is not registered for type %s", typeName) } return parser(data) } func (c *Communicator) removePending(nonce int32) *pendingRequest { c.mu.Lock() defer c.mu.Unlock() pending := c.pendingRequests[nonce] delete(c.pendingRequests, nonce) return pending } func AddListenerTyped[T proto.Message](c *Communicator, fn func(T)) { typeName := TypeNameOf(newMessageOf[T]()) c.AddListener(typeName, func(m proto.Message) { typed, ok := m.(T) if !ok { panic(fmt.Sprintf("received %T for listener %s", m, typeName)) } fn(typed) }) } func AddRequestListenerTyped[Req proto.Message, Res proto.Message](c *Communicator, fn func(Req) (Res, error)) { reqTypeName := TypeNameOf(newMessageOf[Req]()) c.AddRequestListener(reqTypeName, func(m proto.Message, requestNonce int32) { req, ok := m.(Req) if !ok { return } res, err := fn(req) if err != nil { return } select { case <-c.receiveDone: return default: } data, err := proto.Marshal(res) if err != nil { return } _ = c.QueuePacket(&packets.PacketBase{ TypeName: TypeNameOf(res), Nonce: c.nextNonce(), ResponseNonce: requestNonce, Data: data, }) }) } func SendRequestTyped[Req proto.Message, Res proto.Message](c *Communicator, req Req, timeout time.Duration) (Res, error) { resType := newMessageOf[Res]() msg, err := c.SendRequest(req, resType, timeout) if err != nil { var zero Res return zero, err } res, ok := msg.(Res) if !ok { var zero Res return zero, fmt.Errorf("received %T, expected %T", msg, resType) } return res, nil } func newMessageOf[T proto.Message]() T { var zero T t := reflect.TypeOf(zero) if t == nil || t.Kind() != reflect.Ptr { panic("protobuf type parameter must be a pointer message type") } return reflect.New(t.Elem()).Interface().(T) }