379 lines
8.5 KiB
Go
379 lines
8.5 KiB
Go
package toki_socket
|
|
|
|
import (
|
|
"errors"
|
|
"fmt"
|
|
"reflect"
|
|
"sync"
|
|
"sync/atomic"
|
|
"time"
|
|
|
|
"google.golang.org/protobuf/proto"
|
|
|
|
"git.toki-labs.com/toki/common-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 Communicator struct {
|
|
mu sync.RWMutex
|
|
nonce atomic.Int32
|
|
isAlive atomic.Bool
|
|
parserMap ParserMap
|
|
handlers map[string][]func(proto.Message)
|
|
reqHandlers map[string]func(proto.Message, int32)
|
|
pendingRequests map[int32]*pendingRequest
|
|
writeQueue chan queuedPacket
|
|
closed chan struct{}
|
|
closeOnce sync.Once
|
|
transport Transport
|
|
writeErrHandler 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.closed = make(chan struct{})
|
|
c.isAlive.Store(true)
|
|
go c.writeLoop()
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
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) Close() error {
|
|
c.shutdown()
|
|
if c.transport == nil {
|
|
return nil
|
|
}
|
|
return c.transport.Close()
|
|
}
|
|
|
|
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:
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
func (c *Communicator) QueuePacket(base *packets.PacketBase) error {
|
|
if !c.IsAlive() {
|
|
return ErrNotConnected
|
|
}
|
|
done := make(chan error, 1)
|
|
select {
|
|
case c.writeQueue <- queuedPacket{base: base, done: done}:
|
|
case <-c.closed:
|
|
return ErrNotConnected
|
|
}
|
|
select {
|
|
case err := <-done:
|
|
return err
|
|
case <-c.closed:
|
|
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.mu.RLock()
|
|
reqHandler := c.reqHandlers[typeName]
|
|
listeners := append([]func(proto.Message){}, c.handlers[typeName]...)
|
|
c.mu.RUnlock()
|
|
|
|
if reqHandler != nil {
|
|
msg, err := c.parse(typeName, data)
|
|
if err != nil {
|
|
return
|
|
}
|
|
go reqHandler(msg, incomingNonce)
|
|
return
|
|
}
|
|
|
|
if len(listeners) == 0 {
|
|
return
|
|
}
|
|
msg, err := c.parse(typeName, 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 || !c.IsAlive() {
|
|
return
|
|
}
|
|
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)
|
|
}
|