package protosocket import ( "context" "log/slog" "net/http" "sync" "time" 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]clientDiagnostics } type clientDiagnostics struct { connectionID string } type clientSnapshot struct { client *toki.WsClient diagnostic clientDiagnostics } func NewServer(cfg Config, logger *slog.Logger) *Server { return &Server{ cfg: cfg, dispatcher: NewDispatcher(), logger: logger, clients: make(map[*toki.WsClient]clientDiagnostics), } } 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.logError("failed to accept websocket connection", "error", err) return } connectionID := "conn-" + generateID() diagnostic := clientDiagnostics{connectionID: connectionID} client := toki.NewWsClient( conn, s.cfg.HeartbeatIntervalSec, s.cfg.HeartbeatWaitSec, ParserMap(), ) s.mu.Lock() s.clients[client] = diagnostic s.mu.Unlock() client.AddDisconnectListener(func(c *toki.WsClient) { s.mu.Lock() diagnostic := s.clients[c] delete(s.clients, c) s.mu.Unlock() s.logInfo("proto-socket client disconnected", diagnosticsLogAttrs(Envelope{ProtocolVersion: ProtocolVersion}, diagnostic.connectionID)...) }) toki.AddRequestListenerTyped[*structpb.Struct, *structpb.Struct](&client.Communicator, func(req *structpb.Struct) (*structpb.Struct, error) { env, err := EnvelopeFromStruct(req) if err != nil { errEnv := withDiagnosticsMeta(Envelope{ ProtocolVersion: ProtocolVersion, ID: generateID(), CorrelationID: "", Type: "error", Error: &EnvelopeError{ Code: "INVALID_ENVELOPE", Message: "failed to parse envelope: " + err.Error(), }, }, connectionID) s.logError("proto-socket request rejected", append(diagnosticsLogAttrs(errEnv, connectionID), "error", err)...) res, _ := errEnv.ToStruct() return res, nil } resEnv := withDiagnosticsMeta(s.dispatcher.Dispatch(r.Context(), env), connectionID) resStruct, err := resEnv.ToStruct() if err != nil { errEnv := withDiagnosticsMeta(Envelope{ ProtocolVersion: ProtocolVersion, ID: generateID(), CorrelationID: env.ID, Type: "error", Channel: env.Channel, Action: env.Action, Error: &EnvelopeError{ Code: "SERIALIZATION_ERROR", Message: "failed to serialize response envelope: " + err.Error(), }, }, connectionID) s.logError("proto-socket response serialization failed", append(diagnosticsLogAttrs(errEnv, connectionID), "error", err)...) res, _ := errEnv.ToStruct() return res, nil } s.logInfo("proto-socket request handled", diagnosticsLogAttrs(resEnv, connectionID)...) return resStruct, nil }) s.logInfo("proto-socket client connected", diagnosticsLogAttrs(Envelope{ProtocolVersion: ProtocolVersion}, connectionID)...) } 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 { s.mu.Lock() clients := make([]clientSnapshot, 0, len(s.clients)) for cl, diagnostic := range s.clients { clients = append(clients, clientSnapshot{ client: cl, diagnostic: diagnostic, }) } s.mu.Unlock() for _, snapshot := range clients { if snapshot.client.IsAlive() { eventEnv := withDiagnosticsMeta(env, snapshot.diagnostic.connectionID) structMsg, err := eventEnv.ToStruct() if err != nil { s.logError("proto-socket broadcast serialization failed", append(diagnosticsLogAttrs(eventEnv, snapshot.diagnostic.connectionID), "error", err)...) return err } if err := snapshot.client.Send(structMsg); err != nil { s.logError("failed to send proto-socket broadcast", append(diagnosticsLogAttrs(eventEnv, snapshot.diagnostic.connectionID), "error", err)...) continue } s.logInfo("proto-socket broadcast sent", diagnosticsLogAttrs(eventEnv, snapshot.diagnostic.connectionID)...) } } 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) }, } } func withDiagnosticsMeta(env Envelope, connectionID string) Envelope { if env.ProtocolVersion == "" { env.ProtocolVersion = ProtocolVersion } meta := make(map[string]any, len(env.Meta)+6) for k, v := range env.Meta { meta[k] = v } meta["connection_id"] = connectionID meta["protocol_version"] = env.ProtocolVersion meta["channel"] = env.Channel meta["action"] = env.Action meta["error_code"] = errorCode(env) meta["timestamp"] = diagnosticsTimestamp() env.Meta = meta return env } func errorCode(env Envelope) string { if env.Error == nil { return "" } return env.Error.Code } func diagnosticsLogAttrs(env Envelope, connectionID string) []any { protocolVersion := env.ProtocolVersion if protocolVersion == "" { protocolVersion = ProtocolVersion } return []any{ "connection_id", connectionID, "protocol_version", protocolVersion, "channel", env.Channel, "action", env.Action, "error_code", errorCode(env), "timestamp", diagnosticsTimestamp(), } } func diagnosticsTimestamp() string { return time.Now().UTC().Format(time.RFC3339Nano) } func (s *Server) logInfo(msg string, args ...any) { if s.logger == nil { return } s.logger.Info(msg, args...) } func (s *Server) logError(msg string, args ...any) { if s.logger == nil { return } s.logger.Error(msg, args...) }