Files
ai-gateway-go/internal/workbench/notifications.go
T
superidou 5759c1862e AI Gateway Go 0.10.0 源码快照 + 旗舰版需求规划报告
M0-M7 已完成:核心网关(身份/RBAC/TOTP/OIDC/SAML/Provider/配额/路由/内容策略/审计/定价)+ 资源市场(MCP/Skills/数字员工)。
含 22 个 PostgreSQL 迁移、管理端/门户端前端源码、OpenAPI 契约、部署 compose。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-08-12 11:45:54 +08:00

438 lines
16 KiB
Go

package workbench
import (
"bytes"
"context"
"crypto/hmac"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"time"
"aigateway.local/core/internal/platform/cryptox"
"aigateway.local/core/internal/provider"
"github.com/jackc/pgx/v5"
"github.com/redis/go-redis/v9"
)
type NotificationInput struct {
Name, WebhookURL string
SigningSecret *string
EventPatterns []string
Enabled bool
}
type NotificationService struct {
assets *Service
cipher cryptox.Cipher
allowPrivate bool
client *http.Client
}
func NewNotificationService(assets *Service, cipher cryptox.Cipher, allowPrivate bool) *NotificationService {
// One shared client for all deliveries: http.Client is safe for concurrent
// use, and reusing the transport keeps TLS sessions and connections
// warm instead of re-handshaking for every webhook.
client := &http.Client{
Timeout: 15 * time.Second,
Transport: &http.Transport{
DialContext: safeToolDial(allowPrivate),
TLSHandshakeTimeout: 5 * time.Second,
ResponseHeaderTimeout: 10 * time.Second,
MaxIdleConns: 64,
MaxIdleConnsPerHost: 16,
IdleConnTimeout: 90 * time.Second,
},
CheckRedirect: func(*http.Request, []*http.Request) error { return errors.New("Webhook 不允许重定向") },
}
return &NotificationService{assets: assets, cipher: cipher, allowPrivate: allowPrivate, client: client}
}
const channelSelect = `SELECT id::text,name,webhook_url,event_patterns,enabled,octet_length(encrypted_signing_secret)>0,revision,created_at,updated_at,encrypted_signing_secret,signing_secret_kek_version FROM gateway.notification_channels`
func scanChannel(row pgx.Row) (NotificationChannel, error) {
var c NotificationChannel
err := row.Scan(&c.ID, &c.Name, &c.WebhookURL, &c.EventPatterns, &c.Enabled, &c.HasSigningSecret, &c.Revision, &c.CreatedAt, &c.UpdatedAt, &c.EncryptedSigningSecret, &c.SigningSecretKEKVersion)
return c, mapNotFound(err)
}
func (s *NotificationService) ListChannels(ctx context.Context) ([]NotificationChannel, error) {
rows, err := s.assets.pool.Query(ctx, channelSelect+` ORDER BY updated_at DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
items := []NotificationChannel{}
for rows.Next() {
c, err := scanChannel(rows)
if err != nil {
return nil, err
}
items = append(items, c)
}
return items, rows.Err()
}
func (s *NotificationService) GetChannel(ctx context.Context, id string) (NotificationChannel, error) {
return scanChannel(s.assets.pool.QueryRow(ctx, channelSelect+` WHERE id=$1`, id))
}
func (s *NotificationService) SaveChannel(ctx context.Context, id string, input NotificationInput, actorID string, create bool) (NotificationChannel, error) {
input.Name = strings.TrimSpace(input.Name)
input.WebhookURL = strings.TrimSpace(input.WebhookURL)
if input.Name == "" || len(input.Name) > 128 {
return NotificationChannel{}, errors.New("通知通道名称格式无效")
}
validated, err := provider.ValidateBaseURL(ctx, input.WebhookURL, s.allowPrivate)
if err != nil {
return NotificationChannel{}, fmt.Errorf("Webhook 地址校验失败: %w", err)
}
input.WebhookURL = validated
input.EventPatterns, err = normalizeStrings(input.EventPatterns, 100)
if err != nil {
return NotificationChannel{}, err
}
if len(input.EventPatterns) == 0 {
return NotificationChannel{}, errors.New("至少配置一个事件模式")
}
for _, pattern := range input.EventPatterns {
if strings.Count(pattern, "*") > 1 || (strings.Contains(pattern, "*") && !strings.HasSuffix(pattern, "*")) {
return NotificationChannel{}, errors.New("事件模式仅允许末尾 * 通配")
}
}
tx, err := s.assets.pool.Begin(ctx)
if err != nil {
return NotificationChannel{}, err
}
defer rollback(ctx, tx)
var encrypted []byte
var version int
if input.SigningSecret != nil && *input.SigningSecret != "" {
if len(*input.SigningSecret) > 4096 {
return NotificationChannel{}, errors.New("签名密钥过长")
}
encrypted, version, err = s.cipher.Encrypt([]byte(*input.SigningSecret))
if err != nil {
return NotificationChannel{}, err
}
}
if create {
id, err = newUUID()
if err != nil {
return NotificationChannel{}, err
}
_, err = tx.Exec(ctx, `INSERT INTO gateway.notification_channels(id,name,webhook_url,encrypted_signing_secret,signing_secret_kek_version,event_patterns,enabled,created_by) VALUES($1,$2,$3,$4,$5,$6,$7,$8)`, id, input.Name, input.WebhookURL, encrypted, version, input.EventPatterns, input.Enabled, actorID)
} else if input.SigningSecret == nil {
tag, updateErr := tx.Exec(ctx, `UPDATE gateway.notification_channels SET name=$2,webhook_url=$3,event_patterns=$4,enabled=$5,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1`, id, input.Name, input.WebhookURL, input.EventPatterns, input.Enabled)
err = updateErr
if err == nil && tag.RowsAffected() == 0 {
return NotificationChannel{}, ErrNotFound
}
} else {
tag, updateErr := tx.Exec(ctx, `UPDATE gateway.notification_channels SET name=$2,webhook_url=$3,encrypted_signing_secret=$4,signing_secret_kek_version=$5,event_patterns=$6,enabled=$7,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1`, id, input.Name, input.WebhookURL, encrypted, version, input.EventPatterns, input.Enabled)
err = updateErr
if err == nil && tag.RowsAffected() == 0 {
return NotificationChannel{}, ErrNotFound
}
}
if err != nil {
return NotificationChannel{}, err
}
event := "notification_channel.updated"
if create {
event = "notification_channel.created"
}
if err = emit(ctx, tx, event, "notification_channel", id, actorID, nil); err != nil {
return NotificationChannel{}, err
}
if err = tx.Commit(ctx); err != nil {
return NotificationChannel{}, err
}
return s.GetChannel(ctx, id)
}
func (s *NotificationService) DeleteChannel(ctx context.Context, id, actorID string) error {
tx, err := s.assets.pool.Begin(ctx)
if err != nil {
return err
}
defer rollback(ctx, tx)
tag, err := tx.Exec(ctx, `DELETE FROM gateway.notification_channels WHERE id=$1`, id)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrNotFound
}
if err = emit(ctx, tx, "notification_channel.deleted", "notification_channel", id, actorID, nil); err != nil {
return err
}
return tx.Commit(ctx)
}
const deliverySelect = `SELECT d.id::text,d.channel_id::text,c.name,d.event_id::text,d.event_type,d.status,d.last_error,d.payload,d.attempts,d.response_status,d.delivered_at,d.created_at,d.updated_at FROM gateway.notification_deliveries d JOIN gateway.notification_channels c ON c.id=d.channel_id`
func (s *NotificationService) ListDeliveries(ctx context.Context, limit int) ([]NotificationDelivery, error) {
if limit < 1 {
limit = 100
}
if limit > 500 {
limit = 500
}
rows, err := s.assets.pool.Query(ctx, deliverySelect+` ORDER BY d.created_at DESC LIMIT $1`, limit)
if err != nil {
return nil, err
}
defer rows.Close()
items := []NotificationDelivery{}
for rows.Next() {
var d NotificationDelivery
if err = rows.Scan(&d.ID, &d.ChannelID, &d.ChannelName, &d.EventID, &d.EventType, &d.Status, &d.LastError, &d.Payload, &d.Attempts, &d.ResponseStatus, &d.DeliveredAt, &d.CreatedAt, &d.UpdatedAt); err != nil {
return nil, err
}
items = append(items, d)
}
return items, rows.Err()
}
func matchesEvent(patterns []string, eventType string) bool {
for _, pattern := range patterns {
if pattern == "*" || pattern == eventType {
return true
}
if strings.HasSuffix(pattern, "*") && strings.HasPrefix(eventType, strings.TrimSuffix(pattern, "*")) {
return true
}
}
return false
}
func (s *NotificationService) ensureDelivery(ctx context.Context, channel NotificationChannel, eventID, eventType string, payload json.RawMessage) (NotificationDelivery, error) {
id, err := newUUID()
if err != nil {
return NotificationDelivery{}, err
}
var d NotificationDelivery
err = s.assets.pool.QueryRow(ctx, `INSERT INTO gateway.notification_deliveries(id,channel_id,event_id,event_type,payload,status) VALUES($1,$2,$3,$4,$5,'pending') ON CONFLICT(channel_id,event_id) DO UPDATE SET updated_at=gateway.notification_deliveries.updated_at RETURNING id::text,channel_id::text,$6,event_id::text,event_type,status,last_error,payload,attempts,response_status,delivered_at,created_at,updated_at`, id, channel.ID, eventID, eventType, payload, channel.Name).Scan(&d.ID, &d.ChannelID, &d.ChannelName, &d.EventID, &d.EventType, &d.Status, &d.LastError, &d.Payload, &d.Attempts, &d.ResponseStatus, &d.DeliveredAt, &d.CreatedAt, &d.UpdatedAt)
return d, err
}
func (s *NotificationService) deliver(ctx context.Context, channel NotificationChannel, delivery NotificationDelivery) error {
body, _ := json.Marshal(map[string]any{"event_id": delivery.EventID, "event_type": delivery.EventType, "occurred_at": delivery.CreatedAt.UTC().Format(time.RFC3339Nano), "payload": json.RawMessage(delivery.Payload)})
request, err := http.NewRequestWithContext(ctx, http.MethodPost, channel.WebhookURL, bytes.NewReader(body))
if err != nil {
return err
}
request.Header.Set("Content-Type", "application/json")
request.Header.Set("X-Gateway-Event-ID", delivery.EventID)
if len(channel.EncryptedSigningSecret) > 0 {
secret, decryptErr := s.cipher.Decrypt(channel.EncryptedSigningSecret, channel.SigningSecretKEKVersion)
if decryptErr != nil {
return decryptErr
}
mac := hmac.New(sha256.New, secret)
_, _ = mac.Write(body)
request.Header.Set("X-Gateway-Signature", "sha256="+hex.EncodeToString(mac.Sum(nil)))
}
response, requestErr := s.client.Do(request)
status := 0
if response != nil {
status = response.StatusCode
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 64<<10))
response.Body.Close()
}
success := requestErr == nil && status >= 200 && status < 300
message := ""
if requestErr != nil {
message = requestErr.Error()
} else if !success {
message = fmt.Sprintf("Webhook 返回 HTTP %d", status)
}
if len(message) > 1000 {
message = message[:1000]
}
_, dbErr := s.assets.pool.Exec(context.WithoutCancel(ctx), `UPDATE gateway.notification_deliveries SET status=$2,attempts=attempts+1,response_status=nullif($3,0),last_error=$4,delivered_at=CASE WHEN $2='delivered' THEN clock_timestamp() ELSE delivered_at END,updated_at=clock_timestamp() WHERE id=$1`, delivery.ID, map[bool]string{true: "delivered", false: "failed"}[success], status, message)
if dbErr != nil {
return dbErr
}
if !success {
return errors.New(message)
}
return nil
}
func (s *NotificationService) RetryDelivery(ctx context.Context, id string) error {
var d NotificationDelivery
err := s.assets.pool.QueryRow(ctx, deliverySelect+` WHERE d.id=$1`, id).Scan(&d.ID, &d.ChannelID, &d.ChannelName, &d.EventID, &d.EventType, &d.Status, &d.LastError, &d.Payload, &d.Attempts, &d.ResponseStatus, &d.DeliveredAt, &d.CreatedAt, &d.UpdatedAt)
if err != nil {
return mapNotFound(err)
}
channel, err := s.GetChannel(ctx, d.ChannelID)
if err != nil {
return err
}
return s.deliver(ctx, channel, d)
}
type NotificationDispatcher struct {
service *NotificationService
redis *redis.Client
stream, group, consumer string
logger *slog.Logger
}
func NewNotificationDispatcher(service *NotificationService, client *redis.Client, stream, consumer string, logger *slog.Logger) *NotificationDispatcher {
return &NotificationDispatcher{service: service, redis: client, stream: stream, group: "gateway-notifications-v1", consumer: consumer, logger: logger}
}
func (d *NotificationDispatcher) Run(ctx context.Context) error {
if err := d.redis.XGroupCreateMkStream(ctx, d.stream, d.group, "$").Err(); err != nil && !strings.Contains(err.Error(), "BUSYGROUP") {
return err
}
delay := time.Duration(0)
for {
// Compensate for events claimed by a previous iteration (or a crashed
// worker) that were never acknowledged. XReadGroup ">" never re-reads
// the pending list, so without this pass those events would be silently
// dropped for good.
if err := d.reclaim(ctx); err != nil {
if ctx.Err() != nil {
return nil
}
if d.logger != nil {
d.logger.Error("notification pending reclaim failed; will retry", "error", err)
}
}
streams, err := d.redis.XReadGroup(ctx, &redis.XReadGroupArgs{Group: d.group, Consumer: d.consumer, Streams: []string{d.stream, ">"}, Count: 20, Block: 5 * time.Second}).Result()
if errors.Is(err, redis.Nil) {
delay = 0
continue
}
if err != nil {
if ctx.Err() != nil {
return nil
}
if d.logger != nil {
d.logger.Error("notification stream read failed; backing off", "error", err)
}
// Reconnect with backoff instead of killing the worker: a transient
// Redis blip must not take the notification worker down and let the
// stream trim every queued event.
if !d.wait(ctx, &delay) {
return nil
}
continue
}
delay = 0
for _, stream := range streams {
for _, message := range stream.Messages {
if err = d.handle(ctx, message); err != nil {
// Do not acknowledge: the event stays in the pending list and
// is retried by the reclaim pass above.
if d.logger != nil {
d.logger.Error("notification event failed; will retry", "stream_id", message.ID, "error", err)
}
continue
}
if err = d.redis.XAck(ctx, d.stream, d.group, message.ID).Err(); err != nil {
if ctx.Err() != nil {
return nil
}
// An ack failure must not drop a successfully handled event;
// it is reclaimed and re-acked on the next pass (handle is
// idempotent via the delivery table).
if d.logger != nil {
d.logger.Error("notification ack failed; will retry via reclaim", "stream_id", message.ID, "error", err)
}
continue
}
}
}
}
}
// reclaim re-processes pending stream entries that have been idle longer than
// the threshold. handle is idempotent (deliveries are keyed by channel+event),
// so re-running it only creates the delivery rows that were never created or
// updates attempts on ones already recorded.
func (d *NotificationDispatcher) reclaim(ctx context.Context) error {
for {
messages, cursor, err := d.redis.XAutoClaim(ctx, &redis.XAutoClaimArgs{
Stream: d.stream,
Group: d.group,
Consumer: d.consumer,
MinIdle: 30 * time.Second,
Start: "0",
Count: 20,
}).Result()
if err != nil {
return err
}
for _, message := range messages {
if err = d.handle(ctx, message); err != nil {
if d.logger != nil {
d.logger.Warn("notification pending retry failed", "stream_id", message.ID, "error", err)
}
continue
}
if err = d.redis.XAck(ctx, d.stream, d.group, message.ID).Err(); err != nil {
return err
}
}
if cursor == "0-0" {
break
}
}
return nil
}
// wait sleeps with exponential backoff, capping at 30s. It returns false when
// the context is cancelled.
func (d *NotificationDispatcher) wait(ctx context.Context, delay *time.Duration) bool {
const max = 30 * time.Second
if *delay == 0 {
*delay = 200 * time.Millisecond
} else {
*delay *= 2
if *delay > max {
*delay = max
}
}
timer := time.NewTimer(*delay)
defer timer.Stop()
select {
case <-ctx.Done():
return false
case <-timer.C:
return true
}
}
func (d *NotificationDispatcher) handle(ctx context.Context, message redis.XMessage) error {
eventID := fmt.Sprint(message.Values["event_id"])
eventType := fmt.Sprint(message.Values["event_type"])
payload := json.RawMessage(fmt.Sprint(message.Values["payload"]))
channels, err := d.service.ListChannels(ctx)
if err != nil {
return err
}
for _, channel := range channels {
if !channel.Enabled || !matchesEvent(channel.EventPatterns, eventType) {
continue
}
delivery, deliveryErr := d.service.ensureDelivery(ctx, channel, eventID, eventType, payload)
if deliveryErr != nil {
return deliveryErr
}
if delivery.Status == "delivered" {
continue
}
if deliveryErr = d.service.deliver(ctx, channel, delivery); deliveryErr != nil && d.logger != nil {
d.logger.Warn("webhook delivery failed", "channel", channel.Name, "event_id", eventID, "error", deliveryErr)
}
}
return nil
}