package websocket import ( "context" "encoding/json" "errors" "fmt" "net/http" "net/http/httptest" "strconv" "strings" "sync" "time" gorillawebsocket "github.com/gorilla/websocket" melody "github.com/olahol/melody" "kra/internal/integration/runtimeconfig" platformws "kra/pkg/websocket" ) const ProviderMelody = "melody" // Server owns the database-configured WebSocket endpoint. type Server struct { mu sync.RWMutex current *platformws.Server path string messageHandlers []func(*melody.Session, []byte) binaryHandlers []func(*melody.Session, []byte) connectHandlers []func(*melody.Session) disconnectHandlers []func(*melody.Session) stop func() closed bool } func New(store *runtimeconfig.Store) (*Server, func(), error) { s := &Server{} if store != nil { s.apply(storeConfig(store)) s.stop = store.Subscribe("websocket", ProviderMelody, func(config runtimeconfig.Config) { s.apply(config) }) } return s, func() { if s.stop != nil { s.stop() } s.mu.Lock() s.closed = true current := s.current s.current = nil s.path = "" s.mu.Unlock() if current != nil { _ = current.Close() } }, nil } func storeConfig(store *runtimeconfig.Store) runtimeconfig.Config { config, _ := store.Get("websocket", ProviderMelody) return config } // TestConfig performs a local WebSocket handshake using a temporary server // built from the candidate settings. It does not touch the live endpoint. func TestConfig(ctx context.Context, raw json.RawMessage) error { values := map[string]any{} if err := json.Unmarshal(raw, &values); err != nil { return fmt.Errorf("decode websocket configuration: %w", err) } path := text(values, "path") if path == "" { path = "/ws" } if !strings.HasPrefix(path, "/") { path = "/" + path } temporary := platformws.New(platformws.Config{ WriteWait: durationValue(values, "write_wait", 10*time.Second), PongWait: durationValue(values, "pong_wait", 60*time.Second), PingPeriod: durationValue(values, "ping_period", 54*time.Second), MaxMessageSize: int64Value(values, "max_message_size"), MessageBufferSize: int(intValue(values, "message_buffer_size")), ConcurrentMessageHandling: boolValue(values, "concurrent_message_handling"), AllowOrigins: stringList(values, "allow_origins"), }) defer temporary.Close() httpServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != path { http.NotFound(w, r) return } if err := temporary.HandleRequest(w, r); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) } })) defer httpServer.Close() if ctx == nil { ctx = context.Background() } select { case <-ctx.Done(): return ctx.Err() default: } headers := http.Header{} origins := stringList(values, "allow_origins") if len(origins) > 0 { origin := origins[0] if origin == "*" { origin = "http://localhost" } headers.Set("Origin", origin) } wsURL := "ws" + strings.TrimPrefix(httpServer.URL, "http") + path connection, response, err := gorillawebsocket.DefaultDialer.DialContext(ctx, wsURL, headers) if response != nil && response.Body != nil { _ = response.Body.Close() } if err != nil { return fmt.Errorf("websocket handshake failed: %w", err) } if connection == nil { return errors.New("websocket handshake returned an empty connection") } return connection.Close() } func (s *Server) apply(config runtimeconfig.Config) { if s == nil { return } values := map[string]any{} if len(config.Values) > 0 { if err := json.Unmarshal(config.Values, &values); err != nil { return } } s.mu.Lock() if s.closed { s.mu.Unlock() return } previous := s.current if !config.Enabled { s.current = nil s.path = "" s.mu.Unlock() if previous != nil { _ = previous.Close() } return } path := text(values, "path") if path == "" { path = "/ws" } if !strings.HasPrefix(path, "/") { path = "/" + path } next := platformws.New(platformws.Config{ WriteWait: durationValue(values, "write_wait", 10*time.Second), PongWait: durationValue(values, "pong_wait", 60*time.Second), PingPeriod: durationValue(values, "ping_period", 54*time.Second), MaxMessageSize: int64Value(values, "max_message_size"), MessageBufferSize: int(intValue(values, "message_buffer_size")), ConcurrentMessageHandling: boolValue(values, "concurrent_message_handling"), AllowOrigins: stringList(values, "allow_origins"), }) for _, handler := range s.messageHandlers { next.OnMessage(handler) } for _, handler := range s.binaryHandlers { next.OnBinaryMessage(handler) } for _, handler := range s.connectHandlers { next.OnConnect(handler) } for _, handler := range s.disconnectHandlers { next.OnDisconnect(handler) } s.current = next s.path = path s.mu.Unlock() if previous != nil { _ = previous.Close() } } func text(values map[string]any, key string) string { value, ok := values[key] if !ok || value == nil { return "" } return strings.TrimSpace(fmt.Sprint(value)) } func intValue(values map[string]any, key string) int64 { switch value := values[key].(type) { case float64: return int64(value) case int: return int64(value) case json.Number: parsed, _ := strconv.ParseInt(string(value), 10, 64) return parsed default: parsed, _ := strconv.ParseInt(text(values, key), 10, 64) return parsed } } func int64Value(values map[string]any, key string) int64 { return intValue(values, key) } func boolValue(values map[string]any, key string) bool { value, _ := values[key].(bool); return value } func stringList(values map[string]any, key string) []string { value, ok := values[key].([]any) if ok { result := make([]string, 0, len(value)) for _, item := range value { if item != nil && strings.TrimSpace(fmt.Sprint(item)) != "" { result = append(result, strings.TrimSpace(fmt.Sprint(item))) } } return result } if value, ok := values[key].([]string); ok { return append([]string(nil), value...) } return nil } func durationValue(values map[string]any, key string, fallback time.Duration) time.Duration { value := text(values, key) if value == "" { return fallback } parsed, err := time.ParseDuration(value) if err != nil || parsed <= 0 { return fallback } return parsed } func (s *Server) Enabled() bool { if s == nil { return false } s.mu.RLock() defer s.mu.RUnlock() return s.current != nil } func (s *Server) Path() string { if s == nil { return "" } s.mu.RLock() defer s.mu.RUnlock() return s.path } func (s *Server) HandleRequest(w http.ResponseWriter, r *http.Request) error { current, err := s.active() if err != nil { return err } return current.HandleRequest(w, r) } func (s *Server) HandleRequestWithKeys(w http.ResponseWriter, r *http.Request, keys map[string]any) error { current, err := s.active() if err != nil { return err } return current.HandleRequestWithKeys(w, r, keys) } func (s *Server) Broadcast(message []byte) error { current, err := s.active() if err != nil { return err } return current.Broadcast(message) } func (s *Server) BroadcastBinary(message []byte) error { current, err := s.active() if err != nil { return err } return current.BroadcastBinary(message) } func (s *Server) Sessions() ([]*melody.Session, error) { current, err := s.active() if err != nil { return nil, err } return current.Sessions() } func (s *Server) Len() int { current, err := s.active() if err != nil { return 0 } return current.Len() } func (s *Server) Send(session *melody.Session, message []byte) error { current, err := s.active() if err != nil { return err } return current.Send(session, message) } func (s *Server) SendBinary(session *melody.Session, message []byte) error { current, err := s.active() if err != nil { return err } return current.SendBinary(session, message) } func (s *Server) active() (*platformws.Server, error) { if s == nil { return nil, errors.New("websocket server is disabled") } s.mu.RLock() current := s.current s.mu.RUnlock() if current == nil { return nil, errors.New("websocket server is disabled") } return current, nil } func (s *Server) OnMessage(handler func(*melody.Session, []byte)) { if s == nil || handler == nil { return } s.mu.Lock() s.messageHandlers = append(s.messageHandlers, handler) current := s.current if current != nil { current.OnMessage(handler) } s.mu.Unlock() } func (s *Server) OnBinaryMessage(handler func(*melody.Session, []byte)) { if s == nil || handler == nil { return } s.mu.Lock() s.binaryHandlers = append(s.binaryHandlers, handler) current := s.current if current != nil { current.OnBinaryMessage(handler) } s.mu.Unlock() } func (s *Server) OnConnect(handler func(*melody.Session)) { if s == nil || handler == nil { return } s.mu.Lock() s.connectHandlers = append(s.connectHandlers, handler) current := s.current if current != nil { current.OnConnect(handler) } s.mu.Unlock() } func (s *Server) OnDisconnect(handler func(*melody.Session)) { if s == nil || handler == nil { return } s.mu.Lock() s.disconnectHandlers = append(s.disconnectHandlers, handler) current := s.current if current != nil { current.OnDisconnect(handler) } s.mu.Unlock() }