feat(message-gateway): run adapters on worker and handle pairing inbound

This commit is contained in:
ryan
2026-08-16 12:14:17 +08:00
parent 7ca6dbe272
commit 09ec9d0af3
6 changed files with 538 additions and 0 deletions
@@ -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,
})
}
@@ -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)
}
}
@@ -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())
}
+78
View File
@@ -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)
}
+8
View File
@@ -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() {
+7
View File
@@ -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)