proto-socket/go/communicator.go

577 lines
14 KiB
Go

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)
}