diff --git a/internal/apps/message_gateway/runner/inbound.go b/internal/apps/message_gateway/runner/inbound.go new file mode 100644 index 00000000..f341bedd --- /dev/null +++ b/internal/apps/message_gateway/runner/inbound.go @@ -0,0 +1,67 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package runner + +import ( + "context" + "errors" + "fmt" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/pkg/message_gateway" + "gorm.io/gorm" +) + +const pairingTTL = 15 * time.Minute + +type inboundDeps struct { + LookupBinding func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) + UpsertCode func(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) + GenerateCode func() (string, error) + Emit func(context.Context, message_gateway.InboundMessage) error + Send func(context.Context, message_gateway.Recipient, message_gateway.OutboundMessage) error +} + +// Handle pairs unbound senders or emits bound inbound messages. +func (d inboundDeps) Handle(ctx context.Context, msg message_gateway.InboundMessage) error { + binding, err := d.LookupBinding(ctx, msg.ChannelID, msg.PlatformUserID) + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + to := message_gateway.Recipient{ChatID: msg.ChatID, PlatformUserID: msg.PlatformUserID} + if binding == nil || errors.Is(err, gorm.ErrRecordNotFound) { + gen := d.GenerateCode + if gen == nil { + gen = message_gateway.GenerateCode + } + code, err := gen() + if err != nil { + return err + } + row, err := d.UpsertCode(ctx, msg.ChannelID, msg.PlatformUserID, code, time.Now().Add(pairingTTL)) + if err != nil { + return err + } + display := message_gateway.FormatCode(row.Code) + return d.Send(ctx, to, message_gateway.OutboundMessage{ + Text: fmt.Sprintf("Your pairing code is %s. Open Settings → Profile → Bind a bot and enter this code. It expires in 15 minutes.", display), + ReplyToID: msg.MessageID, + }) + } + + uid := binding.UserID + msg.BindingUserID = &uid + if err := d.Emit(ctx, msg); err != nil { + _ = d.Send(ctx, to, message_gateway.OutboundMessage{ + Text: "could not save your message", + ReplyToID: msg.MessageID, + }) + return err + } + return d.Send(ctx, to, message_gateway.OutboundMessage{ + Text: "received", + ReplyToID: msg.MessageID, + }) +} diff --git a/internal/apps/message_gateway/runner/inbound_test.go b/internal/apps/message_gateway/runner/inbound_test.go new file mode 100644 index 00000000..68f00a1e --- /dev/null +++ b/internal/apps/message_gateway/runner/inbound_test.go @@ -0,0 +1,138 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package runner + +import ( + "context" + "errors" + "strings" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/pkg/message_gateway" + "gorm.io/gorm" +) + +func TestHandle_UnboundMintsCodeAndDoesNotEmit(t *testing.T) { + var sent []message_gateway.OutboundMessage + var emitted int + var upserted string + d := inboundDeps{ + LookupBinding: func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { + return nil, gorm.ErrRecordNotFound + }, + GenerateCode: func() (string, error) { return "ABCD2345", nil }, + UpsertCode: func(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) { + upserted = code + return &model.MessagePairingCode{Code: code, ChannelID: channelID, PlatformUserID: platformUserID, ExpiresAt: expiresAt}, nil + }, + Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { + emitted++ + return nil + }, + Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { + sent = append(sent, msg) + return nil + }, + } + err := d.Handle(context.Background(), message_gateway.InboundMessage{ + ChannelID: 1, PlatformUserID: "u1", ChatID: "u1", Text: "hi", + }) + if err != nil { + t.Fatalf("Handle() error = %v", err) + } + if emitted != 0 { + t.Fatalf("Handle() emitted = %d, want 0", emitted) + } + if upserted != "ABCD2345" { + t.Fatalf("UpsertCode() code = %q, want %q", upserted, "ABCD2345") + } + if len(sent) != 1 { + t.Fatalf("Send() calls = %d, want 1", len(sent)) + } + if !strings.Contains(sent[0].Text, "ABCD-2345") { + t.Fatalf("Handle() send text = %q, want pairing code ABCD-2345", sent[0].Text) + } +} + +func TestHandle_UnboundReusesExistingCode(t *testing.T) { + var sent string + d := inboundDeps{ + LookupBinding: func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { + return nil, gorm.ErrRecordNotFound + }, + GenerateCode: func() (string, error) { return "NEWCODE1", nil }, + UpsertCode: func(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) { + return &model.MessagePairingCode{Code: "OLDCODE2"}, nil + }, + Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { + t.Fatal("must not emit") + return nil + }, + Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { + sent = msg.Text + return nil + }, + } + if err := d.Handle(context.Background(), message_gateway.InboundMessage{ChannelID: 1, PlatformUserID: "u1"}); err != nil { + t.Fatalf("Handle() error = %v", err) + } + if !strings.Contains(sent, "OLDC-ODE2") { + t.Fatalf("Handle() send text = %q, want reused code OLDC-ODE2", sent) + } +} + +func TestHandle_BoundEmitsAndAcks(t *testing.T) { + var got message_gateway.InboundMessage + var sent string + d := inboundDeps{ + LookupBinding: func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { + return &model.MessageBinding{UserID: 9, ChannelID: 1, PlatformUserID: "u1"}, nil + }, + Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { + got = msg + return nil + }, + Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { + sent = msg.Text + return nil + }, + } + err := d.Handle(context.Background(), message_gateway.InboundMessage{ + ChannelID: 1, PlatformUserID: "u1", Text: "hello", + }) + if err != nil { + t.Fatalf("Handle() error = %v", err) + } + if got.Text != "hello" || got.BindingUserID == nil || *got.BindingUserID != 9 { + t.Fatalf("Handle() emit = %+v, want text=hello user=9", got) + } + if sent != "received" { + t.Fatalf("Handle() ack = %q, want %q", sent, "received") + } +} + +func TestHandle_BoundEmitError(t *testing.T) { + var sent string + d := inboundDeps{ + LookupBinding: func(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { + return &model.MessageBinding{UserID: 9, ChannelID: 1, PlatformUserID: "u1"}, nil + }, + Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { + return errors.New("listener failed") + }, + Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { + sent = msg.Text + return nil + }, + } + err := d.Handle(context.Background(), message_gateway.InboundMessage{ChannelID: 1, PlatformUserID: "u1", Text: "x"}) + if err == nil { + t.Fatal("Handle() error = nil, want listener error") + } + if sent != "could not save your message" { + t.Fatalf("Handle() send = %q, want could not save your message", sent) + } +} diff --git a/internal/apps/message_gateway/runner/runner.go b/internal/apps/message_gateway/runner/runner.go new file mode 100644 index 00000000..fb419400 --- /dev/null +++ b/internal/apps/message_gateway/runner/runner.go @@ -0,0 +1,240 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package runner + +import ( + "context" + "fmt" + "os" + "sync" + "time" + + appgw "github.com/Rain-kl/Wavelet/internal/apps/message_gateway" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/listener" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/pkg/logger" + "github.com/Rain-kl/Wavelet/pkg/message_gateway" + "github.com/Rain-kl/Wavelet/pkg/message_gateway/channel/qq" + "github.com/Rain-kl/Wavelet/pkg/message_gateway/channel/telegram" +) + +const ( + reloadInterval = 5 * time.Second + lockTTL = 30 * time.Second + lockRefresh = 10 * time.Second +) + +var registerFactoriesOnce sync.Once + +func registerFactories() { + registerFactoriesOnce.Do(func() { + message_gateway.Register(message_gateway.ChannelTypeTelegram, telegram.New) + message_gateway.Register(message_gateway.ChannelTypeQQ, qq.New) + }) +} + +type runningChannel struct { + ch message_gateway.Channel + updatedAt time.Time + cancel context.CancelFunc +} + +// Start loads enabled channels, connects adapters, and reloads on change. +func Start(ctx context.Context) error { + registerFactories() + r := &gateway{ + node: nodeID(), + running: map[uint64]*runningChannel{}, + } + if err := r.sync(ctx); err != nil { + logger.ErrorF(ctx, "message-gateway initial sync: %v", err) + } + ticker := time.NewTicker(reloadInterval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + r.stopAll(ctx) + return ctx.Err() + case <-ticker.C: + if err := repository.DeleteExpiredPairingCodes(ctx); err != nil { + logger.WarnF(ctx, "message-gateway expire pairing: %v", err) + } + if err := r.sync(ctx); err != nil { + logger.ErrorF(ctx, "message-gateway sync: %v", err) + } + } + } +} + +type gateway struct { + node string + mu sync.Mutex + running map[uint64]*runningChannel + seen map[uint64]time.Time +} + +func (r *gateway) sync(ctx context.Context) error { + rows, err := repository.ListMessageChannels(ctx) + if err != nil { + return err + } + live := make(map[uint64]model.MessageChannel, len(rows)) + for _, row := range rows { + live[row.ID] = row + } + + r.mu.Lock() + defer r.mu.Unlock() + + for id, run := range r.running { + row, ok := live[id] + if !ok || !row.Enabled || !row.UpdatedAt.Equal(run.updatedAt) { + r.stopLocked(ctx, id) + } + } + + for id, row := range live { + if !row.Enabled { + continue + } + if _, ok := r.running[id]; ok { + continue + } + if err := r.startLocked(ctx, row); err != nil { + logger.ErrorF(ctx, "message-gateway start channel %d: %v", id, err) + } + } + return nil +} + +func (r *gateway) startLocked(ctx context.Context, row model.MessageChannel) error { + if !r.acquireLock(ctx, row.ID) { + logger.InfoF(ctx, "message-gateway skip channel %d: lock held", row.ID) + return nil + } + + creds, err := appgw.DecryptCredentials(row.Credentials) + if err != nil { + logger.ErrorF(ctx, "message-gateway decrypt channel %d: %v", row.ID, err) + return nil + } + factory, ok := message_gateway.Lookup(row.Type) + if !ok { + return fmt.Errorf("unknown channel type %q", row.Type) + } + + runCtx, cancel := context.WithCancel(ctx) + var live message_gateway.Channel + deps := inboundDeps{ + LookupBinding: repository.GetBindingByChannelPlatform, + UpsertCode: repository.UpsertPairingCode, + GenerateCode: message_gateway.GenerateCode, + Emit: func(ctx context.Context, msg message_gateway.InboundMessage) error { + listener.EmitMessageGatewayInbound(ctx, msg) + return nil + }, + Send: func(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error { + if live == nil { + return fmt.Errorf("qq/telegram channel not ready") + } + return live.Send(ctx, to, msg) + }, + } + ch, err := factory(message_gateway.ChannelConfig{ + ID: row.ID, + Type: row.Type, + Name: row.Name, + Credentials: creds, + Extra: appgw.ParseExtra(row.Extra), + }, func(ctx context.Context, msg message_gateway.InboundMessage) error { + defer func() { + if rec := recover(); rec != nil { + logger.ErrorF(ctx, "message-gateway inbound panic channel %d: %v", row.ID, rec) + } + }() + return deps.Handle(ctx, msg) + }) + if err != nil { + cancel() + return err + } + live = ch + if err := ch.Connect(runCtx); err != nil { + cancel() + _ = ch.Disconnect(ctx) + return err + } + r.running[row.ID] = &runningChannel{ch: ch, updatedAt: row.UpdatedAt, cancel: cancel} + go r.keepLock(runCtx, row.ID) + logger.InfoF(ctx, "message-gateway connected channel %d type=%s", row.ID, row.Type) + return nil +} + +func (r *gateway) stopLocked(ctx context.Context, id uint64) { + run, ok := r.running[id] + if !ok { + return + } + run.cancel() + if err := run.ch.Disconnect(ctx); err != nil { + logger.WarnF(ctx, "message-gateway disconnect channel %d: %v", id, err) + } + delete(r.running, id) + r.releaseLock(ctx, id) +} + +func (r *gateway) stopAll(ctx context.Context) { + r.mu.Lock() + defer r.mu.Unlock() + for id := range r.running { + r.stopLocked(ctx, id) + } +} + +func (r *gateway) acquireLock(ctx context.Context, id uint64) bool { + if db.Redis == nil { + return true + } + ok, err := db.Redis.SetNX(ctx, lockKey(id), r.node, lockTTL).Result() + if err != nil { + logger.WarnF(ctx, "message-gateway lock channel %d: %v", id, err) + return true + } + return ok +} + +func (r *gateway) keepLock(ctx context.Context, id uint64) { + if db.Redis == nil { + return + } + ticker := time.NewTicker(lockRefresh) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + _ = db.Redis.Expire(ctx, lockKey(id), lockTTL).Err() + } + } +} + +func (r *gateway) releaseLock(ctx context.Context, id uint64) { + if db.Redis == nil { + return + } + _ = db.Redis.Del(ctx, lockKey(id)).Err() +} + +func lockKey(id uint64) string { + return db.PrefixedKey(fmt.Sprintf("wg:channel:%d", id)) +} + +func nodeID() string { + host, _ := os.Hostname() + return fmt.Sprintf("%s:%d", host, os.Getpid()) +} diff --git a/internal/apps/message_gateway/secret.go b/internal/apps/message_gateway/secret.go new file mode 100644 index 00000000..62f0490e --- /dev/null +++ b/internal/apps/message_gateway/secret.go @@ -0,0 +1,78 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package message_gateway + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + + "github.com/Rain-kl/Wavelet/internal/infra/config" + "github.com/Rain-kl/Wavelet/pkg/util" +) + +// CredentialKey is AES-256 hex derived from the session secret. +func CredentialKey() string { + secret := "" + if config.Config != nil { + secret = config.Config.App.SessionSecret + } + sum := sha256.Sum256([]byte(secret)) + return hex.EncodeToString(sum[:]) +} + +// EncryptCredentials encrypts a credential map as JSON. +func EncryptCredentials(creds map[string]string) (string, error) { + if creds == nil { + creds = map[string]string{} + } + raw, err := json.Marshal(creds) + if err != nil { + return "", err + } + return util.Encrypt(CredentialKey(), string(raw)) +} + +// DecryptCredentials decrypts a credential map. +func DecryptCredentials(ciphertext string) (map[string]string, error) { + if ciphertext == "" { + return map[string]string{}, nil + } + plain, err := util.Decrypt(CredentialKey(), ciphertext) + if err != nil { + return nil, err + } + var out map[string]string + if err := json.Unmarshal([]byte(plain), &out); err != nil { + return nil, err + } + if out == nil { + out = map[string]string{} + } + return out, nil +} + +// ParseExtra decodes optional extra JSON into a string map. +func ParseExtra(raw string) map[string]string { + if raw == "" { + return map[string]string{} + } + var out map[string]string + if err := json.Unmarshal([]byte(raw), &out); err != nil || out == nil { + return map[string]string{} + } + return out +} + +// EncodeExtra encodes extra fields as JSON. +func EncodeExtra(extra map[string]string) string { + if extra == nil { + return "" + } + raw, err := json.Marshal(extra) + if err != nil { + return "" + } + return string(raw) +} diff --git a/internal/cmd/all.go b/internal/cmd/all.go index db7f0799..21cc2e4a 100644 --- a/internal/cmd/all.go +++ b/internal/cmd/all.go @@ -6,9 +6,11 @@ package cmd import ( + "context" "log" "sync" + gwrunner "github.com/Rain-kl/Wavelet/internal/apps/message_gateway/runner" "github.com/Rain-kl/Wavelet/internal/infra/task/scheduler" "github.com/Rain-kl/Wavelet/internal/infra/task/worker" "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" @@ -36,6 +38,12 @@ var allCmd = &cobra.Command{ }) }() + go func() { + if err := gwrunner.Start(context.Background()); err != nil { + log.Printf("[All] message gateway stopped: %v", err) + } + }() + // 启动 Asynq Worker 任务处理服务 wg.Add(1) go func() { diff --git a/internal/cmd/worker.go b/internal/cmd/worker.go index 0c4d42be..d1ca1054 100644 --- a/internal/cmd/worker.go +++ b/internal/cmd/worker.go @@ -5,8 +5,10 @@ package cmd import ( + "context" "log" + gwrunner "github.com/Rain-kl/Wavelet/internal/apps/message_gateway/runner" "github.com/Rain-kl/Wavelet/internal/infra/task/worker" "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" @@ -19,6 +21,11 @@ var workerCmd = &cobra.Command{ Run: func(_ *cobra.Command, _ []string) { runBootstrap(bootstrap.Options{}) printStartupBanner(startupState{mode: "Worker", relationalDB: latestMigrationState.relationalDB, clickHouseDB: latestMigrationState.clickHouseDB}) + go func() { + if err := gwrunner.Start(context.Background()); err != nil { + log.Printf("[Worker] message gateway stopped: %v", err) + } + }() log.Println("[Worker] 启动任务处理服务") if err := worker.StartWorker(); err != nil { log.Fatalf("[工作器] 启动失败: %v", err)