Files
aiaa-notification-server/internal/store/rule.go
T
ryan d2e8476398 feat(规则): 支持同一事件多规则按条件区分并全部发送
Motivation:
止盈、止损等不同策略信号会映射到同一事件(如 trade.close),但原有唯一约束要求每个事件只能有一条规则,无法按策略区分处理。放开该约束后,同一事件可配置多条规则,通过规则条件与精确/通配优先级区分,命中条件的规则全部发送;同时将去重键从消息原文改为信号维度,避免同一信号因时间戳等无关字段差异被误判为重复。

Changes:

* 移除规则 source_id+event 的唯一约束,改为普通索引
* 事件匹配改为返回命中优先级内所有启用规则,并按 ID 逐条派发
* 通知服务遍历多条规则,按条件过滤后聚合发送渠道
* 去重键由消息 body 哈希改为策略/币种/周期/方向/价格信号维度哈希
2026-08-17 01:20:49 +08:00

197 lines
6.5 KiB
Go

package store
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"aiaa-notification-service/internal/model"
)
func (s *Store) CreateRule(ctx context.Context, r *model.Rule, channelIDs []int) error {
tx, err := s.DB.BeginTxx(ctx, nil)
if err != nil {
return fmt.Errorf("begin tx: %w", err)
}
defer tx.Rollback()
query := `INSERT INTO notification_rule (name, source_id, event, template_id, conditions, enabled) VALUES (?, ?, ?, ?, ?, ?)`
condsJSON, err := marshalJSON(r.Conditions)
if err != nil {
return fmt.Errorf("marshal conditions: %w", err)
}
result, err := tx.ExecContext(ctx, query, r.Name, r.SourceID, r.Event, r.TemplateID, condsJSON, r.Enabled)
if err != nil {
return fmt.Errorf("create rule: %w", err)
}
id, _ := result.LastInsertId()
r.ID = int(id)
for _, chID := range channelIDs {
_, err := tx.ExecContext(ctx, `INSERT INTO notification_rule_channel (rule_id, channel_id, enabled) VALUES (?, ?, 1)`, r.ID, chID)
if err != nil {
return fmt.Errorf("add rule_channel: %w", err)
}
}
return tx.Commit()
}
func (s *Store) GetRule(ctx context.Context, id int) (*model.Rule, error) {
var r model.Rule
var condsBytes []byte
row := s.DB.QueryRowContext(ctx, `SELECT id, name, source_id, event, template_id, conditions, enabled, created_at, updated_at FROM notification_rule WHERE id = ?`, id)
if err := row.Scan(&r.ID, &r.Name, &r.SourceID, &r.Event, &r.TemplateID, &condsBytes, &r.Enabled, &r.CreatedAt, &r.UpdatedAt); err != nil {
return nil, fmt.Errorf("get rule %d: %w", id, err)
}
if len(condsBytes) > 0 && string(condsBytes) != "null" {
raw := json.RawMessage(condsBytes)
r.Conditions = &raw
}
return &r, nil
}
func (s *Store) GetRuleBySourceEvent(ctx context.Context, sourceID int, event string) (*model.Rule, error) {
var r model.Rule
var condsBytes []byte
query := `SELECT id, name, source_id, event, template_id, conditions, enabled, created_at, updated_at FROM notification_rule WHERE source_id = ? AND event = ? AND enabled = 1 ORDER BY id LIMIT 1`
row := s.DB.QueryRowContext(ctx, query, sourceID, event)
if err := row.Scan(&r.ID, &r.Name, &r.SourceID, &r.Event, &r.TemplateID, &condsBytes, &r.Enabled, &r.CreatedAt, &r.UpdatedAt); err != nil {
return nil, fmt.Errorf("get rule by source+event: %w", err)
}
if len(condsBytes) > 0 && string(condsBytes) != "null" {
raw := json.RawMessage(condsBytes)
r.Conditions = &raw
}
return &r, nil
}
func (s *Store) ListEnabledRulesBySource(ctx context.Context, sourceID int) ([]model.Rule, error) {
rows, err := s.DB.QueryContext(ctx, `SELECT id, name, source_id, event, template_id, conditions, enabled, created_at, updated_at FROM notification_rule WHERE source_id = ? AND enabled = 1`, sourceID)
if err != nil {
return nil, fmt.Errorf("list enabled rules by source: %w", err)
}
defer rows.Close()
rules, err := scanRules(rows)
if err != nil {
return nil, err
}
return rules, nil
}
func (s *Store) ListRules(ctx context.Context, page PageFilter) ([]model.Rule, int, error) {
var count int
if err := s.DB.GetContext(ctx, &count, `SELECT COUNT(*) FROM notification_rule`); err != nil {
return nil, 0, fmt.Errorf("count rules: %w", err)
}
page.Normalize()
rows, err := s.DB.QueryContext(ctx, `SELECT id, name, source_id, event, template_id, conditions, enabled, created_at, updated_at FROM notification_rule ORDER BY id LIMIT ? OFFSET ?`, page.PageSize, page.Offset())
if err != nil {
return nil, 0, fmt.Errorf("list rules: %w", err)
}
defer rows.Close()
rules, err := scanRules(rows)
if err != nil {
return nil, 0, err
}
return rules, count, nil
}
func (s *Store) UpdateRule(ctx context.Context, id int, r *model.Rule, channelIDs []int) error {
tx, err := s.DB.BeginTxx(ctx, nil)
if err != nil {
return fmt.Errorf("begin tx: %w", err)
}
defer tx.Rollback()
condsJSON, err := marshalJSON(r.Conditions)
if err != nil {
return fmt.Errorf("marshal conditions: %w", err)
}
_, err = tx.ExecContext(ctx, `UPDATE notification_rule SET name=?, source_id=?, event=?, template_id=?, conditions=?, enabled=? WHERE id=?`,
r.Name, r.SourceID, r.Event, r.TemplateID, condsJSON, r.Enabled, id)
if err != nil {
return fmt.Errorf("update rule: %w", err)
}
if channelIDs != nil {
if _, err := tx.ExecContext(ctx, `DELETE FROM notification_rule_channel WHERE rule_id = ?`, id); err != nil {
return fmt.Errorf("delete rule channels: %w", err)
}
for _, chID := range channelIDs {
_, err := tx.ExecContext(ctx, `INSERT INTO notification_rule_channel (rule_id, channel_id, enabled) VALUES (?, ?, 1)`, id, chID)
if err != nil {
return fmt.Errorf("add rule_channel: %w", err)
}
}
}
return tx.Commit()
}
func (s *Store) DeleteRule(ctx context.Context, id int) error {
_, err := s.DB.ExecContext(ctx, `DELETE FROM notification_rule WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete rule %d: %w", id, err)
}
return nil
}
func (s *Store) SetRuleEnabled(ctx context.Context, id int, enabled bool) error {
v := 0
if enabled {
v = 1
}
_, err := s.DB.ExecContext(ctx, `UPDATE notification_rule SET enabled = ? WHERE id = ?`, v, id)
if err != nil {
return fmt.Errorf("set rule enabled %d: %w", id, err)
}
return nil
}
func (s *Store) GetRuleChannels(ctx context.Context, ruleID int) ([]model.RuleChannel, error) {
rcs := make([]model.RuleChannel, 0)
err := s.DB.SelectContext(ctx, &rcs, `SELECT id, rule_id, channel_id, enabled FROM notification_rule_channel WHERE rule_id = ? AND enabled = 1`, ruleID)
if err != nil {
return nil, fmt.Errorf("get rule channels: %w", err)
}
return rcs, nil
}
func (s *Store) SetRuleChannelEnabled(ctx context.Context, ruleID, channelID int, enabled bool) error {
v := 0
if enabled {
v = 1
}
_, err := s.DB.ExecContext(ctx, `UPDATE notification_rule_channel SET enabled = ? WHERE rule_id = ? AND channel_id = ?`, v, ruleID, channelID)
if err != nil {
return fmt.Errorf("set rule channel enabled %d/%d: %w", ruleID, channelID, err)
}
return nil
}
// helpers
func marshalJSON(v *json.RawMessage) (interface{}, error) {
if v == nil {
return nil, nil
}
return json.Marshal(v)
}
func scanRules(rows *sql.Rows) ([]model.Rule, error) {
rules := make([]model.Rule, 0)
for rows.Next() {
var r model.Rule
var condsBytes []byte
if err := rows.Scan(&r.ID, &r.Name, &r.SourceID, &r.Event, &r.TemplateID, &condsBytes, &r.Enabled, &r.CreatedAt, &r.UpdatedAt); err != nil {
return nil, err
}
if len(condsBytes) > 0 && string(condsBytes) != "null" {
raw := json.RawMessage(condsBytes)
r.Conditions = &raw
}
rules = append(rules, r)
}
return rules, rows.Err()
}