package protosocket import ( "context" "log/slog" "net/http" "sync" toki "git.toki-labs.com/toki/proto-socket/go" "google.golang.org/protobuf/proto" "google.golang.org/protobuf/types/known/structpb" "nhooyr.io/websocket" ) type Config struct { HeartbeatIntervalSec int HeartbeatWaitSec int } type Server struct { cfg Config dispatcher *Dispatcher logger *slog.Logger mu sync.Mutex clients map[*toki.WsClient]struct{} } func NewServer(cfg Config, logger *slog.Logger) *Server { return &Server{ cfg: cfg, dispatcher: NewDispatcher(), logger: logger, clients: make(map[*toki.WsClient]struct{}), } } func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { opts := &websocket.AcceptOptions{ InsecureSkipVerify: true, } conn, err := websocket.Accept(w, r, opts) if err != nil { s.logger.Error("failed to accept websocket connection", "error", err) return } client := toki.NewWsClient( conn, s.cfg.HeartbeatIntervalSec, s.cfg.HeartbeatWaitSec, ParserMap(), ) s.mu.Lock() s.clients[client] = struct{}{} s.mu.Unlock() client.AddDisconnectListener(func(c *toki.WsClient) { s.mu.Lock() delete(s.clients, c) s.mu.Unlock() s.logger.Info("websocket client disconnected") }) toki.AddRequestListenerTyped[*structpb.Struct, *structpb.Struct](&client.Communicator, func(req *structpb.Struct) (*structpb.Struct, error) { env, err := EnvelopeFromStruct(req) if err != nil { errEnv := Envelope{ ProtocolVersion: ProtocolVersion, ID: generateID(), CorrelationID: "", Type: "error", Error: &EnvelopeError{ Code: "INVALID_ENVELOPE", Message: "failed to parse envelope: " + err.Error(), }, } res, _ := errEnv.ToStruct() return res, nil } resEnv := s.dispatcher.Dispatch(r.Context(), env) resStruct, err := resEnv.ToStruct() if err != nil { errEnv := Envelope{ ProtocolVersion: ProtocolVersion, ID: generateID(), CorrelationID: env.ID, Type: "error", Error: &EnvelopeError{ Code: "SERIALIZATION_ERROR", Message: "failed to serialize response envelope: " + err.Error(), }, } res, _ := errEnv.ToStruct() return res, nil } return resStruct, nil }) s.logger.Info("websocket client connected successfully") } func (s *Server) Close() error { s.mu.Lock() clients := make([]*toki.WsClient, 0, len(s.clients)) for cl := range s.clients { clients = append(clients, cl) } s.mu.Unlock() for _, cl := range clients { _ = cl.Close() } return nil } func (s *Server) BroadcastEnvelope(ctx context.Context, env Envelope) error { structMsg, err := env.ToStruct() if err != nil { return err } s.mu.Lock() clients := make([]*toki.WsClient, 0, len(s.clients)) for cl := range s.clients { clients = append(clients, cl) } s.mu.Unlock() for _, cl := range clients { if cl.IsAlive() { if err := cl.Send(structMsg); err != nil { s.logger.Error("failed to send broadcast to client", "error", err) } } } return nil } func (s *Server) Dispatcher() *Dispatcher { return s.dispatcher } func ParserMap() toki.ParserMap { return toki.ParserMap{ toki.TypeNameOf(&structpb.Struct{}): func(b []byte) (proto.Message, error) { msg := &structpb.Struct{} return msg, proto.Unmarshal(b, msg) }, } }