kra-new/internal/integration/mq/emqx.go

711 lines
20 KiB
Go

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
ProviderRabbitMQ = platformmq.ProviderRabbitMQ
retryTick = time.Second
)
// 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 {
r.apply(ProviderEMQX, storeConfig(store, ProviderEMQX))
r.apply(ProviderRabbitMQ, storeConfig(store, ProviderRabbitMQ))
r.stop = append(r.stop,
store.Subscribe("mq", ProviderEMQX, func(config runtimeconfig.Config) { r.apply(ProviderEMQX, config) }),
store.Subscribe("mq", ProviderRabbitMQ, func(config runtimeconfig.Config) { r.apply(ProviderRabbitMQ, 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.
// For RabbitMQ this also checks the configured exchange and queue topology.
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"),
})
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 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 []string{ProviderEMQX, ProviderRabbitMQ} {
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 != 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()
}