Files
aiaa-notification-server/internal/subscriber/subscriber.go
T
ryan b396638438 feat(交易信号): 持仓状态持久化至 Redis 以保障重启与多实例下均价计算准确
Motivation:
持仓与已处理信号的状态此前仅存于进程内存,服务重启或多实例部署后会丢失,导致加仓均价、仓位大小等计算失真,重复信号也无法跨实例幂等。通过将状态持久化到 Redis,保证持仓跟踪跨重启、跨实例连续一致,提升通知内容的准确性与可靠性。

Changes:

* 新增持仓存储抽象,支持内存与 Redis 两种实现,持仓状态与已处理信号快照按 TTL 持久化
* 持仓变更通过 Redis 事务管道原子提交,保证状态更新与幂等记录一致写入
* 存储写入失败时消息进入重试而非直接确认,避免状态丢失导致通知失真
* 缓存层新增原始值读取与批量事务写入能力,并在订阅器初始化时注入 Redis 依赖
* 补充跨实例持久化、幂等去重与存储失败场景的测试覆盖
2026-08-23 14:43:53 +08:00

216 lines
5.4 KiB
Go

package subscriber
import (
"context"
"fmt"
"log/slog"
"net/url"
"time"
"aiaa-notification-service/internal/cache"
"aiaa-notification-service/internal/config"
"aiaa-notification-service/internal/subscriber/cryptostrategy"
"aiaa-notification-service/internal/subscriber/tradesignal"
amqp "github.com/rabbitmq/amqp091-go"
)
type Subscriber struct {
cfg config.SubscriptionConfig
conv MessageConverter
lookup SourceLookup
process ProcessFunc
deduper Deduper
}
func New(cfg config.SubscriptionConfig, lookup SourceLookup, process ProcessFunc, deduper Deduper, redisCache *cache.Cache) (*Subscriber, error) {
var conv MessageConverter
switch cfg.Formatter {
case "trade_signal":
conv = tradesignal.NewConverterWithCache(cfg.StrategyOverrides, redisCache)
case "crypto_strategy":
conv = cryptostrategy.NewConverter()
default:
return nil, fmt.Errorf("unknown formatter %q", cfg.Formatter)
}
return &Subscriber{
cfg: cfg,
conv: conv,
lookup: lookup,
process: process,
deduper: deduper,
}, nil
}
func (s *Subscriber) Run(ctx context.Context) error {
for {
select {
case <-ctx.Done():
return ctx.Err()
default:
}
if err := s.consumeOnce(ctx); err != nil {
slog.Error("subscriber error, reconnecting",
"name", s.cfg.Name,
"host", brokerHost(s.cfg.URL),
"error", err.Error())
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(5 * time.Second):
}
}
}
}
func (s *Subscriber) consumeOnce(ctx context.Context) error {
conn, err := amqp.Dial(s.cfg.URL)
if err != nil {
return fmt.Errorf("dial rabbitmq: %w", err)
}
defer conn.Close()
ch, err := conn.Channel()
if err != nil {
return fmt.Errorf("open channel: %w", err)
}
defer ch.Close()
if err := ch.Qos(1, 0, false); err != nil {
return fmt.Errorf("set qos: %w", err)
}
if err := s.ensureQueue(ch); err != nil {
return err
}
deliveries, err := ch.Consume(s.cfg.Queue, s.cfg.Name, false, false, false, false, nil)
if err != nil {
return fmt.Errorf("consume queue %s: %w", s.cfg.Queue, err)
}
slog.Info("listening on queue", "name", s.cfg.Name, "queue", s.cfg.Queue)
for {
select {
case <-ctx.Done():
return ctx.Err()
case d, ok := <-deliveries:
if !ok {
return fmt.Errorf("delivery channel closed")
}
s.handleDelivery(ch, d)
}
}
}
func (s *Subscriber) ensureQueue(ch *amqp.Channel) error {
if s.cfg.Exchange != "" {
if err := ch.ExchangeDeclare(s.cfg.Exchange, s.cfg.ExchangeType, true, false, false, false, nil); err != nil {
return fmt.Errorf("declare exchange %q: %w", s.cfg.Exchange, err)
}
}
if _, err := ch.QueueDeclare(s.cfg.Queue, true, false, false, false, nil); err != nil {
return fmt.Errorf("declare queue %q: %w", s.cfg.Queue, err)
}
if s.cfg.Exchange != "" {
if err := ch.QueueBind(s.cfg.Queue, s.cfg.RoutingKey, s.cfg.Exchange, false, nil); err != nil {
return fmt.Errorf("bind queue %q to exchange %q: %w", s.cfg.Queue, s.cfg.Exchange, err)
}
slog.Info("queue bound", "queue", s.cfg.Queue, "exchange", s.cfg.Exchange, "type", s.cfg.ExchangeType, "routing_key", s.cfg.RoutingKey)
}
if s.cfg.DeadLetterQueue != "" {
if _, err := ch.QueueDeclare(s.cfg.DeadLetterQueue, true, false, false, false, nil); err != nil {
slog.Warn("declare dead letter queue failed", "queue", s.cfg.DeadLetterQueue, "error", err.Error())
}
}
return nil
}
func (s *Subscriber) handleDelivery(ch *amqp.Channel, d amqp.Delivery) {
slog.Info("raw message received",
"name", s.cfg.Name,
"queue", s.cfg.Queue,
"source", s.cfg.Source,
"delivery_tag", d.DeliveryTag,
"body", string(d.Body),
)
disp := HandleMessage(context.Background(), HandleInput{
Body: d.Body,
Headers: map[string]any(d.Headers),
Name: s.cfg.Name,
Queue: s.cfg.Queue,
SourceName: s.cfg.Source,
MaxRetry: s.cfg.MaxRetry,
Deduper: s.deduper,
}, s.conv, s.lookup, s.process)
switch disp {
case DispositionRetry:
s.republish(ch, d, s.cfg.Queue)
case DispositionDLQ:
if s.cfg.DeadLetterQueue == "" {
slog.Warn("max retry reached, discarded", "name", s.cfg.Name)
_ = d.Ack(false)
return
}
s.republish(ch, d, s.cfg.DeadLetterQueue)
default:
_ = d.Ack(false)
}
}
func (s *Subscriber) republish(ch *amqp.Channel, d amqp.Delivery, queue string) {
headers := copyAMQPHeaders(d.Headers)
headers[retryHeader] = RetryCount(map[string]any(d.Headers)) + 1
if err := publishToQueue(ch, queue, d.Body, headers); err != nil {
slog.Error("requeue failed", "queue", queue, "error", err.Error())
_ = d.Nack(false, true)
return
}
_ = d.Ack(false)
if queue == s.cfg.DeadLetterQueue {
slog.Warn("message moved to dlq", "name", s.cfg.Name, "retries", s.cfg.MaxRetry)
return
}
slog.Info("message requeued", "name", s.cfg.Name, "retry", headers[retryHeader], "max", s.cfg.MaxRetry)
}
func brokerHost(raw string) string {
u, err := url.Parse(raw)
if err != nil || u.Hostname() == "" {
return ""
}
if u.Port() != "" {
return u.Hostname() + ":" + u.Port()
}
return u.Hostname()
}
func copyAMQPHeaders(headers amqp.Table) amqp.Table {
out := amqp.Table{}
for k, v := range headers {
out[k] = v
}
return out
}
func publishToQueue(ch *amqp.Channel, queue string, body []byte, headers amqp.Table) error {
return ch.Publish(
"",
queue,
false,
false,
amqp.Publishing{
ContentType: "application/json",
DeliveryMode: amqp.Persistent,
Headers: headers,
Body: body,
Timestamp: time.Now(),
},
)
}