- Add code review log for cloud G07 - Update plan and code review documents for core diagnostics - Improve protosocket server implementation and tests
235 lines
6.1 KiB
Go
235 lines
6.1 KiB
Go
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...)
|
|
}
|