package toki_socket import ( "errors" "fmt" "reflect" "sync" "sync/atomic" "time" "google.golang.org/protobuf/proto" "toki-labs.com/toki_socket/go/packets" ) var ErrNotConnected = errors.New("not connected") 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 { return c.nonce.Add(1) } func (c *Communicator) shutdown() { c.isAlive.Store(false) c.closeOnce.Do(func() { close(c.closed) }) } 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) }