package mq import ( "context" "encoding/json" "fmt" "log/slog" "strconv" "strings" "sync" "time" "kra/internal/integration/runtimeconfig" platformmq "kra/pkg/mq" ) const ( ProviderEMQX = platformmq.ProviderEMQX ProviderKafka = platformmq.ProviderKafka ProviderRabbitMQ = platformmq.ProviderRabbitMQ retryTick = time.Second ) var providers = []string{ProviderEMQX, ProviderKafka, ProviderRabbitMQ} // Reloadable owns process-wide message clients and the logical subscription // declarations used to restore them after reconnects or configuration reloads. type Reloadable struct { mu sync.RWMutex opMu sync.Mutex clients map[string]platformmq.Client configs map[string]runtimeconfig.Config subscriptions map[string]map[string]map[string]subscription bindings map[string]map[string]byte pending map[string]bool nextRetry map[string]time.Time legacySeq uint64 stop []func() retryStop chan struct{} retryDone chan struct{} logger *slog.Logger closed bool } type subscription struct { qos byte handler platformmq.Handler } type namedClient struct { owner *Reloadable provider string } func New(store *runtimeconfig.Store, logger *slog.Logger) (*Reloadable, func(), error) { if logger == nil { logger = slog.Default() } r := &Reloadable{ clients: make(map[string]platformmq.Client), configs: make(map[string]runtimeconfig.Config), subscriptions: make(map[string]map[string]map[string]subscription), bindings: make(map[string]map[string]byte), pending: make(map[string]bool), nextRetry: make(map[string]time.Time), retryStop: make(chan struct{}), retryDone: make(chan struct{}), logger: logger, } if store != nil { for _, provider := range providers { r.apply(provider, storeConfig(store, provider)) provider := provider r.stop = append(r.stop, store.Subscribe("mq", provider, func(config runtimeconfig.Config) { r.apply(provider, config) })) } } go r.retryLoop() cleanup := func() { for _, stop := range r.stop { stop() } _ = r.Close() } return r, cleanup, nil } func storeConfig(store *runtimeconfig.Store, provider string) runtimeconfig.Config { config, _ := store.Get("mq", provider) return config } // TestConfig creates a short-lived provider client and closes it immediately. // RabbitMQ also checks topology; Kafka reads cluster metadata. func TestConfig(ctx context.Context, provider string, raw json.RawMessage) error { if ctx != nil { select { case <-ctx.Done(): return ctx.Err() default: } } provider = strings.ToLower(strings.TrimSpace(provider)) if provider == ProviderEMQX { values := map[string]any{} if err := json.Unmarshal(raw, &values); err != nil { return fmt.Errorf("decode %s configuration: %w", provider, err) } baseID := configText(values, "client_id") values["client_id"] = fmt.Sprintf("%s-test-%d", baseID, time.Now().UnixNano()) encoded, err := json.Marshal(values) if err != nil { return fmt.Errorf("encode %s test configuration: %w", provider, err) } raw = encoded } client, err := newProviderClient(provider, raw) if err != nil { return err } if client == nil || !client.Connected() { if client != nil { _ = client.Close() } return platformmq.ErrUnavailable } closeErr := client.Close() if ctx != nil { select { case <-ctx.Done(): return ctx.Err() default: } } return closeErr } func (r *Reloadable) apply(provider string, config runtimeconfig.Config) { provider = strings.ToLower(strings.TrimSpace(provider)) r.opMu.Lock() defer r.opMu.Unlock() if r.closed { return } r.ensureStateLocked() config.Kind = "mq" config.Provider = provider config.Values = append(json.RawMessage(nil), config.Values...) r.configs[provider] = config r.replaceClientLocked(provider, nil) if !config.Enabled { delete(r.pending, provider) delete(r.nextRetry, provider) return } if err := r.activateLocked(provider, config); err != nil { r.scheduleRetryLocked(provider, config.Values) if r.logger != nil { r.logger.Warn("message integration unavailable", "mod", "mq", "provider", provider, "error", err) } } } func (r *Reloadable) activateLocked(provider string, config runtimeconfig.Config) error { client, err := newProviderClient(provider, config.Values) if err != nil { return err } bindings, err := r.restoreSubscriptionsLocked(provider, client) if err != nil { _ = client.Close() return fmt.Errorf("restore message subscriptions: %w", err) } r.replaceClientLocked(provider, client) r.mu.Lock() if r.bindings == nil { r.bindings = make(map[string]map[string]byte) } r.bindings[provider] = bindings r.mu.Unlock() delete(r.pending, provider) delete(r.nextRetry, provider) return nil } func newProviderClient(provider string, raw json.RawMessage) (platformmq.Client, error) { values := map[string]any{} if err := json.Unmarshal(raw, &values); err != nil { return nil, fmt.Errorf("decode %s configuration: %w", provider, err) } switch provider { case ProviderEMQX: return platformmq.NewMQTT(platformmq.Config{ Enabled: true, Broker: configText(values, "broker"), ClientID: configText(values, "client_id"), Username: configText(values, "username"), Password: configText(values, "password"), KeepAlive: configSeconds(values, "keep_alive"), CleanSession: configBool(values, "clean_session"), ConnectTimeout: configSeconds(values, "connect_timeout"), ReconnectInterval: configSeconds(values, "reconnect_interval"), }) case ProviderRabbitMQ: return platformmq.NewRabbitMQ(platformmq.RabbitMQConfig{ Enabled: true, Host: configText(values, "host"), Port: configInt(values, "port"), Username: configText(values, "username"), Password: configText(values, "password"), VHost: configText(values, "vhost"), Exchange: configText(values, "exchange"), ExchangeType: configText(values, "exchange_type"), Queue: configText(values, "queue"), RoutingKey: configText(values, "routing_key"), Durable: configBool(values, "durable"), AutoDelete: configBool(values, "auto_delete"), PrefetchCount: configInt(values, "prefetch_count"), Heartbeat: configSeconds(values, "heartbeat"), ConnectTimeout: configSeconds(values, "connect_timeout"), ReconnectInterval: configSeconds(values, "reconnect_interval"), TLS: configBool(values, "tls"), }) case ProviderKafka: return platformmq.NewKafka(platformmq.KafkaConfig{ Enabled: true, Brokers: configStrings(values, "brokers"), ClientID: configText(values, "client_id"), GroupID: configText(values, "group_id"), Username: configText(values, "username"), Password: configText(values, "password"), TLS: configBool(values, "tls"), TLSSkipVerify: configBool(values, "tls_skip_verify"), StartOffset: configText(values, "start_offset"), MinBytes: configInt(values, "min_bytes"), MaxBytes: configInt(values, "max_bytes"), MaxWait: configSeconds(values, "max_wait"), ConnectTimeout: configSeconds(values, "connect_timeout"), ReconnectInterval: configSeconds(values, "reconnect_interval"), AllowAutoTopicCreation: configBool(values, "allow_auto_topic_creation"), }) default: return nil, fmt.Errorf("unsupported message provider %q", provider) } } func configText(values map[string]any, key string) string { value, ok := values[key] if !ok || value == nil { return "" } return strings.TrimSpace(fmt.Sprint(value)) } func configInt(values map[string]any, key string) int { switch value := values[key].(type) { case float64: return int(value) case int: return value case json.Number: parsed, _ := strconv.Atoi(string(value)) return parsed default: parsed, _ := strconv.Atoi(configText(values, key)) return parsed } } func configStrings(values map[string]any, key string) []string { items, ok := values[key].([]any) if ok { result := make([]string, 0, len(items)) for _, item := range items { if value := strings.TrimSpace(fmt.Sprint(item)); value != "" { result = append(result, value) } } return result } stringsValue, ok := values[key].([]string) if ok { return append([]string(nil), stringsValue...) } return nil } func configSeconds(values map[string]any, key string) time.Duration { seconds := configInt(values, key) if seconds <= 0 { return 0 } return time.Duration(seconds) * time.Second } func configBool(values map[string]any, key string) bool { value, _ := values[key].(bool) return value } func (r *Reloadable) retryLoop() { defer close(r.retryDone) ticker := time.NewTicker(retryTick) defer ticker.Stop() for { select { case <-r.retryStop: return case <-ticker.C: r.retryOnce() } } } func (r *Reloadable) retryOnce() { r.opMu.Lock() defer r.opMu.Unlock() if r.closed { return } r.ensureStateLocked() now := time.Now() for _, provider := range providers { config, exists := r.configs[provider] if !exists || !config.Enabled { continue } client := r.clientLocked(provider) if client != nil && client.Connected() { if r.pending[provider] { if err := r.reconcileProviderLocked(provider); err != nil { r.scheduleRetryLocked(provider, config.Values) } } continue } if client != nil { if recovering, ok := client.(interface{ Reconnecting() bool }); ok && recovering.Reconnecting() { continue } } if retryAt := r.nextRetry[provider]; !retryAt.IsZero() && now.Before(retryAt) { continue } r.replaceClientLocked(provider, nil) if err := r.activateLocked(provider, config); err != nil { r.scheduleRetryLocked(provider, config.Values) if r.logger != nil { r.logger.Warn("message integration reconnect failed", "mod", "mq", "provider", provider, "error", err) } } } } func (r *Reloadable) scheduleRetryLocked(provider string, raw json.RawMessage) { if r.pending == nil { r.pending = make(map[string]bool) } if r.nextRetry == nil { r.nextRetry = make(map[string]time.Time) } r.pending[provider] = true r.nextRetry[provider] = time.Now().Add(configRetryInterval(raw)) } func configRetryInterval(raw json.RawMessage) time.Duration { values := map[string]any{} _ = json.Unmarshal(raw, &values) interval := configSeconds(values, "reconnect_interval") if interval <= 0 { return 5 * time.Second } return interval } func (r *Reloadable) replaceClientLocked(provider string, next platformmq.Client) { r.mu.Lock() if r.clients == nil { r.clients = make(map[string]platformmq.Client) } if r.bindings == nil { r.bindings = make(map[string]map[string]byte) } old := r.clients[provider] if next == nil { delete(r.clients, provider) delete(r.bindings, provider) } else { r.clients[provider] = next } r.mu.Unlock() if old != nil && old != next { _ = old.Close() } } func (r *Reloadable) restoreSubscriptionsLocked(provider string, client platformmq.Client) (map[string]byte, error) { desired := r.desiredSubscriptions(provider) bindings := make(map[string]byte, len(desired)) for topic, qos := range desired { if err := client.Subscribe(context.Background(), topic, qos, r.dispatcher(provider, topic)); err != nil { return nil, err } bindings[topic] = qos } return bindings, nil } func (r *Reloadable) desiredSubscriptions(provider string) map[string]byte { result := make(map[string]byte) r.mu.RLock() defer r.mu.RUnlock() for topic, owners := range r.subscriptions[provider] { for _, item := range owners { if qos, exists := result[topic]; !exists || item.qos > qos { result[topic] = item.qos } } } return result } func (r *Reloadable) dispatcher(provider, topic string) platformmq.Handler { return func(ctx context.Context, message platformmq.Message) { r.mu.RLock() owners := r.subscriptions[provider][topic] handlers := make([]platformmq.Handler, 0, len(owners)) for _, item := range owners { handlers = append(handlers, item.handler) } r.mu.RUnlock() for _, handler := range handlers { handler(ctx, message) } } } func (r *Reloadable) clientLocked(provider string) platformmq.Client { r.mu.RLock() defer r.mu.RUnlock() return r.clients[provider] } func (r *Reloadable) reconcileProviderLocked(provider string) (err error) { client := r.clientLocked(provider) desired := r.desiredSubscriptions(provider) if client == nil || !client.Connected() { if len(desired) == 0 { delete(r.pending, provider) delete(r.nextRetry, provider) return nil } return platformmq.ErrUnavailable } r.mu.RLock() current := make(map[string]byte, len(r.bindings[provider])) for topic, qos := range r.bindings[provider] { current[topic] = qos } r.mu.RUnlock() // Broker operations can partially succeed. Persist every successful // change even when a later operation fails, otherwise the next retry will // repeat stale unbinds and may never reach the remaining subscriptions. defer func() { r.setBindings(provider, current) }() for topic := range current { if _, exists := desired[topic]; exists { continue } if err := client.Unsubscribe(context.Background(), topic); err != nil { return err } delete(current, topic) } for topic, qos := range desired { if oldQoS, exists := current[topic]; exists && oldQoS == qos { continue } if _, exists := current[topic]; exists { if err := client.Unsubscribe(context.Background(), topic); err != nil { return err } delete(current, topic) } if err := client.Subscribe(context.Background(), topic, qos, r.dispatcher(provider, topic)); err != nil { return err } current[topic] = qos } delete(r.pending, provider) delete(r.nextRetry, provider) return nil } func (r *Reloadable) setBindings(provider string, bindings map[string]byte) { copyOfBindings := make(map[string]byte, len(bindings)) for topic, qos := range bindings { copyOfBindings[topic] = qos } r.mu.Lock() if r.bindings == nil { r.bindings = make(map[string]map[string]byte) } r.bindings[provider] = copyOfBindings r.mu.Unlock() } func (r *Reloadable) Register(set platformmq.SubscriptionSet) error { normalized, err := platformmq.NormalizeSubscriptionSet(set) if err != nil { return err } r.opMu.Lock() defer r.opMu.Unlock() if r.closed { return platformmq.ErrUnavailable } r.ensureStateLocked() r.mu.Lock() if r.subscriptions[normalized.Provider] == nil { r.subscriptions[normalized.Provider] = make(map[string]map[string]subscription) } for topic, owners := range r.subscriptions[normalized.Provider] { delete(owners, normalized.Owner) if len(owners) == 0 { delete(r.subscriptions[normalized.Provider], topic) } } for _, item := range normalized.Topics { if r.subscriptions[normalized.Provider][item.Topic] == nil { r.subscriptions[normalized.Provider][item.Topic] = make(map[string]subscription) } r.subscriptions[normalized.Provider][item.Topic][normalized.Owner] = subscription{qos: item.QoS, handler: item.Handler} } r.mu.Unlock() if err := r.reconcileProviderLocked(normalized.Provider); err != nil { config := r.configs[normalized.Provider] r.scheduleRetryLocked(normalized.Provider, config.Values) if r.logger != nil { r.logger.Warn("message subscription bind deferred", "mod", "mq", "provider", normalized.Provider, "owner", normalized.Owner, "error", err) } } return nil } func (r *Reloadable) Unregister(owner string) error { owner = strings.TrimSpace(owner) if owner == "" { return fmt.Errorf("mq subscription owner is empty") } r.opMu.Lock() defer r.opMu.Unlock() if r.closed { return nil } r.ensureStateLocked() r.mu.Lock() providers := make([]string, 0, len(r.subscriptions)) for provider, topics := range r.subscriptions { providers = append(providers, provider) for topic, owners := range topics { delete(owners, owner) if len(owners) == 0 { delete(topics, topic) } } } r.mu.Unlock() for _, provider := range providers { if err := r.reconcileProviderLocked(provider); err != nil { if config, ok := r.configs[provider]; ok { r.scheduleRetryLocked(provider, config.Values) } } } return nil } func (r *Reloadable) Publish(ctx context.Context, topic string, payload []byte, qos byte, retain bool) error { return r.PublishTo(ctx, ProviderEMQX, topic, payload, qos, retain) } func (r *Reloadable) PublishTo(ctx context.Context, provider, topic string, payload []byte, qos byte, retain bool) error { if ctx == nil { ctx = context.Background() } provider = strings.ToLower(strings.TrimSpace(provider)) client := r.clientLocked(provider) if client == nil { return platformmq.ErrUnavailable } return client.Publish(ctx, topic, payload, qos, retain) } func (r *Reloadable) Subscribe(ctx context.Context, topic string, qos byte, handler platformmq.Handler) error { return r.SubscribeTo(ctx, ProviderEMQX, topic, qos, handler) } func (r *Reloadable) SubscribeTo(ctx context.Context, provider, topic string, qos byte, handler platformmq.Handler) error { if ctx == nil { ctx = context.Background() } if ctx != nil { select { case <-ctx.Done(): return ctx.Err() default: } } provider = strings.ToLower(strings.TrimSpace(provider)) r.opMu.Lock() r.legacySeq++ owner := fmt.Sprintf("legacy/%s/%d", provider, r.legacySeq) r.opMu.Unlock() if err := r.Register(platformmq.SubscriptionSet{Owner: owner, Provider: provider, Topics: []platformmq.TopicSubscription{{Topic: topic, QoS: qos, Handler: handler}}}); err != nil { return err } if !r.ConnectedTo(provider) { // The legacy API reports unavailable when no live broker exists. Do not // leave a durable declaration behind when that call failed. _ = r.Unregister(owner) return platformmq.ErrUnavailable } return nil } func (r *Reloadable) Unsubscribe(ctx context.Context, topics ...string) error { return r.UnsubscribeFrom(ctx, ProviderEMQX, topics...) } func (r *Reloadable) UnsubscribeFrom(ctx context.Context, provider string, topics ...string) error { if ctx == nil { ctx = context.Background() } if ctx != nil { select { case <-ctx.Done(): return ctx.Err() default: } } provider = strings.ToLower(strings.TrimSpace(provider)) if len(topics) == 0 { return fmt.Errorf("mq topics are empty") } r.opMu.Lock() defer r.opMu.Unlock() r.ensureStateLocked() r.mu.Lock() for _, rawTopic := range topics { topic := strings.TrimSpace(rawTopic) owners := r.subscriptions[provider][topic] for owner := range owners { if strings.HasPrefix(owner, "legacy/") { delete(owners, owner) } } if len(owners) == 0 { delete(r.subscriptions[provider], topic) } } r.mu.Unlock() if err := r.reconcileProviderLocked(provider); err != nil { if config, ok := r.configs[provider]; ok { r.scheduleRetryLocked(provider, config.Values) } return err } return nil } func (r *Reloadable) Client(provider string) platformmq.Client { provider = strings.ToLower(strings.TrimSpace(provider)) if provider != ProviderEMQX && provider != ProviderKafka && provider != ProviderRabbitMQ { return nil } return &namedClient{owner: r, provider: provider} } func (c *namedClient) Publish(ctx context.Context, topic string, payload []byte, qos byte, retain bool) error { return c.owner.PublishTo(ctx, c.provider, topic, payload, qos, retain) } func (c *namedClient) Subscribe(ctx context.Context, topic string, qos byte, handler platformmq.Handler) error { return c.owner.SubscribeTo(ctx, c.provider, topic, qos, handler) } func (c *namedClient) Unsubscribe(ctx context.Context, topics ...string) error { return c.owner.UnsubscribeFrom(ctx, c.provider, topics...) } func (c *namedClient) Connected() bool { return c.owner.ConnectedTo(c.provider) } func (*namedClient) Close() error { return nil } func (r *Reloadable) Connected() bool { return r.ConnectedTo(ProviderEMQX) } func (r *Reloadable) ConnectedTo(provider string) bool { provider = strings.ToLower(strings.TrimSpace(provider)) client := r.clientLocked(provider) return client != nil && client.Connected() } func (r *Reloadable) Close() error { r.opMu.Lock() if r.closed { r.opMu.Unlock() return nil } r.closed = true if r.retryStop != nil { close(r.retryStop) } r.mu.Lock() clients := make([]platformmq.Client, 0, len(r.clients)) for provider, client := range r.clients { clients = append(clients, client) delete(r.clients, provider) } r.mu.Unlock() r.opMu.Unlock() if r.retryDone != nil { <-r.retryDone } for _, client := range clients { if client != nil { _ = client.Close() } } return nil } func (r *Reloadable) ensureStateLocked() { if r.configs == nil { r.configs = make(map[string]runtimeconfig.Config) } if r.pending == nil { r.pending = make(map[string]bool) } if r.nextRetry == nil { r.nextRetry = make(map[string]time.Time) } r.mu.Lock() if r.clients == nil { r.clients = make(map[string]platformmq.Client) } if r.subscriptions == nil { r.subscriptions = make(map[string]map[string]map[string]subscription) } if r.bindings == nil { r.bindings = make(map[string]map[string]byte) } r.mu.Unlock() }