This commit is contained in:
2026-09-16 06:26:16 +02:00
parent 95b0779db0
commit 4f8e0bbb04
51 changed files with 6290 additions and 0 deletions
+101
View File
@@ -0,0 +1,101 @@
package auth
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"fmt"
"golang.org/x/crypto/bcrypt"
"strconv"
"strings"
"time"
)
const passwordRounds = 150000
func HashPassword(password string) (string, error) {
b, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
return string(b), err
}
func RandomSecret() string {
b := make([]byte, 32)
if _, err := rand.Read(b); err != nil {
panic(err)
}
return base64.RawURLEncoding.EncodeToString(b)
}
func CheckPassword(encoded, password string) bool {
if strings.HasPrefix(encoded, "$2") {
return bcrypt.CompareHashAndPassword([]byte(encoded), []byte(password)) == nil
}
parts := strings.Split(encoded, "$")
if len(parts) != 4 || parts[0] != "sha256" {
return false
}
rounds, err := strconv.Atoi(parts[1])
if err != nil || rounds < 10000 || rounds > 1000000 {
return false
}
salt, err := base64.RawStdEncoding.DecodeString(parts[2])
if err != nil {
return false
}
want, err := base64.RawStdEncoding.DecodeString(parts[3])
if err != nil {
return false
}
got := stretch([]byte(password), salt, rounds)
return len(got) == len(want) && subtle.ConstantTimeCompare(got, want) == 1
}
func stretch(password, salt []byte, rounds int) []byte {
buf := append(append([]byte{}, salt...), password...)
s := sha256.Sum256(buf)
out := s[:]
for i := 1; i < rounds; i++ {
h := sha256.New()
h.Write(out)
h.Write(salt)
h.Write(password)
out = h.Sum(nil)
}
return out
}
func SignSession(secret, username string, expires time.Time) string {
payload := fmt.Sprintf("%s|%d", username, expires.Unix())
mac := hmac.New(sha256.New, []byte(secret))
mac.Write([]byte(payload))
sig := base64.RawURLEncoding.EncodeToString(mac.Sum(nil))
return base64.RawURLEncoding.EncodeToString([]byte(payload)) + "." + sig
}
func VerifySession(secret, token string, now time.Time) (string, bool) {
parts := strings.Split(token, ".")
if len(parts) != 2 {
return "", false
}
payloadBytes, err := base64.RawURLEncoding.DecodeString(parts[0])
if err != nil {
return "", false
}
mac := hmac.New(sha256.New, []byte(secret))
mac.Write(payloadBytes)
want := base64.RawURLEncoding.EncodeToString(mac.Sum(nil))
if subtle.ConstantTimeCompare([]byte(want), []byte(parts[1])) != 1 {
return "", false
}
fields := strings.Split(string(payloadBytes), "|")
if len(fields) != 2 {
return "", false
}
exp, err := strconv.ParseInt(fields[1], 10, 64)
if err != nil || now.Unix() > exp {
return "", false
}
return fields[0], true
}
+39
View File
@@ -0,0 +1,39 @@
package auth
import (
"encoding/base64"
"fmt"
"strings"
"testing"
"time"
)
func TestPassword(t *testing.T) {
h, e := HashPassword("secret")
if e != nil || !CheckPassword(h, "secret") || CheckPassword(h, "wrong") {
t.Fatal("password hash check failed")
}
if !strings.HasPrefix(h, "$2") {
t.Fatal("new hashes must use bcrypt")
}
}
func TestLegacyPasswordCompatibility(t *testing.T) {
salt := []byte("sixteen-byte-salt")
password := "legacy password"
hash := stretch([]byte(password), salt, passwordRounds)
encoded := fmt.Sprintf("sha256$%d$%s$%s", passwordRounds, base64.RawStdEncoding.EncodeToString(salt), base64.RawStdEncoding.EncodeToString(hash))
if !CheckPassword(encoded, password) || CheckPassword(encoded, "wrong") {
t.Fatal("legacy hash compatibility")
}
if CheckPassword(strings.Replace(encoded, "150000", "999999999", 1), password) {
t.Fatal("unbounded legacy rounds")
}
}
func TestSession(t *testing.T) {
tok := SignSession("s", "admin", time.Now().Add(time.Hour))
u, ok := VerifySession("s", tok, time.Now())
if !ok || u != "admin" {
t.Fatal("session verify failed")
}
}
+277
View File
@@ -0,0 +1,277 @@
package config
import (
"crypto/rand"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"regexp"
"strings"
"sync"
)
const CurrentVersion = 1
type Config struct {
Version int `json:"version"`
Server ServerConfig `json:"server"`
Divera DiveraConfig `json:"divera"`
Ingress IngressConfig `json:"ingress"`
Mappings []Mapping `json:"mappings"`
Outbounds []OutboundConfig `json:"outbounds,omitempty"`
}
type ServerConfig struct {
Listen string `json:"listen"`
AdminUsername string `json:"admin_username"`
AdminPasswordHash string `json:"admin_password_hash"`
SessionSecret string `json:"session_secret"`
}
type DiveraConfig struct {
BaseURL string `json:"base_url"`
AccessKey string `json:"access_key"`
UCR int64 `json:"ucr,omitempty"`
TimeoutS int `json:"timeout_seconds"`
DryRun bool `json:"dry_run"`
}
type IngressConfig struct {
Mail []MailIngress `json:"mail,omitempty"`
Discord DiscordIngress `json:"discord"`
AllowUnauthenticated bool `json:"allow_unauthenticated"`
NtfyTokens map[string]string `json:"ntfy_tokens"`
WebhookTokens map[string]string `json:"webhook_tokens"`
GotifyTokens []string `json:"gotify_tokens"`
}
type Mapping struct {
OutboundID string `json:"outbound_id,omitempty"`
ID string `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
Source string `json:"source"`
ChannelRegex string `json:"channel_regex,omitempty"`
TitleRegex string `json:"title_regex,omitempty"`
MessageRegex string `json:"message_regex,omitempty"`
MinPriority int `json:"min_priority,omitempty"`
Target string `json:"target"`
TitleTemplate string `json:"title_template"`
TextTemplate string `json:"text_template"`
AddressTemplate string `json:"address_template,omitempty"`
NotificationType int `json:"notification_type"`
Clusters []int64 `json:"clusters,omitempty"`
ClusterRoutes map[string]int `json:"cluster_routes,omitempty"`
Groups []int64 `json:"groups,omitempty"`
Users []int64 `json:"users,omitempty"`
Vehicles []int64 `json:"vehicles,omitempty"`
SendPush bool `json:"send_push"`
SendSMS bool `json:"send_sms"`
SendCall bool `json:"send_call"`
SendMail bool `json:"send_mail"`
SendPager bool `json:"send_pager"`
PrivateMode bool `json:"private_mode"`
Extra map[string]any `json:"extra,omitempty"`
}
func Default() Config {
return Config{
Version: CurrentVersion,
Server: ServerConfig{Listen: ":8080", AdminUsername: "admin", SessionSecret: randomSecret()},
Divera: DiveraConfig{BaseURL: "https://app.divera247.com", TimeoutS: 15, DryRun: true},
Ingress: IngressConfig{
NtfyTokens: make(map[string]string), WebhookTokens: make(map[string]string), GotifyTokens: []string{},
},
Mappings: []Mapping{
{
ID: "example-ntfy-alarm", Name: "Beispiel ntfy -> Alarm", Enabled: false,
Source: "ntfy", ChannelRegex: "^alarm$", Target: "alarm",
TitleTemplate: "{{if .Title}}{{.Title}}{{else}}ALARM{{end}}", TextTemplate: "{{.Message}}",
NotificationType: 2, SendPush: true,
},
},
}
}
func randomSecret() string {
b := make([]byte, 32)
if _, err := rand.Read(b); err != nil {
panic(err)
}
return base64.RawURLEncoding.EncodeToString(b)
}
type Store struct {
path string
mu sync.RWMutex
cfg Config
}
func Open(path string) (*Store, error) {
s := &Store{path: path}
b, err := os.ReadFile(path)
if errors.Is(err, os.ErrNotExist) {
s.cfg = Default()
if err := s.saveLocked(); err != nil {
return nil, err
}
return s, nil
}
if err != nil {
return nil, err
}
if err := json.Unmarshal(b, &s.cfg); err != nil {
return nil, fmt.Errorf("parse config: %w", err)
}
normalize(&s.cfg)
return s, nil
}
func normalize(c *Config) {
if c.Version == 0 {
c.Version = CurrentVersion
}
if c.Server.Listen == "" {
c.Server.Listen = ":8080"
}
if c.Server.AdminUsername == "" {
c.Server.AdminUsername = "admin"
}
if c.Server.SessionSecret == "" {
c.Server.SessionSecret = randomSecret()
}
if c.Divera.BaseURL == "" {
c.Divera.BaseURL = "https://app.divera247.com"
}
if c.Divera.TimeoutS <= 0 {
c.Divera.TimeoutS = 15
}
if c.Ingress.NtfyTokens == nil {
c.Ingress.NtfyTokens = map[string]string{}
}
if c.Ingress.WebhookTokens == nil {
c.Ingress.WebhookTokens = map[string]string{}
}
if c.Ingress.GotifyTokens == nil {
c.Ingress.GotifyTokens = []string{}
}
}
func (s *Store) Get() Config {
s.mu.RLock()
defer s.mu.RUnlock()
b, _ := json.Marshal(s.cfg)
var out Config
_ = json.Unmarshal(b, &out)
return out
}
func (s *Store) Replace(c Config) error {
normalize(&c)
if err := Validate(c); err != nil {
return err
}
if c.Version != CurrentVersion {
return fmt.Errorf("unsupported config version %d", c.Version)
}
s.mu.Lock()
defer s.mu.Unlock()
old := s.cfg
s.cfg = c
if err := s.saveLocked(); err != nil {
s.cfg = old
return err
}
return nil
}
func (s *Store) Update(fn func(*Config) error) error {
s.mu.Lock()
defer s.mu.Unlock()
copyBytes, _ := json.Marshal(s.cfg)
var next Config
_ = json.Unmarshal(copyBytes, &next)
if err := fn(&next); err != nil {
return err
}
normalize(&next)
old := s.cfg
s.cfg = next
if err := s.saveLocked(); err != nil {
s.cfg = old
return err
}
return nil
}
func (s *Store) saveLocked() error {
if err := os.MkdirAll(filepath.Dir(s.path), 0700); err != nil {
return err
}
b, err := json.MarshalIndent(s.cfg, "", " ")
if err != nil {
return err
}
tmp := s.path + ".tmp"
if err := os.WriteFile(tmp, append(b, '\n'), 0600); err != nil {
return err
}
if err := os.Chmod(tmp, 0600); err != nil {
return err
}
return os.Rename(tmp, s.path)
}
// Validate checks user-editable configuration before it is persisted.
func Validate(c Config) error {
if err := validateIngress(c.Ingress); err != nil {
return err
}
if strings.TrimSpace(c.Server.Listen) == "" {
return fmt.Errorf("server.listen darf nicht leer sein")
}
if strings.TrimSpace(c.Server.AdminUsername) == "" {
return fmt.Errorf("server.admin_username darf nicht leer sein")
}
if c.Divera.TimeoutS < 1 || c.Divera.TimeoutS > 300 {
return fmt.Errorf("divera.timeout_seconds muss zwischen 1 und 300 liegen")
}
if err := validateOutbounds(c); err != nil {
return err
}
seen := map[string]bool{}
for i, m := range c.Mappings {
if strings.TrimSpace(m.ID) == "" {
return fmt.Errorf("mapping %d: id fehlt", i+1)
}
if seen[m.ID] {
return fmt.Errorf("mapping id %q ist doppelt", m.ID)
}
seen[m.ID] = true
switch m.Source {
case "", "any", "ntfy", "gotify", "webhook", "mail", "discord":
default:
return fmt.Errorf("mapping %q: unbekannte quelle %q", m.ID, m.Source)
}
switch strings.ToLower(m.Target) {
case "alarm", "alarms", "news", "message", "mitteilung", "event", "termin", "discord", "webhook", "smtp", "ntfy", "gotify":
default:
return fmt.Errorf("mapping %q: unbekanntes ziel %q", m.ID, m.Target)
}
for field, pattern := range map[string]string{"channel_regex": m.ChannelRegex, "title_regex": m.TitleRegex, "message_regex": m.MessageRegex} {
if pattern == "" {
continue
}
if _, err := regexp.Compile(pattern); err != nil {
return fmt.Errorf("mapping %q: %s ist kein gültiger regulärer Ausdruck: %v", m.ID, field, err)
}
}
if m.NotificationType < 0 || m.NotificationType > 4 {
return fmt.Errorf("mapping %q: notification_type muss 0..4 sein", m.ID)
}
}
return nil
}
+66
View File
@@ -0,0 +1,66 @@
package config
import (
"crypto/ed25519"
"encoding/hex"
"fmt"
"net"
"strconv"
"strings"
)
type MailIngress struct {
ID string `json:"id"`
Enabled bool `json:"enabled"`
Address string `json:"address"`
Username string `json:"username"`
Password string `json:"password"`
Folder string `json:"folder"`
Channel string `json:"channel"`
PollSeconds int `json:"poll_seconds"`
From []string `json:"from,omitempty"`
To []string `json:"to,omitempty"`
ImportExisting bool `json:"import_existing"`
}
type DiscordIngress struct {
Enabled bool `json:"enabled"`
PublicKey string `json:"public_key"`
ApplicationID string `json:"application_id"`
GuildIDs []string `json:"guild_ids"`
ChannelIDs []string `json:"channel_ids"`
Command string `json:"command"`
}
func validateIngress(c IngressConfig) error {
seen := map[string]bool{}
for _, m := range c.Mail {
if m.ID == "" || seen[m.ID] {
return fmt.Errorf("Mail-Eingang benötigt eindeutige ID")
}
seen[m.ID] = true
if !m.Enabled {
continue
}
host, port, err := net.SplitHostPort(m.Address)
p, _ := strconv.Atoi(port)
if err != nil || host == "" || p < 1 || p > 65535 || m.Username == "" || m.Password == "" || m.Channel == "" {
return fmt.Errorf("Mail-Eingang %q: Adresse (Host:Port), Benutzer, Passwort und Kanal erforderlich", m.ID)
}
if m.PollSeconds < 10 || m.PollSeconds > 86400 {
return fmt.Errorf("Mail-Abfrageintervall muss 10..86400 Sekunden sein")
}
}
d := c.Discord
if d.Enabled {
for _, id := range append(append([]string{}, d.GuildIDs...), d.ChannelIDs...) {
if strings.TrimSpace(id) == "" {
return fmt.Errorf("Discord-Allowlists dürfen keine leeren IDs enthalten")
}
}
key, err := hex.DecodeString(d.PublicKey)
if err != nil || len(key) != ed25519.PublicKeySize || d.ApplicationID == "" || d.Command == "" || len(d.GuildIDs) == 0 || len(d.ChannelIDs) == 0 {
return fmt.Errorf("Discord benötigt Public Key, Application-ID, Command und Guild-/Channel-Allowlist")
}
}
return nil
}
+138
View File
@@ -0,0 +1,138 @@
package config
import (
"fmt"
"net"
"net/mail"
"net/url"
"os"
"regexp"
"strings"
"text/template"
)
var topicPattern = regexp.MustCompile(`^[-_A-Za-z0-9]{1,64}$`)
func Secret(value string) (string, error) {
if strings.HasPrefix(value, "env:") {
v := os.Getenv(strings.TrimPrefix(value, "env:"))
if v == "" {
return "", fmt.Errorf("required secret environment variable is missing")
}
return v, nil
}
return value, nil
}
// Live is opt-in: omitted fields in new or existing configurations never send.
type OutboundConfig struct {
Topic string `json:"topic,omitempty"`
SMTPHost string `json:"smtp_host,omitempty"`
SMTPPort int `json:"smtp_port,omitempty"`
TLSMode string `json:"tls_mode,omitempty"`
Username string `json:"username,omitempty"`
Password string `json:"password,omitempty"`
From string `json:"from,omitempty"`
To []string `json:"to,omitempty"`
ID string `json:"id"`
Name string `json:"name"`
Provider string `json:"provider"`
URL string `json:"url"`
BearerToken string `json:"bearer_token,omitempty"`
TimeoutS int `json:"timeout_seconds"`
Live bool `json:"live"`
}
func IsOutbound(kind string) bool {
switch kind {
case "discord", "webhook", "smtp", "ntfy", "gotify":
return true
}
return false
}
func ValidateOutbound(o OutboundConfig) error {
if strings.TrimSpace(o.ID) == "" {
return fmt.Errorf("Ausgang: id fehlt")
}
if !IsOutbound(o.Provider) {
return fmt.Errorf("Ausgang %q: unbekannter Provider", o.ID)
}
if o.TimeoutS < 0 || o.TimeoutS > 300 {
return fmt.Errorf("Ausgang %q: Timeout muss 0..300 sein", o.ID)
}
if o.Provider == "smtp" {
if o.SMTPHost == "" || strings.ContainsAny(o.SMTPHost, "\r\n /:") && net.ParseIP(o.SMTPHost) == nil || o.SMTPPort < 1 || o.SMTPPort > 65535 {
return fmt.Errorf("Ausgang %q: SMTP-Host/Port ungültig", o.ID)
}
if o.TLSMode != "tls" && o.TLSMode != "starttls" {
return fmt.Errorf("SMTP benötigt tls oder starttls")
}
if len(o.To) == 0 {
return fmt.Errorf("SMTP benötigt Empfänger")
}
for _, address := range append([]string{o.From}, o.To...) {
if strings.ContainsAny(address, "\r\n") {
return fmt.Errorf("ungültige Mailadresse")
}
a, err := mail.ParseAddress(address)
if err != nil || a.Address != address {
return fmt.Errorf("SMTP benötigt reine Mailadressen ohne Anzeigenamen")
}
}
return nil
}
if o.Provider == "ntfy" && !topicPattern.MatchString(o.Topic) {
return fmt.Errorf("ntfy benötigt Topic mit 1..64 Buchstaben, Ziffern, _ oder -")
}
u, err := url.Parse(o.URL)
if err != nil || u.Hostname() == "" || (u.Scheme != "http" && u.Scheme != "https") || u.User != nil || u.Fragment != "" {
return fmt.Errorf("Ausgang %q: gültige HTTP(S)-URL ohne Benutzerinfo oder Fragment erforderlich", o.ID)
}
if o.Provider == "discord" {
parts := strings.Split(strings.Trim(u.Path, "/"), "/")
if u.Scheme != "https" || u.Host != "discord.com" || len(parts) != 4 || parts[0] != "api" || parts[1] != "webhooks" || parts[2] == "" || parts[3] == "" {
return fmt.Errorf("Ausgang %q: Discord-URL muss https://discord.com/api/webhooks/ID/TOKEN entsprechen", o.ID)
}
if o.BearerToken != "" {
return fmt.Errorf("Ausgang %q: Discord benötigt keinen Bearer-Token", o.ID)
}
}
if strings.ContainsAny(o.BearerToken, "\r\n") {
return fmt.Errorf("Ausgang %q: ungültiger Bearer-Token", o.ID)
}
if o.TimeoutS < 0 || o.TimeoutS > 300 {
return fmt.Errorf("Ausgang %q: timeout_seconds muss 0..300 sein (0 = 15 Sekunden)", o.ID)
}
return nil
}
func validateOutbounds(c Config) error {
outputs := map[string]OutboundConfig{}
for _, o := range c.Outbounds {
if err := ValidateOutbound(o); err != nil {
return err
}
if _, exists := outputs[o.ID]; exists {
return fmt.Errorf("Ausgang id %q ist doppelt", o.ID)
}
outputs[o.ID] = o
}
for _, m := range c.Mappings {
target := strings.ToLower(m.Target)
if IsOutbound(target) {
o, ok := outputs[m.OutboundID]
if !ok || o.Provider != target {
return fmt.Errorf("mapping %q: passender Ausgang für %s erforderlich", m.ID, target)
}
} else if m.OutboundID != "" {
return fmt.Errorf("mapping %q: Divera247 verwendet keinen outbound_id", m.ID)
}
for field, value := range map[string]string{"title_template": m.TitleTemplate, "text_template": m.TextTemplate, "address_template": m.AddressTemplate} {
if _, err := template.New(field).Parse(value); err != nil {
return fmt.Errorf("mapping %q: ungültiges %s: %w", m.ID, field, err)
}
}
}
return nil
}
+94
View File
@@ -0,0 +1,94 @@
package config
import (
"encoding/json"
"os"
"testing"
)
func TestValidateOutbounds(t *testing.T) {
valid := OutboundConfig{ID: "ops", Provider: "discord", URL: "https://discord.com/api/webhooks/123/secret"}
for _, tc := range []struct {
name string
change func(*Config)
}{
{"duplicate", func(c *Config) { c.Outbounds = append(c.Outbounds, valid) }},
{"missing destination", func(c *Config) { c.Mappings[0].OutboundID = "missing" }},
{"mismatched provider", func(c *Config) { c.Mappings[0].Target = "webhook" }},
{"unsafe scheme", func(c *Config) { c.Outbounds[0].URL = "file:///secret" }},
{"discord host", func(c *Config) { c.Outbounds[0].URL = "https://evil.example/api/webhooks/123/secret" }},
{"discord credentials", func(c *Config) { c.Outbounds[0].BearerToken = "secret" }},
{"timeout", func(c *Config) { c.Outbounds[0].TimeoutS = -1 }},
{"template", func(c *Config) { c.Mappings[0].TextTemplate = "{{invalid" }},
{"divera reference", func(c *Config) { c.Mappings[0].Target = "alarm" }},
} {
t.Run(tc.name, func(t *testing.T) {
c := Default()
c.Outbounds = []OutboundConfig{valid}
c.Mappings[0].Target = "discord"
c.Mappings[0].OutboundID = "ops"
if err := Validate(c); err != nil {
t.Fatal(err)
}
tc.change(&c)
if Validate(c) == nil {
t.Fatal("expected validation error")
}
})
}
}
func TestExistingConfigurationCompatibility(t *testing.T) {
c := Default()
if err := Validate(c); err != nil {
t.Fatal(err)
}
var o OutboundConfig
if err := json.Unmarshal([]byte(`{"id":"test","provider":"webhook","url":"https://example.invalid"}`), &o); err != nil {
t.Fatal(err)
}
if o.Live {
t.Fatal("omitted live flag must be dry-run")
}
}
func TestExampleConfig(t *testing.T) {
b, err := os.ReadFile("../../config.example.json")
if err != nil {
t.Fatal(err)
}
var c Config
if err := json.Unmarshal(b, &c); err != nil {
t.Fatal(err)
}
if err := Validate(c); err != nil {
t.Fatal(err)
}
for _, o := range c.Outbounds {
if o.Live {
t.Fatal("example must use dry-run")
}
}
}
func TestMailAndPushConfigurationValidation(t *testing.T) {
for _, o := range []OutboundConfig{
{ID: "smtp", Provider: "smtp", SMTPHost: "smtp.example.org", SMTPPort: 587, TLSMode: "plain", From: "sender@example.org", To: []string{"to@example.org"}},
{ID: "smtp", Provider: "smtp", SMTPHost: "smtp.example.org", SMTPPort: 587, TLSMode: "starttls", From: "sender@example.org\r\nBcc: bad@example.org", To: []string{"to@example.org"}},
{ID: "ntfy", Provider: "ntfy", URL: "https://ntfy.sh", Topic: "bad/topic"},
} {
if ValidateOutbound(o) == nil {
t.Fatalf("invalid destination accepted: %s", o.ID)
}
}
c := Default()
c.Ingress.Mail = []MailIngress{{ID: "mail", Enabled: true, Address: "host:993", Username: "user", Password: "env:PASS", Channel: "ops", PollSeconds: 1}}
if Validate(c) == nil {
t.Fatal("invalid polling interval")
}
c.Ingress.Mail = nil
c.Ingress.Discord = DiscordIngress{Enabled: true}
if Validate(c) == nil {
t.Fatal("incomplete Discord credentials")
}
}
+208
View File
@@ -0,0 +1,208 @@
package divera
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"mime/multipart"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/outbound"
)
type Client struct{ cfg func() config.DiveraConfig }
type Response struct {
RetryAfter time.Duration
StatusCode int
ContentType string
Body []byte
}
func New(cfg func() config.DiveraConfig) *Client { return &Client{cfg: cfg} }
func (c *Client) Request(ctx context.Context, method, path string, query url.Values, body any) (Response, error) {
cfg := c.cfg()
if cfg.DryRun && method != http.MethodGet {
b, _ := json.Marshal(map[string]any{"dry_run": true, "method": method, "path": path, "body": body})
return Response{StatusCode: 200, ContentType: "application/json", Body: b}, nil
}
if cfg.AccessKey == "" {
return Response{}, fmt.Errorf("Divera247 access_key is empty")
}
key, err := config.Secret(cfg.AccessKey)
if err != nil {
return Response{}, err
}
cfg.AccessKey = key
base := strings.TrimRight(cfg.BaseURL, "/")
u, err := url.Parse(base + path)
if err != nil {
return Response{}, err
}
q := u.Query()
for k, vals := range query {
for _, v := range vals {
q.Add(k, v)
}
}
q.Set("accesskey", cfg.AccessKey)
// The v2 pull API requires the configured UCR. v3 uses the accesskey and
// clusterId instead; sending an unrelated ucr parameter there can lead to
// confusing responses, so only inject it for v2 calls.
if strings.HasPrefix(path, "/api/v2/") && cfg.UCR != 0 && q.Get("ucr") == "" {
q.Set("ucr", strconv.FormatInt(cfg.UCR, 10))
}
u.RawQuery = q.Encode()
var r io.Reader
if body != nil {
b, err := json.Marshal(body)
if err != nil {
return Response{}, err
}
r = bytes.NewReader(b)
}
req, err := http.NewRequestWithContext(ctx, method, u.String(), r)
if err != nil {
return Response{}, err
}
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
to := time.Duration(cfg.TimeoutS) * time.Second
if to <= 0 {
to = 15 * time.Second
}
resp, err := (&http.Client{Timeout: to, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}).Do(req)
if err != nil {
return Response{}, fmt.Errorf("Divera247 HTTP request failed (connection or timeout)")
}
defer resp.Body.Close()
b, err := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
if err != nil {
return Response{}, err
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return Response{StatusCode: resp.StatusCode, RetryAfter: outbound.ParseRetryAfter(resp.Header.Get("Retry-After"), time.Now()), ContentType: resp.Header.Get("Content-Type")}, fmt.Errorf("Divera247 API HTTP %d", resp.StatusCode)
}
return Response{StatusCode: resp.StatusCode, ContentType: resp.Header.Get("Content-Type"), Body: b}, nil
}
func (c *Client) List(ctx context.Context, kind string) (Response, error) {
return c.Request(ctx, http.MethodGet, "/api/v2/"+kind, nil, nil)
}
func (c *Client) Create(ctx context.Context, kind string, body any) (Response, error) {
return c.Request(ctx, http.MethodPost, "/api/v2/"+kind, nil, body)
}
func (c *Client) Get(ctx context.Context, kind string, id int64) (Response, error) {
return c.Request(ctx, http.MethodGet, fmt.Sprintf("/api/v2/%s/%d", kind, id), nil, nil)
}
func (c *Client) Update(ctx context.Context, kind string, id int64, body any) (Response, error) {
return c.Request(ctx, http.MethodPut, fmt.Sprintf("/api/v2/%s/%d", kind, id), nil, body)
}
func (c *Client) Delete(ctx context.Context, kind string, id int64) (Response, error) {
return c.Request(ctx, http.MethodDelete, fmt.Sprintf("/api/v2/%s/%d", kind, id), nil, nil)
}
func (c *Client) Archive(ctx context.Context, kind string, id int64) (Response, error) {
return c.Request(ctx, http.MethodPost, fmt.Sprintf("/api/v2/%s/archive/%d", kind, id), nil, nil)
}
func (c *Client) Read(ctx context.Context, kind string, id int64) (Response, error) {
return c.Request(ctx, http.MethodPost, fmt.Sprintf("/api/v2/%s/read/%d", kind, id), nil, nil)
}
func (c *Client) Reach(ctx context.Context, kind string, id int64) (Response, error) {
return c.Request(ctx, http.MethodGet, fmt.Sprintf("/api/v2/%s/reach/%d", kind, id), nil, nil)
}
func (c *Client) ResetResponses(ctx context.Context, kind string, id int64) (Response, error) {
return c.Request(ctx, http.MethodDelete, fmt.Sprintf("/api/v2/%s/reset-responses/%d", kind, id), nil, nil)
}
func (c *Client) Confirm(ctx context.Context, kind string, id int64, body any) (Response, error) {
return c.Request(ctx, http.MethodPost, fmt.Sprintf("/api/v2/%s/confirm/%d", kind, id), nil, body)
}
func (c *Client) Download(ctx context.Context, kind string, id int64) (Response, error) {
return c.Request(ctx, http.MethodGet, fmt.Sprintf("/api/v2/%s/download/%d", kind, id), nil, nil)
}
func (c *Client) CloseAlarm(ctx context.Context, id int64, body any) (Response, error) {
return c.Request(ctx, http.MethodPost, fmt.Sprintf("/api/v2/alarms/close/%d", id), nil, body)
}
func (c *Client) AlarmList(ctx context.Context, closed *int) (Response, error) {
q := url.Values{}
if closed != nil {
q.Set("closed", strconv.Itoa(*closed))
}
return c.Request(ctx, http.MethodGet, "/api/v2/alarms/list", q, nil)
}
func (c *Client) EventICS(ctx context.Context) (Response, error) {
return c.Request(ctx, http.MethodGet, "/api/v2/events/ics", nil, nil)
}
func (c *Client) PullAll(ctx context.Context, q url.Values) (Response, error) {
return c.Request(ctx, http.MethodGet, "/api/v2/pull/all", q, nil)
}
func (c *Client) PullVehicleStatus(ctx context.Context) (Response, error) {
return c.Request(ctx, http.MethodGet, "/api/v2/pull/vehicle-status", nil, nil)
}
// ListUserClusterRelations returns the users of a Divera247 unit. The returned
// id is the User-Cluster-Relation id (UCR id), i.e. the identifier expected by
// alarm recipient filters. clusterID is required for some PRO structures.
func (c *Client) ListUserClusterRelations(ctx context.Context, clusterID *int64) (Response, error) {
q := url.Values{}
if clusterID != nil && *clusterID > 0 {
q.Set("clusterId", strconv.FormatInt(*clusterID, 10))
}
return c.Request(ctx, http.MethodGet, "/api/v3/user-cluster-relations", q, nil)
}
func (c *Client) AddAttachment(ctx context.Context, kind string, id int64, filename, title, description string, data []byte) (Response, error) {
cfg := c.cfg()
if cfg.DryRun {
return Response{StatusCode: 200, ContentType: "application/json", Body: []byte(`{"dry_run":true}`)}, nil
}
key, err := config.Secret(cfg.AccessKey)
if err != nil {
return Response{}, err
}
cfg.AccessKey = key
var buf bytes.Buffer
mw := multipart.NewWriter(&buf)
part, err := mw.CreateFormFile("Attachment[upload]", filename)
if err != nil {
return Response{}, err
}
if _, err := part.Write(data); err != nil {
return Response{}, err
}
_ = mw.WriteField("Attachment[title]", title)
_ = mw.WriteField("Attachment[description]", description)
_ = mw.Close()
u, _ := url.Parse(strings.TrimRight(cfg.BaseURL, "/") + fmt.Sprintf("/api/v2/%s/attachment/%d", kind, id))
q := u.Query()
q.Set("accesskey", cfg.AccessKey)
if cfg.UCR != 0 {
q.Set("ucr", strconv.FormatInt(cfg.UCR, 10))
}
u.RawQuery = q.Encode()
req, err := http.NewRequestWithContext(ctx, http.MethodPost, u.String(), &buf)
if err != nil {
return Response{}, err
}
req.Header.Set("Content-Type", mw.FormDataContentType())
resp, err := (&http.Client{Timeout: time.Duration(cfg.TimeoutS) * time.Second, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}).Do(req)
if err != nil {
return Response{}, fmt.Errorf("Divera247 HTTP request failed (connection or timeout)")
}
defer resp.Body.Close()
b, _ := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
out := Response{StatusCode: resp.StatusCode, ContentType: resp.Header.Get("Content-Type"), Body: b}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
out.Body = nil
return out, fmt.Errorf("Divera247 API HTTP %d", resp.StatusCode)
}
return out, nil
}
+35
View File
@@ -0,0 +1,35 @@
package divera
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/example/notify-gateway/internal/config"
)
func TestRedirectsAndErrorsDoNotLeakCredentials(t *testing.T) {
calls := 0
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
if r.URL.Query().Get("accesskey") != "private-key" {
t.Error("missing key")
}
w.Header().Set("Location", "/other")
w.Header().Set("Retry-After", "90")
w.WriteHeader(307)
w.Write([]byte("private-key"))
}))
defer srv.Close()
t.Setenv("TEST_DIVERA_KEY", "private-key")
c := New(func() config.DiveraConfig {
return config.DiveraConfig{BaseURL: srv.URL, AccessKey: "env:TEST_DIVERA_KEY", TimeoutS: 2}
})
resp, err := c.Create(context.Background(), "news", map[string]any{})
if err == nil || calls != 1 || strings.Contains(err.Error(), "private-key") || len(resp.Body) != 0 || resp.RetryAfter != 90*time.Second {
t.Fatalf("response=%+v error=%v calls=%d", resp, err, calls)
}
}
+195
View File
@@ -0,0 +1,195 @@
package gateway
import (
"context"
"encoding/json"
"errors"
"fmt"
"regexp"
"strings"
"text/template"
"unicode/utf16"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/divera"
"github.com/example/notify-gateway/internal/model"
"github.com/example/notify-gateway/internal/outbound"
)
type Dispatcher struct {
store *config.Store
divera *divera.Client
}
func New(store *config.Store, client *divera.Client) *Dispatcher {
return &Dispatcher{store: store, divera: client}
}
func (d *Dispatcher) Dispatch(ctx context.Context, msg model.InboundMessage) ([]model.DeliveryResult, error) {
cfg := d.store.Get()
var results []model.DeliveryResult
matched := false
var failures []error
for _, m := range cfg.Mappings {
if !m.Enabled || !matches(m, msg) {
continue
}
matched = true
r := model.DeliveryResult{MappingID: m.ID, MappingName: m.Name, Target: m.Target, OutboundID: m.OutboundID}
payload, kind, err := buildPayload(m, msg)
if err == nil {
if config.IsOutbound(kind) {
var destination *config.OutboundConfig
for _, o := range cfg.Outbounds {
if o.ID == m.OutboundID && o.Provider == kind {
destination = &o
break
}
}
if destination == nil {
err = fmt.Errorf("outbound destination missing or provider mismatch")
} else {
var resp outbound.Response
resp, err = outbound.Send(ctx, *destination, payload)
r.StatusCode, r.Response = resp.StatusCode, resp.Body
}
} else {
var resp divera.Response
resp, err = d.divera.Create(ctx, kind, payload)
r.StatusCode, r.Response = resp.StatusCode, string(resp.Body)
}
}
if err != nil {
r.Error = err.Error()
failures = append(failures, fmt.Errorf("mapping %s: %w", m.ID, err))
}
results = append(results, r)
}
if !matched {
return results, fmt.Errorf("no mapping matched source=%s channel=%s", msg.Source, msg.Channel)
}
return results, errors.Join(failures...)
}
func matches(m config.Mapping, msg model.InboundMessage) bool {
if m.Source != "" && m.Source != "any" && m.Source != msg.Source {
return false
}
if m.MinPriority != 0 && msg.Priority < m.MinPriority {
return false
}
for _, condition := range [][2]string{{m.ChannelRegex, msg.Channel}, {m.TitleRegex, msg.Title}, {m.MessageRegex, msg.Message}} {
pat, value := condition[0], condition[1]
if pat == "" {
continue
}
ok, err := regexp.MatchString(pat, value)
if err != nil || !ok {
return false
}
}
return true
}
func render(s string, msg model.InboundMessage) (string, error) {
if s == "" {
return "", nil
}
t, err := template.New("m").Option("missingkey=zero").Parse(s)
if err != nil {
return "", err
}
var b strings.Builder
if err := t.Execute(&b, msg); err != nil {
return "", err
}
return b.String(), nil
}
func buildPayload(m config.Mapping, msg model.InboundMessage) (map[string]any, string, error) {
title, err := render(m.TitleTemplate, msg)
if err != nil {
return nil, "", err
}
text, err := render(m.TextTemplate, msg)
if err != nil {
return nil, "", err
}
address, err := render(m.AddressTemplate, msg)
if err != nil {
return nil, "", err
}
if address == "" {
address = msg.Address
}
switch strings.ToLower(m.Target) {
case "smtp", "ntfy", "gotify":
if strings.TrimSpace(text) == "" {
return nil, "", fmt.Errorf("message must not be empty")
}
return map[string]any{"title": title, "message": text, "priority": msg.Priority}, strings.ToLower(m.Target), nil
case "discord":
content := strings.TrimSpace(title + "\n" + text)
if content == "" {
return nil, "", fmt.Errorf("Discord content must not be empty")
}
if len(utf16.Encode([]rune(content))) > 2000 {
return nil, "", fmt.Errorf("Discord content exceeds 2000 characters; shorten title/text templates")
}
return map[string]any{"content": content, "allowed_mentions": map[string]any{"parse": []string{}}}, "discord", nil
case "webhook":
return map[string]any{"source": msg.Source, "channel": msg.Channel, "title": title, "message": text, "address": address, "priority": msg.Priority, "tags": msg.Tags, "received_at": msg.ReceivedAt}, "webhook", nil
}
obj := map[string]any{"title": title, "text": text, "notification_type": m.NotificationType, "send_push": m.SendPush, "send_sms": m.SendSMS, "send_call": m.SendCall, "send_mail": m.SendMail, "send_pager": m.SendPager, "private_mode": m.PrivateMode}
if address != "" {
obj["address"] = address
}
if len(m.ClusterRoutes) > 0 {
routes := make(map[string]any, len(m.ClusterRoutes))
for id, notificationType := range m.ClusterRoutes {
routes[id] = map[string]any{"notification_type": notificationType}
}
obj["cluster"] = routes
} else if len(m.Clusters) > 0 {
obj["cluster"] = m.Clusters
}
if len(m.Groups) > 0 {
obj["group"] = m.Groups
}
if len(m.Users) > 0 {
obj["user_cluster_relation"] = m.Users
}
if len(m.Vehicles) > 0 {
obj["vehicle"] = m.Vehicles
}
for k, v := range m.Extra {
obj[k] = v
}
var root, kind string
switch strings.ToLower(m.Target) {
case "alarm", "alarms":
root, kind = "Alarm", "alarms"
case "news", "message", "mitteilung":
root, kind = "News", "news"
case "event", "termin":
root, kind = "Event", "events"
default:
return nil, "", fmt.Errorf("unsupported target %q", m.Target)
}
return map[string]any{root: obj}, kind, nil
}
func Pretty(v any) string { b, _ := json.MarshalIndent(v, "", " "); return string(b) }
// Preview evaluates a single mapping against a sample message without sending anything.
func Preview(m config.Mapping, msg model.InboundMessage) (map[string]any, string, bool, error) {
matched := matches(m, msg)
if !matched {
return nil, "", false, nil
}
payload, kind, err := buildPayload(m, msg)
if err != nil {
return nil, "", true, err
}
return payload, kind, true, nil
}
+102
View File
@@ -0,0 +1,102 @@
package gateway
import (
"context"
"encoding/json"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/divera"
"github.com/example/notify-gateway/internal/model"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
)
func TestRepeatedRegexMustMatchEveryField(t *testing.T) {
m := config.Mapping{ChannelRegex: "^alarm$", TitleRegex: "^alarm$"}
if matches(m, model.InboundMessage{Channel: "alarm", Title: "other"}) {
t.Fatal("same pattern must be checked for both fields")
}
}
func TestProviderPayloads(t *testing.T) {
m := config.Mapping{Target: "discord", TitleTemplate: "{{.Title}}", TextTemplate: "{{.Message}}"}
p, kind, err := buildPayload(m, model.InboundMessage{Title: "Test", Message: "@everyone Grüße"})
if err != nil || kind != "discord" || p["content"] != "Test\n@everyone Grüße" {
t.Fatalf("%v %s %v", p, kind, err)
}
b, _ := json.Marshal(p)
if !strings.Contains(string(b), `"allowed_mentions":{"parse":[]}`) {
t.Fatalf("mentions not disabled: %s", b)
}
for _, body := range []string{"", strings.Repeat("a", 2001), strings.Repeat("😀", 1001)} {
if _, _, err := buildPayload(m, model.InboundMessage{Message: body}); err == nil {
t.Fatal("expected empty/oversized payload to fail")
}
}
if _, _, err := buildPayload(m, model.InboundMessage{Message: strings.Repeat("a", 2000)}); err != nil {
t.Fatal(err)
}
m.Target = "webhook"
p, _, err = buildPayload(m, model.InboundMessage{Source: "ntfy", Channel: "ops", Title: "hello", Message: "world", Raw: map[string]any{"secret": "private"}})
if err != nil || p["source"] != "ntfy" || p["message"] != "world" {
t.Fatalf("%v %v", p, err)
}
if _, ok := p["raw"]; ok {
t.Fatal("raw input must not be forwarded automatically")
}
}
func TestDispatchContinuesAfterFailedDestination(t *testing.T) {
store, err := config.Open(filepath.Join(t.TempDir(), "config.json"))
if err != nil {
t.Fatal(err)
}
failing := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(429) }))
defer failing.Close()
c := store.Get()
c.Outbounds = []config.OutboundConfig{{ID: "bad", Provider: "webhook", URL: failing.URL, Live: true}, {ID: "dry", Provider: "discord", URL: "https://discord.com/api/webhooks/123/secret"}}
c.Mappings = []config.Mapping{
{ID: "first", Enabled: true, Target: "webhook", OutboundID: "bad", TextTemplate: "{{.Message}}"},
{ID: "second", Enabled: true, Target: "discord", OutboundID: "dry", TextTemplate: "{{.Message}}"},
{ID: "legacy", Enabled: true, Target: "alarm", TextTemplate: "{{.Message}}"},
}
if err := store.Replace(c); err != nil {
t.Fatal(err)
}
d := New(store, divera.New(func() config.DiveraConfig { return store.Get().Divera }))
results, err := d.Dispatch(context.Background(), model.InboundMessage{Message: "test"})
if err == nil || len(results) != 3 {
t.Fatalf("results=%+v error=%v", results, err)
}
if results[0].StatusCode != 429 || results[0].Error == "" || results[1].Error != "" || results[1].StatusCode != 200 || results[2].StatusCode != 200 {
t.Fatalf("results=%+v", results)
}
}
func TestMatches(t *testing.T) {
m := config.Mapping{Enabled: true, Source: "ntfy", ChannelRegex: "^alarm-.*$", MessageRegex: "Feuer", MinPriority: 3}
msg := model.InboundMessage{Source: "ntfy", Channel: "alarm-1", Message: "Feuer im Gebäude", Priority: 4}
if !matches(m, msg) {
t.Fatal("expected mapping to match")
}
msg.Priority = 2
if matches(m, msg) {
t.Fatal("priority should not match")
}
}
func TestBuildPayload(t *testing.T) {
m := config.Mapping{Target: "alarm", TitleTemplate: "{{.Title}}", TextTemplate: "{{.Message}}", NotificationType: 3, Groups: []int64{12}, SendPush: true}
p, kind, err := buildPayload(m, model.InboundMessage{Title: "B3", Message: "Rauch"})
if err != nil {
t.Fatal(err)
}
if kind != "alarms" {
t.Fatalf("kind=%s", kind)
}
a := p["Alarm"].(map[string]any)
if a["title"] != "B3" {
t.Fatalf("title=%v", a["title"])
}
}
+188
View File
@@ -0,0 +1,188 @@
package gateway
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"log"
"math/rand/v2"
"strings"
"time"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/divera"
"github.com/example/notify-gateway/internal/model"
"github.com/example/notify-gateway/internal/outbound"
"github.com/example/notify-gateway/internal/outbox"
)
type Delivery struct {
Kind string `json:"kind"`
Payload map[string]any `json:"payload"`
Outbound *config.OutboundConfig `json:"outbound,omitempty"`
Divera *config.DiveraConfig `json:"divera,omitempty"`
}
type Queue struct {
Store *outbox.Store
Dispatcher *Dispatcher
Logger *log.Logger
}
type InputError struct{ Err error }
func (e *InputError) Error() string { return e.Err.Error() }
func (e *InputError) Unwrap() error { return e.Err }
func (q *Queue) Accept(ctx context.Context, msg model.InboundMessage, key string) (outbox.Receipt, error) {
if len(key) > 256 {
return outbox.Receipt{}, &InputError{errors.New("Idempotency-Key exceeds 256 bytes")}
}
original := msg
original.ReceivedAt = time.Time{}
b, err := json.Marshal(original)
if err != nil {
return outbox.Receipt{}, &InputError{err}
}
hash := sha256.Sum256(b)
digest := hex.EncodeToString(hash[:])
scopeBytes, _ := json.Marshal([]string{msg.Source, msg.Channel})
scope := string(scopeBytes)
if key == "" {
key = outbox.ID()
} else {
r, err := q.Store.Lookup(ctx, scope, key, digest)
if err != nil || r.ID != "" {
return r, err
}
}
cfg := q.Dispatcher.store.Get()
var jobs []outbox.Job
for _, m := range cfg.Mappings {
if !m.Enabled || !matches(m, msg) {
continue
}
payload, kind, err := buildPayload(m, msg)
if err != nil {
return outbox.Receipt{}, &InputError{fmt.Errorf("mapping %s: %w", m.ID, err)}
}
delivery := Delivery{Kind: kind, Payload: payload}
if config.IsOutbound(kind) {
for _, o := range cfg.Outbounds {
if o.ID == m.OutboundID && o.Provider == kind {
delivery.Outbound = &o
break
}
}
if delivery.Outbound == nil {
return outbox.Receipt{}, &InputError{errors.New("outbound destination missing")}
}
} else {
delivery.Divera = &cfg.Divera
}
data, err := json.Marshal(delivery)
if err != nil {
return outbox.Receipt{}, err
}
jobs = append(jobs, outbox.Job{MappingID: m.ID, Target: m.Target, Data: data})
}
if len(jobs) == 0 {
return outbox.Receipt{}, &InputError{errors.New("no mapping matched")}
}
return q.Store.Enqueue(ctx, scope, key, digest, jobs)
}
func retryDelay(attempt int, after time.Duration) time.Duration {
d := 5 * time.Second * time.Duration(1<<min(max(attempt-1, 0), 8))
d += time.Duration(rand.Int64N(int64(d/5) + 1))
if after > d {
d = after
}
return min(d, 24*time.Hour)
}
// ProcessOne is also used by integration tests. A lease survives process crashes.
func (q *Queue) ProcessOne(ctx context.Context) (bool, error) {
j, err := q.Store.Claim(ctx, time.Now())
if err != nil || j == nil {
return false, err
}
var d Delivery
var resp outbound.Response
err = json.Unmarshal(j.Data, &d)
dry := false
retry := false
if err == nil {
if d.Outbound != nil {
dry = !d.Outbound.Live
resp, err = outbound.SendWithID(ctx, *d.Outbound, d.Payload, j.ID)
retry = resp.Retryable
} else if d.Divera != nil {
dry = d.Divera.DryRun
c := divera.New(func() config.DiveraConfig { return *d.Divera })
r, e := c.Create(ctx, d.Kind, d.Payload)
resp.StatusCode = r.StatusCode
resp.RetryAfter = r.RetryAfter
err = e
retry = outbound.RetryableHTTP(r.StatusCode)
} else {
err = errors.New("invalid persisted destination")
}
}
state := "succeeded"
message := ""
if dry {
state = "dry_run"
}
if err != nil {
// Never persist provider error strings: URLs, SMTP responses and templates may contain secrets.
message = "Zustellung fehlgeschlagen"
if resp.StatusCode != 0 {
message = fmt.Sprintf("Provider-Status %d", resp.StatusCode)
}
state = "dead"
if retry && j.Attempts < 8 {
state = "pending"
}
}
next := time.Now().Add(retryDelay(j.Attempts, resp.RetryAfter))
// Persist the outcome even if shutdown canceled the network request.
finishCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := q.Store.Finish(finishCtx, *j, state, resp.StatusCode, message, next); err != nil {
return true, err
}
if q.Logger != nil {
q.Logger.Printf("delivery id=%s state=%s attempt=%d status=%d", j.ID, state, j.Attempts, resp.StatusCode)
}
return true, nil
}
func (q *Queue) Run(ctx context.Context) {
ticker := time.NewTicker(time.Second)
defer ticker.Stop()
for {
if ctx.Err() != nil {
return
}
worked, err := q.ProcessOne(ctx)
if err != nil && q.Logger != nil {
q.Logger.Print("outbox worker: database operation failed")
}
if worked && err == nil {
continue
}
select {
case <-ctx.Done():
return
case <-ticker.C:
}
}
}
// Message keys should never contain transport credentials.
func ScopedKey(parts ...string) string {
b, _ := json.Marshal(parts)
sum := sha256.Sum256(b)
return strings.ToLower(hex.EncodeToString(sum[:]))
}
+166
View File
@@ -0,0 +1,166 @@
package gateway
import (
"context"
"database/sql"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"time"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/model"
"github.com/example/notify-gateway/internal/outbox"
)
func TestQueueRetriesOnlyFailedDestinationAndFreezesDryRun(t *testing.T) {
ctx := context.Background()
dir := t.TempDir()
store, err := config.Open(filepath.Join(dir, "config.json"))
if err != nil {
t.Fatal(err)
}
db, err := outbox.Open(filepath.Join(dir, "outbox.db"))
if err != nil {
t.Fatal(err)
}
defer db.Close()
good, bad := 0, 0
fail := true
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Header.Get("Idempotency-Key") == "" {
t.Error("missing delivery id")
}
if r.URL.Path == "/good" {
good++
w.WriteHeader(204)
} else {
bad++
if fail {
w.WriteHeader(400)
} else {
w.WriteHeader(200)
}
}
}))
defer srv.Close()
c := store.Get()
c.Outbounds = []config.OutboundConfig{{ID: "good", Provider: "webhook", URL: srv.URL + "/good", Live: true}, {ID: "bad", Provider: "webhook", URL: srv.URL + "/bad", Live: true}, {ID: "dry", Provider: "webhook", URL: srv.URL + "/good"}}
c.Mappings = nil
for _, id := range []string{"good", "bad", "dry"} {
c.Mappings = append(c.Mappings, config.Mapping{ID: id, Enabled: true, Target: "webhook", OutboundID: id, TextTemplate: "{{.Message}}"})
}
if err := store.Replace(c); err != nil {
t.Fatal(err)
}
q := &Queue{Store: db, Dispatcher: New(store, nil)}
msg := model.InboundMessage{Source: "webhook", Channel: "ops", Message: "test", ReceivedAt: time.Now()}
first, err := q.Accept(ctx, msg, "request-1")
if err != nil {
t.Fatal(err)
}
c.Outbounds[2].Live = true
store.Replace(c)
for i := 0; i < 3; i++ {
if ok, err := q.ProcessOne(ctx); err != nil || !ok {
t.Fatalf("process %v %v", ok, err)
}
}
if good != 1 || bad != 1 {
t.Fatalf("calls good=%d bad=%d", good, bad)
}
msg.ReceivedAt = time.Now().Add(time.Hour)
duplicate, err := q.Accept(ctx, msg, "request-1")
if err != nil || !duplicate.Duplicate || first.ID != duplicate.ID {
t.Fatalf("receipt %+v %v", duplicate, err)
}
rows, _ := db.List(ctx, 50, 0)
for _, j := range rows {
if j.MappingID == "bad" {
if j.State != "dead" {
t.Fatal(j.State)
}
if err := db.Retry(ctx, j.ID); err != nil {
t.Fatal(err)
}
}
if j.MappingID == "dry" && j.State != "dry_run" {
t.Fatal("dry-run turned live")
}
}
fail = false
if ok, err := q.ProcessOne(ctx); !ok || err != nil {
t.Fatal(err)
}
if good != 1 || bad != 2 {
t.Fatalf("successful destination repeated: good=%d bad=%d", good, bad)
}
}
func TestQueueBackoffAndAtomicInvalidMapping(t *testing.T) {
ctx := context.Background()
dir := t.TempDir()
store, _ := config.Open(filepath.Join(dir, "config.json"))
db, err := outbox.Open(filepath.Join(dir, "outbox.db"))
if err != nil {
t.Fatal(err)
}
defer db.Close()
calls := 0
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
w.Header().Set("Retry-After", "120")
w.WriteHeader(429)
}))
defer srv.Close()
c := store.Get()
c.Outbounds = []config.OutboundConfig{{ID: "rate", Provider: "webhook", URL: srv.URL, Live: true}}
c.Mappings = []config.Mapping{{ID: "rate", Enabled: true, Target: "webhook", OutboundID: "rate", TextTemplate: "{{.Message}}"}}
if err := store.Replace(c); err != nil {
t.Fatal(err)
}
q := &Queue{Store: db, Dispatcher: New(store, nil)}
if _, err := q.Accept(ctx, model.InboundMessage{Message: "one"}, "one"); err != nil {
t.Fatal(err)
}
if _, err := q.ProcessOne(ctx); err != nil {
t.Fatal(err)
}
rows, _ := db.List(ctx, 50, 0)
if rows[0].State != "pending" || rows[0].Next < time.Now().Add(119*time.Second).Unix() {
t.Fatalf("backoff: %+v", rows)
}
if worked, err := q.ProcessOne(ctx); err != nil || worked || calls != 1 {
t.Fatalf("early retry %v %v %d", worked, err, calls)
}
c.Mappings = append(c.Mappings, config.Mapping{ID: "invalid", Enabled: true, Target: "discord", OutboundID: "discord", TextTemplate: ""})
c.Outbounds = append(c.Outbounds, config.OutboundConfig{ID: "discord", Provider: "discord", URL: "https://discord.com/api/webhooks/123/token"})
store.Replace(c)
if _, err := q.Accept(ctx, model.InboundMessage{Message: "new"}, "new"); err == nil {
t.Fatal("expected invalid empty Discord payload")
}
rows, _ = db.List(ctx, 50, 0)
if len(rows) != 1 {
t.Fatal("partial message was queued")
}
// Advance the due time in the isolated database to exercise the retry cap
// without waiting for the backoff clock in this integration test.
clockDB, err := sql.Open("sqlite", filepath.Join(dir, "outbox.db"))
if err != nil {
t.Fatal(err)
}
defer clockDB.Close()
for i := 1; i < 8; i++ {
if _, err := clockDB.Exec("UPDATE deliveries SET next=0 WHERE state='pending'"); err != nil {
t.Fatal(err)
}
if _, err := q.ProcessOne(ctx); err != nil {
t.Fatal(err)
}
}
rows, _ = db.List(ctx, 50, 0)
if rows[0].State != "dead" || rows[0].Attempts != 8 || calls != 8 {
t.Fatalf("retry limit: %+v calls=%d", rows, calls)
}
}
+673
View File
@@ -0,0 +1,673 @@
package httpserver
import (
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"sort"
"strconv"
"strings"
"sync"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/gateway"
"github.com/example/notify-gateway/internal/model"
)
type editableConfig struct {
Version int `json:"version"`
Server editableServerConfig `json:"server"`
Divera config.DiveraConfig `json:"divera"`
Ingress config.IngressConfig `json:"ingress"`
Mappings []config.Mapping `json:"mappings"`
Outbounds []config.OutboundConfig `json:"outbounds,omitempty"`
}
type editableServerConfig struct {
Listen string `json:"listen"`
AdminUsername string `json:"admin_username"`
}
func makeEditable(c config.Config) editableConfig {
return editableConfig{
Version: c.Version,
Server: editableServerConfig{Listen: c.Server.Listen, AdminUsername: c.Server.AdminUsername},
Divera: c.Divera,
Ingress: c.Ingress,
Mappings: c.Mappings,
Outbounds: c.Outbounds,
}
}
func (s *Server) apiConfig(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet:
writeJSON(w, http.StatusOK, makeEditable(s.store.Get()))
case http.MethodPut:
var in editableConfig
dec := json.NewDecoder(io.LimitReader(r.Body, 2<<20))
dec.DisallowUnknownFields()
if err := dec.Decode(&in); err != nil {
writeJSON(w, http.StatusBadRequest, map[string]any{"error": "Ungültige Konfiguration: " + err.Error()})
return
}
old := s.store.Get()
next := old
next.Version = config.CurrentVersion
next.Server.Listen = strings.TrimSpace(in.Server.Listen)
next.Server.AdminUsername = strings.TrimSpace(in.Server.AdminUsername)
next.Divera = in.Divera
next.Ingress = in.Ingress
next.Mappings = in.Mappings
next.Outbounds = in.Outbounds
if err := config.Validate(next); err != nil {
writeJSON(w, http.StatusBadRequest, map[string]any{"error": err.Error()})
return
}
if err := s.store.Replace(next); err != nil {
writeJSON(w, http.StatusInternalServerError, map[string]any{"error": err.Error()})
return
}
writeJSON(w, http.StatusOK, map[string]any{"ok": true, "config": makeEditable(s.store.Get())})
default:
w.Header().Set("Allow", "GET, PUT")
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
}
}
func (s *Server) apiDiveraCatalog(w http.ResponseWriter, r *http.Request) {
// Divera247 v2 pull/all is the primary catalog source. Besides the active
// unit it returns data.ucr, i.e. every UserClusterRelation the current
// account can switch to. We use those UCR ids to load each reachable unit
// and merge its cluster.consumer/group/vehicle catalog. This works with
// normal v2 access even when the v3 synchronization API is not enabled.
resp, err := s.divera.PullAll(r.Context(), nil)
if err != nil {
status := http.StatusBadGateway
if strings.Contains(strings.ToLower(err.Error()), "access_key") {
status = http.StatusBadRequest
}
writeJSON(w, status, map[string]any{"error": err.Error(), "status_code": resp.StatusCode, "response": string(resp.Body)})
return
}
var raw any
if err := json.Unmarshal(resp.Body, &raw); err != nil {
writeJSON(w, http.StatusBadGateway, map[string]any{"error": "Divera247-Antwort ist kein JSON: " + err.Error()})
return
}
catalog := extractCatalog(raw)
refs := extractUCRRefs(raw)
activeUCR := extractActiveUCR(raw)
for _, ref := range refs {
catalog.Units = mergeCatalogItems(catalog.Units, []catalogItem{{ID: ref.ClusterID, Label: ref.Label}})
}
if ref, ok := findUCRRef(refs, activeUCR); ok {
annotateCatalogWithUnit(&catalog, ref)
}
warnings := []string{}
v2Calls, v2Successes := 1, 1
failedV2 := []string{}
// pull/all requests are independent. Limit concurrency so a PRO account with
// many units does not create a burst of dozens of simultaneous requests.
type pullResult struct {
ref ucrRef
catalog diveraCatalog
err error
}
resultCh := make(chan pullResult, len(refs))
sem := make(chan struct{}, 6)
var wg sync.WaitGroup
for _, ref := range refs {
if ref.UCRID <= 0 || ref.UCRID == activeUCR {
continue
}
ref := ref
v2Calls++
wg.Add(1)
go func() {
defer wg.Done()
select {
case sem <- struct{}{}:
defer func() { <-sem }()
case <-r.Context().Done():
resultCh <- pullResult{ref: ref, err: r.Context().Err()}
return
}
q := url.Values{}
q.Set("ucr", strconv.FormatInt(ref.UCRID, 10))
unitResp, unitErr := s.divera.PullAll(r.Context(), q)
if unitErr != nil {
resultCh <- pullResult{ref: ref, err: unitErr}
return
}
var unitRaw any
if err := json.Unmarshal(unitResp.Body, &unitRaw); err != nil {
resultCh <- pullResult{ref: ref, err: fmt.Errorf("ungültiges JSON: %w", err)}
return
}
unitCatalog := extractCatalog(unitRaw)
annotateCatalogWithUnit(&unitCatalog, ref)
resultCh <- pullResult{ref: ref, catalog: unitCatalog}
}()
}
go func() {
wg.Wait()
close(resultCh)
}()
for result := range resultCh {
if result.err != nil {
failedV2 = append(failedV2, fmt.Sprintf("%s: %v", result.ref.Label, result.err))
continue
}
v2Successes++
mergeDiveraCatalog(&catalog, result.catalog)
}
if len(failedV2) > 0 {
examples := failedV2
if len(examples) > 3 {
examples = examples[:3]
}
warnings = append(warnings, fmt.Sprintf("Stammdaten für %d von %d Divera247-Einheiten konnten über v2 pull/all nicht geladen werden. Beispiele: %s", len(failedV2), len(refs), strings.Join(examples, " | ")))
}
personSource := "v2 pull/all · cluster.consumer"
v3Calls, v3Successes := 0, 0
// v3 is a beta/synchronization API and can legitimately be forbidden for an
// otherwise valid access key. Only use it as a fallback when v2 did not yield
// any consumers. On an explicit 403 stop immediately instead of producing one
// warning per unit.
if len(catalog.Persons) == 0 {
personSource = "keine Personen verfügbar"
probeRefs := refs
if len(probeRefs) == 0 {
probeRefs = []ucrRef{{}}
}
v3Persons := []catalogItem{}
for _, ref := range probeRefs {
var clusterID *int64
if ref.ClusterID > 0 {
id := ref.ClusterID
clusterID = &id
}
v3Calls++
userResp, userErr := s.divera.ListUserClusterRelations(r.Context(), clusterID)
if userErr != nil {
if userResp.StatusCode == http.StatusForbidden {
warnings = append(warnings, "Divera247 v3 Benutzer-Synchronisation ist für diesen Accesskey nicht freigeschaltet (HTTP 403). Personen werden ausschließlich aus v2 pull/all gelesen.")
break
}
warnings = append(warnings, "Divera247-v3-Fallback für Personen fehlgeschlagen: "+userErr.Error())
break
}
persons, parseErr := parseV3Persons(userResp.Body)
if parseErr != nil {
warnings = append(warnings, "Divera247-v3-Benutzerantwort konnte nicht gelesen werden: "+parseErr.Error())
break
}
v3Successes++
v3Persons = mergeCatalogItems(v3Persons, persons)
}
if len(v3Persons) > 0 {
catalog.Persons = v3Persons
personSource = "v3 user-cluster-relations"
}
}
sortCatalog(&catalog)
writeJSON(w, http.StatusOK, map[string]any{
"ok": true,
"catalog": catalog,
"warnings": warnings,
"diagnostics": map[string]any{
"v2_ucr_calls": v2Calls,
"v2_ucr_successes": v2Successes,
"v3_user_calls": v3Calls,
"v3_user_successes": v3Successes,
"person_count": len(catalog.Persons),
"person_source": personSource,
"available_ucr_count": len(refs),
},
})
}
type ucrRef struct {
UCRID int64
ClusterID int64
Label string
}
func extractUCRRefs(raw any) []ucrRef {
root, ok := raw.(map[string]any)
if !ok {
return nil
}
data, ok := root["data"].(map[string]any)
if !ok {
return nil
}
var refs []ucrRef
add := func(m map[string]any, implicit int64) {
ucrID, _ := positiveInt64(m["id"])
if ucrID <= 0 {
ucrID = implicit
}
clusterID, _ := positiveInt64(m["cluster_id"])
if ucrID <= 0 || clusterID <= 0 {
return
}
label := anyString(m["name"])
short := anyString(m["shortname"])
if label == "" {
label = short
} else if short != "" && short != label {
label += " (" + short + ")"
}
if label == "" {
label = fmt.Sprintf("Einheit %d", clusterID)
}
refs = append(refs, ucrRef{UCRID: ucrID, ClusterID: clusterID, Label: label})
}
switch u := data["ucr"].(type) {
case []any:
for _, row := range u {
if m, ok := row.(map[string]any); ok {
add(m, 0)
}
}
case map[string]any:
for key, row := range u {
if m, ok := row.(map[string]any); ok {
implicit, _ := strconv.ParseInt(key, 10, 64)
add(m, implicit)
}
}
}
byUCR := map[int64]ucrRef{}
for _, ref := range refs {
byUCR[ref.UCRID] = ref
}
refs = refs[:0]
for _, ref := range byUCR {
refs = append(refs, ref)
}
sort.Slice(refs, func(i, j int) bool {
li, lj := strings.ToLower(refs[i].Label), strings.ToLower(refs[j].Label)
if li == lj {
return refs[i].UCRID < refs[j].UCRID
}
return li < lj
})
return refs
}
func extractActiveUCR(raw any) int64 {
root, ok := raw.(map[string]any)
if !ok {
return 0
}
data, ok := root["data"].(map[string]any)
if !ok {
return 0
}
id, _ := positiveInt64(data["ucr_active"])
return id
}
func findUCRRef(refs []ucrRef, id int64) (ucrRef, bool) {
for _, ref := range refs {
if ref.UCRID == id {
return ref, true
}
}
return ucrRef{}, false
}
func annotateCatalogWithUnit(c *diveraCatalog, ref ucrRef) {
if ref.ClusterID <= 0 {
return
}
unit := ref.Label
if unit == "" {
unit = fmt.Sprintf("Einheit %d", ref.ClusterID)
}
suffix := fmt.Sprintf("Einheit: %s · Cluster-ID %d", unit, ref.ClusterID)
annotate := func(items []catalogItem) {
for i := range items {
if !strings.Contains(items[i].Subtitle, "Cluster-ID ") {
if items[i].Subtitle == "" {
items[i].Subtitle = suffix
} else {
items[i].Subtitle += " · " + suffix
}
}
}
}
annotate(c.Groups)
annotate(c.Persons)
annotate(c.Vehicles)
}
func mergeDiveraCatalog(dst *diveraCatalog, src diveraCatalog) {
dst.Units = mergeCatalogItems(dst.Units, src.Units)
dst.Groups = mergeCatalogItems(dst.Groups, src.Groups)
dst.Persons = mergeCatalogItems(dst.Persons, src.Persons)
dst.Vehicles = mergeCatalogItems(dst.Vehicles, src.Vehicles)
}
type catalogItem struct {
ID int64 `json:"id"`
Label string `json:"label"`
Subtitle string `json:"subtitle,omitempty"`
}
type diveraCatalog struct {
Units []catalogItem `json:"units"`
Groups []catalogItem `json:"groups"`
Persons []catalogItem `json:"persons"`
Vehicles []catalogItem `json:"vehicles"`
}
// extractCatalog tolerates different pull/all response layouts. Numeric keys in
// collections are preserved separately because in cluster/consumer Divera247 can use
// that key as UCR id while the nested object's "id" is the global user id.
func extractCatalog(raw any) diveraCatalog {
out := diveraCatalog{}
seen := map[string]map[int64]bool{
"units": {}, "groups": {}, "persons": {}, "vehicles": {},
}
var walk func(v any, path string)
walk = func(v any, path string) {
switch x := v.(type) {
case map[string]any:
kind := classifyCatalogPath(path)
if id, ok := objectIDForKind(x, kind); ok && kind != "" && !seen[kind][id] {
seen[kind][id] = true
item := catalogItem{ID: id, Label: objectLabel(x, id), Subtitle: objectSubtitle(x)}
switch kind {
case "units":
out.Units = append(out.Units, item)
case "groups":
out.Groups = append(out.Groups, item)
case "persons":
out.Persons = append(out.Persons, item)
case "vehicles":
out.Vehicles = append(out.Vehicles, item)
}
}
for k, child := range x {
p := strings.ToLower(k)
if path != "" {
p = path + "/" + p
}
// pull/all commonly returns collections as objects keyed by numeric ID.
// Preserve the key even if the child has its own "id" field: for users,
// the child id may be the global User-ID while the collection key is UCR.
if cm, ok := child.(map[string]any); ok {
if implicit, err := strconv.ParseInt(k, 10, 64); err == nil && implicit > 0 {
copyMap := make(map[string]any, len(cm)+1)
for ck, cv := range cm {
copyMap[ck] = cv
}
copyMap["__collection_key_id"] = implicit
child = copyMap
}
}
walk(child, p)
}
case []any:
for _, child := range x {
walk(child, path)
}
}
}
walk(raw, "")
sortCatalog(&out)
return out
}
func classifyCatalogPath(path string) string {
p := strings.ToLower(path)
switch {
case strings.Contains(p, "vehicle") || strings.Contains(p, "fahrzeug"):
return "vehicles"
case strings.Contains(p, "user_cluster_relation") || strings.Contains(p, "user-cluster-relation") || strings.Contains(p, "userclusterrelation") || strings.Contains(p, "consumer") || strings.Contains(p, "member") || strings.Contains(p, "person") || strings.Contains(p, "/user") || strings.HasSuffix(p, "users"):
return "persons"
case strings.Contains(p, "group") || strings.Contains(p, "gruppen"):
return "groups"
case strings.Contains(p, "cluster") || strings.Contains(p, "unit") || strings.Contains(p, "standort") || strings.Contains(p, "einheit"):
return "units"
default:
return ""
}
}
func objectIDForKind(m map[string]any, kind string) (int64, bool) {
if kind == "persons" {
// Prefer explicit UCR fields. If pull/all uses cluster/consumer keyed by UCR,
// prefer the collection key before the nested global user id.
for _, k := range []string{"user_cluster_relation_id", "userClusterRelationId", "ucr_id", "ucr"} {
if id, ok := positiveInt64(m[k]); ok {
return id, true
}
}
if id, ok := positiveInt64(m["__collection_key_id"]); ok {
return id, true
}
}
for _, k := range []string{"id", "cluster_id", "group_id", "vehicle_id", "__collection_key_id"} {
if id, ok := positiveInt64(m[k]); ok {
return id, true
}
}
return 0, false
}
func positiveInt64(v any) (int64, bool) {
switch n := v.(type) {
case int64:
return n, n > 0
case int:
return int64(n), n > 0
case float64:
return int64(n), n > 0
case json.Number:
i, err := n.Int64()
return i, err == nil && i > 0
case string:
i, err := strconv.ParseInt(strings.TrimSpace(n), 10, 64)
return i, err == nil && i > 0
default:
return 0, false
}
}
func objectLabel(m map[string]any, id int64) string {
for _, k := range []string{"title", "name", "display_name", "fullname", "full_name", "label", "shortname", "callname", "number"} {
if s := anyString(m[k]); s != "" {
return s
}
}
first := anyString(m["firstname"])
last := anyString(m["lastname"])
if strings.TrimSpace(first+" "+last) != "" {
return strings.TrimSpace(first + " " + last)
}
if nested, ok := m["user"].(map[string]any); ok {
first = anyString(nested["firstname"])
last = anyString(nested["lastname"])
if strings.TrimSpace(first+" "+last) != "" {
return strings.TrimSpace(first + " " + last)
}
if email := anyString(nested["email"]); email != "" {
return email
}
}
if email := anyString(m["email"]); email != "" {
return email
}
return fmt.Sprintf("ID %d", id)
}
func objectSubtitle(m map[string]any) string {
vals := []string{}
appendUnique := func(s string) {
s = strings.TrimSpace(s)
if s == "" {
return
}
for _, old := range vals {
if old == s {
return
}
}
vals = append(vals, s)
}
for _, k := range []string{"foreign_id", "ric", "issi", "opta", "number", "email"} {
appendUnique(anyString(m[k]))
}
if nested, ok := m["user"].(map[string]any); ok {
appendUnique(anyString(nested["email"]))
}
if clusterID, ok := positiveInt64(m["cluster_id"]); ok {
appendUnique(fmt.Sprintf("Einheit %d", clusterID))
}
return strings.Join(vals, " · ")
}
func anyString(v any) string {
switch x := v.(type) {
case string:
return strings.TrimSpace(x)
case int:
return strconv.Itoa(x)
case int64:
return strconv.FormatInt(x, 10)
case float64:
return strconv.FormatInt(int64(x), 10)
case json.Number:
return x.String()
default:
return ""
}
}
func parseV3Persons(body []byte) ([]catalogItem, error) {
var rows []struct {
ID int64 `json:"id"`
ClusterID int64 `json:"cluster_id"`
ForeignID string `json:"foreign_id"`
User struct {
ID int64 `json:"id"`
Firstname string `json:"firstname"`
Lastname string `json:"lastname"`
Email string `json:"email"`
} `json:"user"`
}
if err := json.Unmarshal(body, &rows); err != nil {
return nil, err
}
items := make([]catalogItem, 0, len(rows))
for _, row := range rows {
if row.ID <= 0 {
continue
}
label := strings.TrimSpace(row.User.Firstname + " " + row.User.Lastname)
if label == "" {
label = strings.TrimSpace(row.User.Email)
}
if label == "" {
label = fmt.Sprintf("UCR %d", row.ID)
}
parts := []string{}
if row.User.Email != "" {
parts = append(parts, row.User.Email)
}
if row.ClusterID > 0 {
parts = append(parts, fmt.Sprintf("Einheit %d", row.ClusterID))
}
if row.ForeignID != "" {
parts = append(parts, row.ForeignID)
}
items = append(items, catalogItem{ID: row.ID, Label: label, Subtitle: strings.Join(parts, " · ")})
}
return mergeCatalogItems(nil, items), nil
}
func mergeCatalogItems(base, add []catalogItem) []catalogItem {
byID := make(map[int64]catalogItem, len(base)+len(add))
for _, item := range base {
if item.ID > 0 {
byID[item.ID] = item
}
}
for _, item := range add {
if item.ID <= 0 {
continue
}
if old, ok := byID[item.ID]; ok {
// Prefer the richer v3/person label over generic fallback labels.
if strings.HasPrefix(old.Label, "ID ") || strings.HasPrefix(old.Label, "UCR ") || len(item.Label) > len(old.Label) {
old.Label = item.Label
}
if item.Subtitle != "" {
old.Subtitle = item.Subtitle
}
byID[item.ID] = old
} else {
byID[item.ID] = item
}
}
out := make([]catalogItem, 0, len(byID))
for _, item := range byID {
out = append(out, item)
}
sort.Slice(out, func(i, j int) bool {
li, lj := strings.ToLower(out[i].Label), strings.ToLower(out[j].Label)
if li == lj {
return out[i].ID < out[j].ID
}
return li < lj
})
return out
}
func sortCatalog(c *diveraCatalog) {
sorter := func(items []catalogItem) {
sort.Slice(items, func(i, j int) bool {
li, lj := strings.ToLower(items[i].Label), strings.ToLower(items[j].Label)
if li == lj {
return items[i].ID < items[j].ID
}
return li < lj
})
}
sorter(c.Units)
sorter(c.Groups)
sorter(c.Persons)
sorter(c.Vehicles)
}
type previewRequest struct {
Mapping config.Mapping `json:"mapping"`
Message model.InboundMessage `json:"message"`
}
func (s *Server) apiPreview(w http.ResponseWriter, r *http.Request) {
var in previewRequest
if err := json.NewDecoder(io.LimitReader(r.Body, 1<<20)).Decode(&in); err != nil {
writeJSON(w, http.StatusBadRequest, map[string]any{"error": err.Error()})
return
}
payload, kind, matched, err := gateway.Preview(in.Mapping, in.Message)
if err != nil {
writeJSON(w, http.StatusBadRequest, map[string]any{"error": err.Error(), "matched": matched})
return
}
writeJSON(w, http.StatusOK, map[string]any{"matched": matched, "kind": kind, "payload": payload})
}
+271
View File
@@ -0,0 +1,271 @@
package httpserver
import (
"encoding/json"
"io"
"log"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"sync"
"testing"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/divera"
)
func TestExtractCatalogFromKeyedPullData(t *testing.T) {
raw := map[string]any{
"data": map[string]any{
"cluster": map[string]any{"101": map[string]any{"title": "Löschzug 1"}},
"group": map[string]any{"202": map[string]any{"title": "Atemschutz"}},
"user_cluster_relation": map[string]any{"303": map[string]any{"firstname": "Max", "lastname": "Muster"}},
"vehicle": map[string]any{"404": map[string]any{"name": "HLF 20", "ric": "1234567"}},
},
}
got := extractCatalog(raw)
if len(got.Units) != 1 || got.Units[0].ID != 101 || got.Units[0].Label != "Löschzug 1" {
t.Fatalf("units = %#v", got.Units)
}
if len(got.Groups) != 1 || got.Groups[0].ID != 202 {
t.Fatalf("groups = %#v", got.Groups)
}
if len(got.Persons) != 1 || got.Persons[0].ID != 303 || got.Persons[0].Label != "Max Muster" {
t.Fatalf("persons = %#v", got.Persons)
}
if len(got.Vehicles) != 1 || got.Vehicles[0].ID != 404 {
t.Fatalf("vehicles = %#v", got.Vehicles)
}
}
func TestExtractCatalogPrefersUCRCollectionKeyOverGlobalUserID(t *testing.T) {
// Divera247 documents cluster.consumer as the user master-data collection.
// It can be keyed by UCR while a nested id refers to the global User-ID.
// Alarm recipients need the UCR key (9876).
raw := map[string]any{
"data": map[string]any{
"cluster": map[string]any{
"id": 42,
"consumer": map[string]any{
"9876": map[string]any{
"id": 1234,
"firstname": "Erika",
"lastname": "Mustermann",
"email": "erika@example.invalid",
},
},
},
},
}
got := extractCatalog(raw)
if len(got.Persons) != 1 {
t.Fatalf("persons = %#v", got.Persons)
}
if got.Persons[0].ID != 9876 {
t.Fatalf("expected UCR id 9876, got %#v", got.Persons[0])
}
if got.Persons[0].Label != "Erika Mustermann" {
t.Fatalf("unexpected label: %#v", got.Persons[0])
}
}
func TestParseV3PersonsUsesUserClusterRelationID(t *testing.T) {
body := []byte(`[
{"id":303,"cluster_id":42,"foreign_id":"ext-1","user":{"id":999,"firstname":"Max","lastname":"Muster","email":"max@example.invalid"}}
]`)
got, err := parseV3Persons(body)
if err != nil {
t.Fatal(err)
}
if len(got) != 1 || got[0].ID != 303 || got[0].Label != "Max Muster" {
t.Fatalf("persons = %#v", got)
}
if !strings.Contains(got[0].Subtitle, "Einheit 42") {
t.Fatalf("subtitle = %q", got[0].Subtitle)
}
}
func TestExtractUCRRefsFromPullAll(t *testing.T) {
raw := map[string]any{
"data": map[string]any{
"ucr_active": float64(555),
"ucr": []any{
map[string]any{"id": float64(555), "cluster_id": float64(42), "name": "Einheit Nord", "shortname": "N"},
map[string]any{"id": float64(666), "cluster_id": float64(43), "name": "Einheit Süd"},
},
},
}
refs := extractUCRRefs(raw)
if len(refs) != 2 {
t.Fatalf("refs=%#v", refs)
}
if extractActiveUCR(raw) != 555 {
t.Fatalf("active=%d", extractActiveUCR(raw))
}
byID := map[int64]ucrRef{}
for _, ref := range refs {
byID[ref.UCRID] = ref
}
if byID[555].ClusterID != 42 || byID[666].ClusterID != 43 {
t.Fatalf("refs=%#v", refs)
}
}
func TestDiveraCatalogLoadsConsumersAcrossV2UCRsWithoutV3(t *testing.T) {
var mu sync.Mutex
pullCalls := []string{}
v3Calls := 0
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/v2/pull/all":
mu.Lock()
pullCalls = append(pullCalls, r.URL.Query().Get("ucr"))
mu.Unlock()
if r.URL.Query().Get("accesskey") != "test-key" {
t.Fatalf("missing accesskey: %s", r.URL.RawQuery)
}
switch r.URL.Query().Get("ucr") {
case "555":
_, _ = io.WriteString(w, `{"data":{"ucr_active":555,"ucr":[{"id":555,"cluster_id":42,"name":"Einheit 42"},{"id":666,"cluster_id":43,"name":"Einheit 43"}],"cluster":{"id":42,"consumer":{"777":{"id":12,"firstname":"Lisa","lastname":"Beispiel","email":"lisa@example.invalid"}},"group":{"11":{"title":"Gruppe Nord"}},"vehicle":{}}}}`)
case "666":
_, _ = io.WriteString(w, `{"data":{"ucr_active":666,"ucr":[{"id":555,"cluster_id":42,"name":"Einheit 42"},{"id":666,"cluster_id":43,"name":"Einheit 43"}],"cluster":{"id":43,"consumer":{"888":{"id":13,"firstname":"Max","lastname":"Süd"}},"group":{"22":{"title":"Gruppe Süd"}},"vehicle":{}}}}`)
default:
http.Error(w, "unexpected ucr", http.StatusBadRequest)
}
case "/api/v3/user-cluster-relations":
v3Calls++
http.Error(w, `{"message":"Nicht autorisiert"}`, http.StatusForbidden)
default:
http.NotFound(w, r)
}
}))
defer upstream.Close()
store, err := config.Open(filepath.Join(t.TempDir(), "config.json"))
if err != nil {
t.Fatal(err)
}
if err := store.Update(func(c *config.Config) error {
c.Divera.BaseURL = upstream.URL
c.Divera.AccessKey = "test-key"
c.Divera.UCR = 555
c.Divera.DryRun = true
return nil
}); err != nil {
t.Fatal(err)
}
client := divera.New(func() config.DiveraConfig { return store.Get().Divera })
srv := New(store, nil, client, log.New(io.Discard, "", 0))
req := httptest.NewRequest(http.MethodGet, "/ui/api/divera247/catalog", nil)
rr := httptest.NewRecorder()
srv.apiDiveraCatalog(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
var out struct {
Catalog diveraCatalog `json:"catalog"`
Warnings []string `json:"warnings"`
Diagnostics struct {
V2Calls int `json:"v2_ucr_calls"`
V3Calls int `json:"v3_user_calls"`
Source string `json:"person_source"`
} `json:"diagnostics"`
}
if err := json.Unmarshal(rr.Body.Bytes(), &out); err != nil {
t.Fatal(err)
}
if len(out.Catalog.Persons) != 2 {
t.Fatalf("persons=%#v warnings=%#v", out.Catalog.Persons, out.Warnings)
}
ids := map[int64]bool{}
for _, person := range out.Catalog.Persons {
ids[person.ID] = true
if !strings.Contains(person.Subtitle, "Cluster-ID") {
t.Fatalf("missing unit subtitle: %#v", person)
}
}
if !ids[777] || !ids[888] {
t.Fatalf("persons=%#v", out.Catalog.Persons)
}
if v3Calls != 0 || out.Diagnostics.V3Calls != 0 {
t.Fatalf("v3 must not be called when v2 consumers exist: upstream=%d diagnostics=%d", v3Calls, out.Diagnostics.V3Calls)
}
if out.Diagnostics.Source != "v2 pull/all · cluster.consumer" {
t.Fatalf("source=%q", out.Diagnostics.Source)
}
mu.Lock()
defer mu.Unlock()
if len(pullCalls) != 2 {
t.Fatalf("pullCalls=%#v", pullCalls)
}
}
func TestDiveraCatalogStopsV3FallbackAfterForbidden(t *testing.T) {
v3Calls := 0
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/v2/pull/all":
ucr := r.URL.Query().Get("ucr")
if ucr == "555" {
_, _ = io.WriteString(w, `{"data":{"ucr_active":555,"ucr":[{"id":555,"cluster_id":42,"name":"Einheit 42"},{"id":666,"cluster_id":43,"name":"Einheit 43"}],"cluster":{"id":42,"consumer":{}}}}`)
} else {
_, _ = io.WriteString(w, `{"data":{"ucr_active":666,"ucr":[{"id":555,"cluster_id":42,"name":"Einheit 42"},{"id":666,"cluster_id":43,"name":"Einheit 43"}],"cluster":{"id":43,"consumer":{}}}}`)
}
case "/api/v3/user-cluster-relations":
v3Calls++
w.WriteHeader(http.StatusForbidden)
_, _ = io.WriteString(w, `{"name":"Forbidden","message":"Nicht autorisiert","code":0,"status":403}`)
default:
http.NotFound(w, r)
}
}))
defer upstream.Close()
store, err := config.Open(filepath.Join(t.TempDir(), "config.json"))
if err != nil {
t.Fatal(err)
}
if err := store.Update(func(c *config.Config) error {
c.Divera.BaseURL = upstream.URL
c.Divera.AccessKey = "test-key"
c.Divera.UCR = 555
c.Divera.DryRun = true
return nil
}); err != nil {
t.Fatal(err)
}
client := divera.New(func() config.DiveraConfig { return store.Get().Divera })
srv := New(store, nil, client, log.New(io.Discard, "", 0))
rr := httptest.NewRecorder()
srv.apiDiveraCatalog(rr, httptest.NewRequest(http.MethodGet, "/ui/api/divera247/catalog", nil))
if rr.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
var out struct {
Warnings []string `json:"warnings"`
}
if err := json.Unmarshal(rr.Body.Bytes(), &out); err != nil {
t.Fatal(err)
}
if v3Calls != 1 {
t.Fatalf("expected one v3 probe, got %d", v3Calls)
}
if len(out.Warnings) != 1 || !strings.Contains(out.Warnings[0], "HTTP 403") {
t.Fatalf("warnings=%#v", out.Warnings)
}
}
func TestEditableConfigDoesNotExposeSessionSecrets(t *testing.T) {
c := config.Default()
c.Server.AdminPasswordHash = "secret-hash"
c.Server.SessionSecret = "session-secret"
b, err := json.Marshal(makeEditable(c))
if err != nil {
t.Fatal(err)
}
s := string(b)
if strings.Contains(s, "secret-hash") || strings.Contains(s, "session-secret") || strings.Contains(s, "admin_password_hash") || strings.Contains(s, "session_secret") {
t.Fatalf("editable config leaked internal auth data: %s", s)
}
}
+98
View File
@@ -0,0 +1,98 @@
package httpserver
import (
"context"
"crypto/ed25519"
"encoding/hex"
"encoding/json"
"io"
"net/http"
"slices"
"strconv"
"time"
"github.com/example/notify-gateway/internal/model"
)
func (s *Server) discord(w http.ResponseWriter, r *http.Request) {
c := s.store.Get().Ingress.Discord
if !c.Enabled {
http.NotFound(w, r)
return
}
body, err := io.ReadAll(io.LimitReader(r.Body, (1<<20)+1))
if err != nil || len(body) > 1<<20 {
http.Error(w, "invalid request", 400)
return
}
stamp := r.Header.Get("X-Signature-Timestamp")
ts, err := strconv.ParseInt(stamp, 10, 64)
key, e1 := hex.DecodeString(c.PublicKey)
sig, e2 := hex.DecodeString(r.Header.Get("X-Signature-Ed25519"))
if err != nil || e1 != nil || e2 != nil || len(key) != ed25519.PublicKeySize || ts < time.Now().Add(-5*time.Minute).Unix() || ts > time.Now().Add(5*time.Minute).Unix() || !ed25519.Verify(key, append([]byte(stamp), body...), sig) {
http.Error(w, "invalid signature", 401)
return
}
var in struct {
ID string `json:"id"`
ApplicationID string `json:"application_id"`
Type int `json:"type"`
GuildID string `json:"guild_id"`
ChannelID string `json:"channel_id"`
Data struct {
Name string `json:"name"`
Type int `json:"type"`
Options []struct {
Name string `json:"name"`
Type int `json:"type"`
Value json.RawMessage `json:"value"`
} `json:"options"`
} `json:"data"`
}
if json.Unmarshal(body, &in) != nil || in.ApplicationID != c.ApplicationID {
http.Error(w, "invalid interaction", 400)
return
}
if in.Type == 1 {
writeJSON(w, 200, map[string]int{"type": 1})
return
}
if in.Type != 2 || in.Data.Type != 1 || in.Data.Name != c.Command || in.ID == "" || !slices.Contains(c.GuildIDs, in.GuildID) || !slices.Contains(c.ChannelIDs, in.ChannelID) {
http.Error(w, "interaction not allowed", 403)
return
}
msg := model.InboundMessage{Source: "discord", Channel: in.ChannelID, ReceivedAt: time.Now().UTC()}
for _, o := range in.Data.Options {
switch o.Name {
case "title":
err = json.Unmarshal(o.Value, &msg.Title)
case "message":
err = json.Unmarshal(o.Value, &msg.Message)
case "priority":
err = json.Unmarshal(o.Value, &msg.Priority)
default:
http.Error(w, "unsupported option", 400)
return
}
if err != nil {
http.Error(w, "invalid option", 400)
return
}
}
if msg.Message == "" {
http.Error(w, "message required", 400)
return
}
if s.queue == nil {
http.Error(w, "outbox unavailable", 503)
return
}
ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second)
defer cancel()
receipt, err := s.queue.Accept(ctx, msg, "discord:"+in.ID)
if err != nil {
http.Error(w, "interaction could not be queued; check routing", 503)
return
}
writeJSON(w, 200, map[string]any{"type": 4, "data": map[string]any{"content": "Meldung angenommen. Referenz: " + receipt.ID, "flags": 64, "allowed_mentions": map[string]any{"parse": []string{}}}})
}
+118
View File
@@ -0,0 +1,118 @@
package httpserver
import (
"context"
"crypto/hmac"
"crypto/sha256"
"crypto/subtle"
"encoding/hex"
"errors"
"fmt"
"net/http"
"os"
"strconv"
"time"
"github.com/example/notify-gateway/internal/gateway"
"github.com/example/notify-gateway/internal/mailingress"
"github.com/example/notify-gateway/internal/outbox"
)
func (s *Server) UseQueue(q *gateway.Queue, p *mailingress.Poller) { s.queue = q; s.mail = p }
func (s *Server) ready(w http.ResponseWriter, r *http.Request) {
ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second)
defer cancel()
if s.queue == nil || s.queue.Store.Ping(ctx) != nil {
writeJSON(w, 503, map[string]any{"ok": false})
return
}
writeJSON(w, 200, map[string]any{"ok": true})
}
func (s *Server) deliveries(w http.ResponseWriter, r *http.Request) {
if s.queue == nil {
http.Error(w, "outbox unavailable", 503)
return
}
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
jobs, err := s.queue.Store.List(r.Context(), 50, offset)
if err != nil {
http.Error(w, "outbox unavailable", 503)
return
}
counts, err := s.queue.Store.Counts(r.Context())
if err != nil {
http.Error(w, "outbox unavailable", 503)
return
}
statuses := []mailingress.Status{}
if s.mail != nil {
statuses = s.mail.Statuses()
}
writeJSON(w, 200, map[string]any{"deliveries": jobs, "counts": counts, "mail": statuses})
}
func (s *Server) deliveryHistory(w http.ResponseWriter, r *http.Request) {
if s.queue == nil {
http.Error(w, "outbox unavailable", 503)
return
}
rows, err := s.queue.Store.History(r.Context(), r.PathValue("id"))
if err != nil {
http.Error(w, "outbox unavailable", 503)
return
}
writeJSON(w, 200, rows)
}
func (s *Server) retryDelivery(w http.ResponseWriter, r *http.Request) {
if s.queue == nil {
http.Error(w, "outbox unavailable", 503)
return
}
err := s.queue.Store.Retry(r.Context(), r.PathValue("id"))
if errors.Is(err, outbox.ErrNotRetryable) {
http.Error(w, err.Error(), 409)
return
}
if err != nil {
http.Error(w, "outbox unavailable", 503)
return
}
writeJSON(w, 200, map[string]any{"ok": true})
}
func (s *Server) metrics(w http.ResponseWriter, r *http.Request) {
token := os.Getenv("GATEWAY_METRICS_TOKEN")
if token == "" {
s.requireAdmin(s.writeMetrics)(w, r)
return
}
got := r.Header.Get("Authorization")
if subtle.ConstantTimeCompare([]byte(got), []byte("Bearer "+token)) != 1 {
http.Error(w, "unauthorized", 401)
return
}
s.writeMetrics(w, r)
}
func (s *Server) writeMetrics(w http.ResponseWriter, r *http.Request) {
if s.queue == nil {
http.Error(w, "outbox unavailable", 503)
return
}
counts, err := s.queue.Store.Counts(r.Context())
if err != nil {
http.Error(w, "outbox unavailable", 503)
return
}
w.Header().Set("Content-Type", "text/plain; version=0.0.4")
fmt.Fprintln(w, "# HELP notify_gateway_deliveries Persisted deliveries by current state.\n# TYPE notify_gateway_deliveries gauge")
for _, state := range []string{"pending", "sending", "succeeded", "dry_run", "dead"} {
fmt.Fprintf(w, "notify_gateway_deliveries{state=%q} %d\n", state, counts[state])
}
}
func csrfToken(secret, session string) string {
mac := hmac.New(sha256.New, []byte(secret))
mac.Write([]byte("csrf:" + session))
return hex.EncodeToString(mac.Sum(nil))
}
func (s *Server) csrf(w http.ResponseWriter, r *http.Request) {
c, _ := r.Cookie("ng_session")
writeJSON(w, 200, map[string]string{"token": csrfToken(s.store.Get().Server.SessionSecret, c.Value)})
}
+181
View File
@@ -0,0 +1,181 @@
package httpserver
import (
"bytes"
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/hex"
"encoding/json"
"io"
"log"
"net/http"
"net/http/httptest"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
"github.com/example/notify-gateway/internal/auth"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/divera"
"github.com/example/notify-gateway/internal/gateway"
"github.com/example/notify-gateway/internal/outbox"
)
func queuedServer(t *testing.T) (*Server, *config.Store, *outbox.Store) {
t.Helper()
dir := t.TempDir()
store, err := config.Open(filepath.Join(dir, "config.json"))
if err != nil {
t.Fatal(err)
}
db, err := outbox.Open(filepath.Join(dir, "outbox.db"))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { db.Close() })
client := divera.New(func() config.DiveraConfig { return store.Get().Divera })
dispatcher := gateway.New(store, client)
s := New(store, dispatcher, client, log.New(io.Discard, "", 0))
s.UseQueue(&gateway.Queue{Store: db, Dispatcher: dispatcher}, nil)
return s, store, db
}
func TestAsyncIngressIdempotencyReadinessAndMetrics(t *testing.T) {
s, store, db := queuedServer(t)
c := store.Get()
c.Ingress.WebhookTokens = map[string]string{"ops": "secret"}
c.Mappings = []config.Mapping{{ID: "route", Enabled: true, Target: "news", TextTemplate: "{{.Message}}"}}
store.Replace(c)
h := s.Handler()
send := func(body string) *httptest.ResponseRecorder {
r := httptest.NewRequest("POST", "/in/webhook/ops", strings.NewReader(body))
r.Header.Set("Authorization", "Bearer secret")
r.Header.Set("Idempotency-Key", "same-request")
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
return w
}
for i := 0; i < 2; i++ {
if w := send("hello"); w.Code != 202 {
t.Fatalf("accepted=%d %s", w.Code, w.Body)
}
}
if w := send("different"); w.Code != 409 {
t.Fatal(w.Code)
}
rows, _ := db.List(context.Background(), 50, 0)
if len(rows) != 1 || rows[0].State != "pending" {
t.Fatal(rows)
}
for _, path := range []string{"/readyz", "/metrics"} {
w := httptest.NewRecorder()
h.ServeHTTP(w, httptest.NewRequest("GET", path, nil))
if path == "/readyz" && w.Code != 200 {
t.Fatal(w.Code)
}
if path == "/metrics" && w.Code != 303 {
t.Fatal("public metrics")
}
}
t.Setenv("GATEWAY_METRICS_TOKEN", "metric-token")
r := httptest.NewRequest("GET", "/metrics", nil)
r.Header.Set("Authorization", "Bearer metric-token")
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
if w.Code != 200 || !strings.Contains(w.Body.String(), `state="pending"} 1`) {
t.Fatalf("metrics %d %s", w.Code, w.Body)
}
}
func TestCSRFRequestLimitsAndLoginRateLimit(t *testing.T) {
s, store, _ := queuedServer(t)
c := store.Get()
h := s.Handler()
session := auth.SignSession(c.Server.SessionSecret, c.Server.AdminUsername, time.Now().Add(time.Hour))
for _, valid := range []bool{false, true} {
body, _ := json.Marshal(makeEditable(c))
r := httptest.NewRequest("PUT", "/ui/api/config", bytes.NewReader(body))
r.AddCookie(&http.Cookie{Name: "ng_session", Value: session})
if valid {
r.Header.Set("X-CSRF-Token", csrfToken(c.Server.SessionSecret, session))
}
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
want := 403
if valid {
want = 200
}
if w.Code != want {
t.Fatalf("CSRF valid=%v status=%d %s", valid, w.Code, w.Body)
}
}
r := httptest.NewRequest("POST", "/login", nil)
r.Header.Set("Origin", "https://evil.example")
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
if w.Code != 403 {
t.Fatal("cross-site login accepted")
}
w = httptest.NewRecorder()
h.ServeHTTP(w, httptest.NewRequest("POST", "/ntfy/ops", strings.NewReader(strings.Repeat("x", (1<<20)+1))))
if w.Code != 413 {
t.Fatal("large body accepted")
}
for i := 0; i < 11; i++ {
w = httptest.NewRecorder()
h.ServeHTTP(w, httptest.NewRequest("POST", "/login", nil))
}
if w.Code != 429 {
t.Fatalf("rate limit %d", w.Code)
}
}
func TestDiscordSignatureAllowlistAndReplay(t *testing.T) {
s, store, db := queuedServer(t)
pub, priv, _ := ed25519.GenerateKey(rand.Reader)
c := store.Get()
c.Ingress.Discord = config.DiscordIngress{Enabled: true, PublicKey: hex.EncodeToString(pub), ApplicationID: "app", GuildIDs: []string{"guild"}, ChannelIDs: []string{"channel"}, Command: "notify"}
c.Mappings = []config.Mapping{{ID: "discord", Enabled: true, Source: "discord", Target: "news", TextTemplate: "{{.Message}}"}}
if err := store.Replace(c); err != nil {
t.Fatal(err)
}
h := s.Handler()
invoke := func(body string, stamp time.Time, valid bool) *httptest.ResponseRecorder {
r := httptest.NewRequest("POST", "/in/discord", strings.NewReader(body))
ts := strconv.FormatInt(stamp.Unix(), 10)
sig := ed25519.Sign(priv, []byte(ts+body))
if !valid {
sig[0] ^= 1
}
r.Header.Set("X-Signature-Timestamp", ts)
r.Header.Set("X-Signature-Ed25519", hex.EncodeToString(sig))
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
return w
}
ping := `{"type":1,"application_id":"app"}`
if w := invoke(ping, time.Now(), true); w.Code != 200 || !strings.Contains(w.Body.String(), `"type":1`) {
t.Fatalf("ping %d %s", w.Code, w.Body)
}
if w := invoke(ping, time.Now(), false); w.Code != 401 {
t.Fatal("bad signature accepted")
}
if w := invoke(ping, time.Now().Add(-6*time.Minute), true); w.Code != 401 {
t.Fatal("stale request accepted")
}
body := `{"id":"interaction","application_id":"app","type":2,"guild_id":"guild","channel_id":"channel","data":{"name":"notify","type":1,"options":[{"name":"message","type":3,"value":"hello"}]}}`
for i := 0; i < 2; i++ {
if w := invoke(body, time.Now(), true); w.Code != 200 {
t.Fatalf("command %d %s", w.Code, w.Body)
}
}
if w := invoke(strings.Replace(body, `"guild"`, `"other"`, 1), time.Now(), true); w.Code != 403 {
t.Fatal("guild allowlist bypass")
}
rows, _ := db.List(context.Background(), 50, 0)
if len(rows) != 1 {
t.Fatalf("duplicate interaction: %+v", rows)
}
}
+106
View File
@@ -0,0 +1,106 @@
package httpserver
import (
"bytes"
"encoding/json"
"io"
"log"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"time"
"github.com/example/notify-gateway/internal/auth"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/divera"
"github.com/example/notify-gateway/internal/gateway"
)
func TestOutboundAdminRoundTripPreviewAndIngress(t *testing.T) {
configPath := filepath.Join(t.TempDir(), "config.json")
store, err := config.Open(configPath)
if err != nil {
t.Fatal(err)
}
client := divera.New(func() config.DiveraConfig { return store.Get().Divera })
h := New(store, gateway.New(store, client), client, log.New(io.Discard, "", 0)).Handler()
c := store.Get()
cookie := &http.Cookie{Name: "ng_session", Value: auth.SignSession(c.Server.SessionSecret, c.Server.AdminUsername, time.Now().Add(time.Hour))}
edit := makeEditable(c)
edit.Outbounds = []config.OutboundConfig{{ID: "discord-ops", Name: "Ops", Provider: "discord", URL: "https://discord.com/api/webhooks/123/secret"}}
edit.Ingress.NtfyTokens = map[string]string{"ops": "ingress-token"}
edit.Mappings = []config.Mapping{{ID: "route", Name: "Route", Enabled: true, Source: "ntfy", Target: "discord", OutboundID: "discord-ops", TitleTemplate: "{{.Title}}", TextTemplate: "{{.Message}}"}}
body, _ := json.Marshal(edit)
request := func(method, path string, body []byte, admin bool) *httptest.ResponseRecorder {
r := httptest.NewRequest(method, path, bytes.NewReader(body))
r.Header.Set("Content-Type", "application/json")
if admin {
r.AddCookie(cookie)
r.Header.Set("X-CSRF-Token", csrfToken(c.Server.SessionSecret, cookie.Value))
}
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
return w
}
if w := request("PUT", "/ui/api/config", body, false); w.Code != 303 {
t.Fatalf("unprotected admin API: %d", w.Code)
}
if w := request("PUT", "/ui/api/config", body, true); w.Code != 200 {
t.Fatalf("save: %d %s", w.Code, w.Body)
}
if store.Get().Server.SessionSecret != c.Server.SessionSecret {
t.Fatal("admin secret changed")
}
reopened, err := config.Open(configPath)
if err != nil {
t.Fatal(err)
}
if len(reopened.Get().Outbounds) != 1 || reopened.Get().Outbounds[0].ID != "discord-ops" {
t.Fatal("outbound was not persisted")
}
w := request("GET", "/ui/api/config", nil, true)
var loaded editableConfig
if err := json.Unmarshal(w.Body.Bytes(), &loaded); err != nil {
t.Fatal(err)
}
if len(loaded.Outbounds) != 1 || loaded.Outbounds[0].Live || loaded.Mappings[0].OutboundID != "discord-ops" {
t.Fatalf("round trip: %+v", loaded)
}
preview, _ := json.Marshal(map[string]any{"mapping": edit.Mappings[0], "message": map[string]any{"source": "ntfy", "title": "Hello", "message": "World"}})
w = request("POST", "/ui/api/preview", preview, true)
if w.Code != 200 || !strings.Contains(w.Body.String(), `"content":"Hello\nWorld"`) {
t.Fatalf("preview: %d %s", w.Code, w.Body)
}
for _, authorized := range []bool{false, true} {
r := httptest.NewRequest("POST", "/ntfy/ops", strings.NewReader("World"))
r.Header.Set("Title", "Hello")
if authorized {
r.Header.Set("Authorization", "Bearer ingress-token")
}
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
if !authorized {
if w.Code != 401 {
t.Fatalf("ingress auth: %d", w.Code)
}
continue
}
if w.Code != 200 || !strings.Contains(w.Body.String(), `"outbound_id":"discord-ops"`) || !strings.Contains(w.Body.String(), `dry_run`) {
t.Fatalf("delivery: %d %s", w.Code, w.Body)
}
if strings.Contains(w.Body.String(), "secret") {
t.Fatal("destination token leaked")
}
}
// A destination still referenced by a mapping cannot be removed.
edit.Outbounds = nil
body, _ = json.Marshal(edit)
if w := request("PUT", "/ui/api/config", body, true); w.Code != 400 {
t.Fatalf("invalid reference: %d %s", w.Code, w.Body)
}
if len(store.Get().Outbounds) != 1 {
t.Fatal("invalid config was persisted")
}
}
+85
View File
@@ -0,0 +1,85 @@
package httpserver
import (
"bytes"
"io"
"net"
"net/http"
"net/url"
"strings"
"time"
)
type loginBucket struct {
Start time.Time
Count int
}
func (s *Server) allowLogin(address string) bool {
ip, _, err := net.SplitHostPort(address)
if err != nil {
ip = address
}
s.loginMu.Lock()
defer s.loginMu.Unlock()
if s.logins == nil {
s.logins = map[string]loginBucket{}
}
now := time.Now()
for k, v := range s.logins {
if now.Sub(v.Start) > time.Minute {
delete(s.logins, k)
}
}
b := s.logins[ip]
if b.Start.IsZero() {
if len(s.logins) >= 1024 {
return false
}
b.Start = now
}
b.Count++
s.logins[ip] = b
return b.Count <= 10
}
func requestProtection(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// No proxy headers are trusted. Configure the proxy to preserve the Host.
admin := strings.HasPrefix(r.URL.Path, "/ui/") || r.URL.Path == "/login" || r.URL.Path == "/logout"
if admin {
w.Header().Set("Cache-Control", "no-store")
}
if admin && r.Method != "GET" && r.Method != "HEAD" {
if origin := r.Header.Get("Origin"); origin != "" {
u, err := url.Parse(origin)
if err != nil || u.Host != r.Host || (u.Scheme != "http" && u.Scheme != "https") {
http.Error(w, "cross-origin request rejected", 403)
return
}
}
if r.Header.Get("Sec-Fetch-Site") == "cross-site" {
http.Error(w, "cross-site request rejected", 403)
return
}
}
if r.Body != nil {
limit := int64(1 << 20)
if admin {
limit = 2 << 20
}
b, err := io.ReadAll(io.LimitReader(r.Body, limit+1))
r.Body.Close()
if err != nil {
http.Error(w, "invalid body", 400)
return
}
if int64(len(b)) > limit {
http.Error(w, "request too large", 413)
return
}
r.Body = io.NopCloser(bytes.NewReader(b))
}
next.ServeHTTP(w, r)
})
}
+513
View File
@@ -0,0 +1,513 @@
package httpserver
import (
"crypto/sha256"
"crypto/subtle"
_ "embed"
"encoding/base64"
"encoding/json"
"errors"
"html/template"
"io"
"log"
"net/http"
"strconv"
"strings"
"sync"
"time"
"github.com/example/notify-gateway/internal/auth"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/divera"
"github.com/example/notify-gateway/internal/gateway"
"github.com/example/notify-gateway/internal/mailingress"
"github.com/example/notify-gateway/internal/model"
"github.com/example/notify-gateway/internal/outbox"
)
type Server struct {
queue *gateway.Queue
mail *mailingress.Poller
loginMu sync.Mutex
logins map[string]loginBucket
store *config.Store
dispatcher *gateway.Dispatcher
divera *divera.Client
logger *log.Logger
}
func New(store *config.Store, dispatcher *gateway.Dispatcher, client *divera.Client, logger *log.Logger) *Server {
return &Server{store: store, dispatcher: dispatcher, divera: client, logger: logger}
}
func (s *Server) Handler() http.Handler {
mux := http.NewServeMux()
mux.HandleFunc("GET /healthz", s.health)
mux.HandleFunc("GET /readyz", s.ready)
mux.HandleFunc("GET /metrics", s.metrics)
mux.HandleFunc("GET /ui/api/csrf", s.requireAdmin(s.csrf))
mux.HandleFunc("GET /ui/api/deliveries", s.requireAdmin(s.deliveries))
mux.HandleFunc("GET /ui/api/deliveries/{id}/attempts", s.requireAdmin(s.deliveryHistory))
mux.HandleFunc("POST /ui/api/deliveries/{id}/retry", s.requireAdmin(s.retryDelivery))
mux.HandleFunc("POST /in/discord", s.discord)
mux.HandleFunc("POST /login", s.login)
mux.HandleFunc("GET /login", s.loginPage)
mux.HandleFunc("POST /logout", s.requireAdmin(s.logout))
mux.HandleFunc("GET /ui/", s.requireAdmin(s.ui))
mux.HandleFunc("GET /ui/api/config", s.requireAdmin(s.apiConfig))
mux.HandleFunc("PUT /ui/api/config", s.requireAdmin(s.apiConfig))
mux.HandleFunc("GET /ui/api/divera247/catalog", s.requireAdmin(s.apiDiveraCatalog))
mux.HandleFunc("POST /ui/api/preview", s.requireAdmin(s.apiPreview))
mux.HandleFunc("POST /ui/password", s.requireAdmin(s.changePassword))
mux.HandleFunc("POST /ui/test-divera", s.requireAdmin(s.testDivera))
mux.HandleFunc("POST /in/webhook/{channel}", s.webhook)
mux.HandleFunc("PUT /in/webhook/{channel}", s.webhook)
mux.HandleFunc("POST /in/ntfy", s.ntfyJSON)
mux.HandleFunc("PUT /in/ntfy", s.ntfyJSON)
mux.HandleFunc("POST /ntfy", s.ntfyJSON)
mux.HandleFunc("PUT /ntfy", s.ntfyJSON)
mux.HandleFunc("POST /in/ntfy/{topic}", s.ntfyTopic)
mux.HandleFunc("PUT /in/ntfy/{topic}", s.ntfyTopic)
mux.HandleFunc("GET /in/ntfy/{topic}/trigger", s.ntfyTrigger)
mux.HandleFunc("GET /in/ntfy/{topic}/send", s.ntfyTrigger)
mux.HandleFunc("GET /in/ntfy/{topic}/publish", s.ntfyTrigger)
mux.HandleFunc("POST /ntfy/{topic}", s.ntfyTopic)
mux.HandleFunc("PUT /ntfy/{topic}", s.ntfyTopic)
mux.HandleFunc("GET /ntfy/{topic}/trigger", s.ntfyTrigger)
mux.HandleFunc("GET /ntfy/{topic}/send", s.ntfyTrigger)
mux.HandleFunc("GET /ntfy/{topic}/publish", s.ntfyTrigger)
mux.HandleFunc("POST /in/gotify/message", s.gotify)
mux.HandleFunc("POST /message", s.gotify)
mux.HandleFunc("GET /", func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "/ui/", http.StatusFound) })
return securityHeaders(requestProtection(mux))
}
func securityHeaders(next http.Handler) http.Handler {
script := strings.Split(strings.Split(uiHTML, "<script>")[1], "</script>")[0]
// HTML parsing normalizes CRLF before CSP hashes are checked.
sum := sha256.Sum256([]byte(strings.ReplaceAll(script, "\r\n", "\n")))
scriptHash := base64.StdEncoding.EncodeToString(sum[:])
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Security-Policy", "default-src 'self'; script-src 'sha256-"+scriptHash+"'; style-src 'self' 'unsafe-inline'; img-src 'self' data:; base-uri 'none'; form-action 'self'; frame-ancestors 'none'")
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("X-Frame-Options", "DENY")
w.Header().Set("Referrer-Policy", "same-origin")
next.ServeHTTP(w, r)
})
}
func (s *Server) health(w http.ResponseWriter, r *http.Request) {
writeJSON(w, 200, map[string]any{"ok": true, "time": time.Now().UTC()})
}
func (s *Server) loginPage(w http.ResponseWriter, r *http.Request) { renderLogin(w, "") }
func (s *Server) login(w http.ResponseWriter, r *http.Request) {
if !s.allowLogin(r.RemoteAddr) {
w.Header().Set("Retry-After", "60")
http.Error(w, "too many login attempts", 429)
return
}
_ = r.ParseForm()
cfg := s.store.Get()
if r.Form.Get("username") != cfg.Server.AdminUsername || !auth.CheckPassword(cfg.Server.AdminPasswordHash, r.Form.Get("password")) {
renderLogin(w, "Login fehlgeschlagen")
return
}
if strings.HasPrefix(cfg.Server.AdminPasswordHash, "sha256$") && len(r.Form.Get("password")) <= 72 {
h, err := auth.HashPassword(r.Form.Get("password"))
if err != nil {
http.Error(w, "password migration failed", 500)
return
}
if err = s.store.Update(func(c *config.Config) error { c.Server.AdminPasswordHash = h; return nil }); err != nil {
http.Error(w, "password migration failed", 500)
return
}
}
token := auth.SignSession(cfg.Server.SessionSecret, cfg.Server.AdminUsername, time.Now().Add(12*time.Hour))
http.SetCookie(w, &http.Cookie{Name: "ng_session", Value: token, Path: "/", HttpOnly: true, SameSite: http.SameSiteStrictMode, Secure: r.TLS != nil, MaxAge: 43200})
http.Redirect(w, r, "/ui/", http.StatusSeeOther)
}
func (s *Server) logout(w http.ResponseWriter, r *http.Request) {
http.SetCookie(w, &http.Cookie{Name: "ng_session", Value: "", Path: "/", HttpOnly: true, MaxAge: -1, SameSite: http.SameSiteStrictMode})
http.Redirect(w, r, "/login", http.StatusSeeOther)
}
func (s *Server) requireAdmin(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
cfg := s.store.Get()
c, err := r.Cookie("ng_session")
if err != nil {
http.Redirect(w, r, "/login", http.StatusSeeOther)
return
}
u, ok := auth.VerifySession(cfg.Server.SessionSecret, c.Value, time.Now())
if !ok || u != cfg.Server.AdminUsername {
http.Redirect(w, r, "/login", http.StatusSeeOther)
return
}
if r.Method != "GET" && r.Method != "HEAD" {
token := r.Header.Get("X-CSRF-Token")
if token == "" && strings.HasPrefix(r.Header.Get("Content-Type"), "application/x-www-form-urlencoded") {
_ = r.ParseForm()
token = r.Form.Get("csrf_token")
}
if subtle.ConstantTimeCompare([]byte(token), []byte(csrfToken(cfg.Server.SessionSecret, c.Value))) != 1 {
http.Error(w, "invalid CSRF token", 403)
return
}
}
next(w, r)
}
}
//go:embed ui.html
var uiHTML string
func (s *Server) ui(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_, _ = io.WriteString(w, uiHTML)
}
func (s *Server) changePassword(w http.ResponseWriter, r *http.Request) {
if err := r.ParseForm(); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
p := r.Form.Get("password")
if len(p) < 10 || len(p) > 72 {
http.Error(w, "Passwort muss 10 bis 72 UTF-8-Bytes enthalten", http.StatusBadRequest)
return
}
h, err := auth.HashPassword(p)
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
if err := s.store.Update(func(c *config.Config) error {
c.Server.AdminPasswordHash = h
c.Server.SessionSecret = auth.RandomSecret()
return nil
}); err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
http.Redirect(w, r, "/ui/", http.StatusSeeOther)
}
func (s *Server) testDivera(w http.ResponseWriter, r *http.Request) {
resp, err := s.divera.PullAll(r.Context(), nil)
if err != nil {
http.Error(w, err.Error(), 502)
return
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(resp.StatusCode)
_, _ = w.Write(resp.Body)
}
func (s *Server) authorized(source, channel string, r *http.Request) bool {
cfg := s.store.Get()
if cfg.Ingress.AllowUnauthenticated {
return true
}
token := bearerOrToken(r)
var allowed []string
switch source {
case "ntfy":
if t := cfg.Ingress.NtfyTokens[channel]; t != "" {
allowed = append(allowed, t)
}
if t := cfg.Ingress.NtfyTokens["*"]; t != "" {
allowed = append(allowed, t)
}
case "webhook":
if t := cfg.Ingress.WebhookTokens[channel]; t != "" {
allowed = append(allowed, t)
}
if t := cfg.Ingress.WebhookTokens["*"]; t != "" {
allowed = append(allowed, t)
}
case "gotify":
allowed = append(allowed, cfg.Ingress.GotifyTokens...)
}
for _, a := range allowed {
if len(a) == len(token) && subtle.ConstantTimeCompare([]byte(a), []byte(token)) == 1 {
return true
}
}
return false
}
func bearerOrToken(r *http.Request) string {
if t := r.URL.Query().Get("token"); t != "" {
return t
}
if t := r.Header.Get("X-Gotify-Key"); t != "" {
return t
}
a := r.Header.Get("Authorization")
if strings.HasPrefix(strings.ToLower(a), "bearer ") {
return strings.TrimSpace(a[7:])
}
if u, p, ok := r.BasicAuth(); ok {
if p != "" {
return p
}
return u
}
return ""
}
func (s *Server) webhook(w http.ResponseWriter, r *http.Request) {
channel := r.PathValue("channel")
if !s.authorized("webhook", channel, r) {
writeJSON(w, 401, map[string]any{"error": "unauthorized"})
return
}
raw, body, err := decodeFlexible(r)
if err != nil {
writeJSON(w, 400, map[string]any{"error": err.Error()})
return
}
msg := model.InboundMessage{Source: "webhook", Channel: channel, Title: str(raw, "title", "subject"), Message: firstNonEmpty(str(raw, "message", "text", "body"), body), Address: str(raw, "address", "location"), Priority: intval(raw, "priority"), Raw: raw, ReceivedAt: time.Now().UTC()}
s.deliver(w, r, msg)
}
func (s *Server) ntfyJSON(w http.ResponseWriter, r *http.Request) {
var raw map[string]any
if err := json.NewDecoder(io.LimitReader(r.Body, 1<<20)).Decode(&raw); err != nil {
writeJSON(w, 400, map[string]any{"error": "invalid ntfy json"})
return
}
topic := str(raw, "topic")
if topic == "" {
writeJSON(w, 400, map[string]any{"error": "topic required"})
return
}
if !s.authorized("ntfy", topic, r) {
writeJSON(w, 401, map[string]any{"error": "unauthorized"})
return
}
msg := model.InboundMessage{Source: "ntfy", Channel: topic, Title: str(raw, "title"), Message: str(raw, "message"), Priority: intval(raw, "priority"), Raw: raw, ReceivedAt: time.Now().UTC()}
s.deliver(w, r, msg)
}
func (s *Server) ntfyTopic(w http.ResponseWriter, r *http.Request) {
topic := r.PathValue("topic")
if !s.authorized("ntfy", topic, r) {
writeJSON(w, 401, map[string]any{"error": "unauthorized"})
return
}
raw, body, err := decodeNtfy(r)
if err != nil {
writeJSON(w, 400, map[string]any{"error": err.Error()})
return
}
title := firstNonEmpty(r.Header.Get("Title"), r.Header.Get("X-Title"), str(raw, "title"))
message := firstNonEmpty(str(raw, "message"), body)
priority := intval(raw, "priority")
if priority == 0 {
priority = parsePriority(firstNonEmpty(r.Header.Get("Priority"), r.Header.Get("X-Priority")))
}
msg := model.InboundMessage{Source: "ntfy", Channel: topic, Title: title, Message: message, Priority: priority, Raw: raw, ReceivedAt: time.Now().UTC()}
s.deliver(w, r, msg)
}
func (s *Server) ntfyTrigger(w http.ResponseWriter, r *http.Request) {
topic := r.PathValue("topic")
if !s.authorized("ntfy", topic, r) {
writeJSON(w, 401, map[string]any{"error": "unauthorized"})
return
}
msg := model.InboundMessage{Source: "ntfy", Channel: topic, Title: r.URL.Query().Get("title"), Message: firstNonEmpty(r.URL.Query().Get("message"), "triggered"), Priority: parsePriority(r.URL.Query().Get("priority")), Raw: map[string]any{"query": r.URL.Query()}, ReceivedAt: time.Now().UTC()}
s.deliver(w, r, msg)
}
func (s *Server) gotify(w http.ResponseWriter, r *http.Request) {
if !s.authorized("gotify", "", r) {
writeJSON(w, 401, map[string]any{"error": "unauthorized"})
return
}
raw, body, err := decodeFlexible(r)
if err != nil {
writeJSON(w, 400, map[string]any{"error": err.Error()})
return
}
if len(raw) == 0 {
_ = r.ParseForm()
raw = map[string]any{"title": r.Form.Get("title"), "message": r.Form.Get("message"), "priority": r.Form.Get("priority")}
}
msg := model.InboundMessage{Source: "gotify", Channel: firstNonEmpty(str(raw, "channel"), "default"), Title: str(raw, "title"), Message: firstNonEmpty(str(raw, "message"), body), Priority: intval(raw, "priority"), Raw: raw, ReceivedAt: time.Now().UTC()}
s.deliver(w, r, msg)
}
func (s *Server) deliver(w http.ResponseWriter, r *http.Request, msg model.InboundMessage) {
if s.queue != nil {
// Never accept a gateway-generated HTTP notification back into the gateway.
if r.Header.Get("X-Notify-Gateway-Delivery") != "" {
writeJSON(w, 409, map[string]any{"error": "gateway loop detected"})
return
}
receipt, err := s.queue.Accept(r.Context(), msg, r.Header.Get("Idempotency-Key"))
if err != nil {
status := 503
var input *gateway.InputError
if errors.As(err, &input) {
status = 400
}
if errors.Is(err, outbox.ErrConflict) {
status = 409
}
message := "outbox unavailable"
if status != 503 {
message = err.Error()
}
writeJSON(w, status, map[string]any{"error": message})
return
}
writeJSON(w, 202, map[string]any{"ok": true, "receipt": receipt})
return
}
results, err := s.dispatcher.Dispatch(r.Context(), msg)
if err != nil {
s.logger.Printf("delivery failed: %v", err)
writeJSON(w, 502, map[string]any{"error": err.Error(), "results": results})
return
}
writeJSON(w, 200, map[string]any{"ok": true, "results": results})
}
func decodeNtfy(r *http.Request) (map[string]any, string, error) {
b, err := io.ReadAll(io.LimitReader(r.Body, 1<<20))
if err != nil {
return nil, "", err
}
raw := map[string]any{}
if strings.Contains(r.Header.Get("Content-Type"), "application/json") || (len(b) > 0 && b[0] == '{') {
if err := json.Unmarshal(b, &raw); err != nil {
return nil, "", err
}
return raw, "", nil
}
return raw, string(b), nil
}
func decodeFlexible(r *http.Request) (map[string]any, string, error) {
b, err := io.ReadAll(io.LimitReader(r.Body, 1<<20))
if err != nil {
return nil, "", err
}
raw := map[string]any{}
ct := r.Header.Get("Content-Type")
if strings.Contains(ct, "application/json") || (len(b) > 0 && b[0] == '{') {
if err := json.Unmarshal(b, &raw); err != nil {
return nil, "", err
}
return raw, "", nil
}
if strings.Contains(ct, "application/x-www-form-urlencoded") {
vals, err := urlParse(string(b))
if err != nil {
return nil, "", err
}
for k, v := range vals {
if len(v) > 0 {
raw[k] = v[0]
}
}
return raw, "", nil
}
return raw, string(b), nil
}
func urlParse(s string) (map[string][]string, error) {
vals, err := netURLParseQuery(s)
return map[string][]string(vals), err
}
var netURLParseQuery = func(s string) (map[string][]string, error) {
out := map[string][]string{}
for _, p := range strings.Split(s, "&") {
if p == "" {
continue
}
kv := strings.SplitN(p, "=", 2)
k := strings.ReplaceAll(kv[0], "+", " ")
v := ""
if len(kv) > 1 {
v = strings.ReplaceAll(kv[1], "+", " ")
}
ku, err := urlQueryUnescape(k)
if err != nil {
return nil, err
}
vu, err := urlQueryUnescape(v)
if err != nil {
return nil, err
}
out[ku] = append(out[ku], vu)
}
return out, nil
}
var urlQueryUnescape = func(s string) (string, error) { return queryUnescape(s) }
func queryUnescape(s string) (string, error) { // minimal wrapper avoids exposing net/url in helpers
var b strings.Builder
for i := 0; i < len(s); i++ {
if s[i] == '%' && i+2 < len(s) {
n, err := strconv.ParseUint(s[i+1:i+3], 16, 8)
if err != nil {
return "", err
}
b.WriteByte(byte(n))
i += 2
} else {
b.WriteByte(s[i])
}
}
return b.String(), nil
}
func str(m map[string]any, keys ...string) string {
for _, k := range keys {
if v, ok := m[k]; ok {
switch x := v.(type) {
case string:
return x
case json.Number:
return x.String()
}
}
}
return ""
}
func intval(m map[string]any, k string) int {
v, ok := m[k]
if !ok {
return 0
}
switch x := v.(type) {
case float64:
return int(x)
case int:
return x
case string:
n, _ := strconv.Atoi(x)
return n
case json.Number:
n, _ := strconv.Atoi(x.String())
return n
}
return 0
}
func parsePriority(s string) int { n, _ := strconv.Atoi(s); return n }
func firstNonEmpty(v ...string) string {
for _, s := range v {
if s != "" {
return s
}
}
return ""
}
func writeJSON(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(v)
}
var loginTpl = template.Must(template.New("login").Parse(`<!doctype html><html lang="de"><head><meta charset="utf-8"><meta name="viewport" content="width=device-width"><title>Notify Gateway Login</title><style>body{font:16px system-ui;max-width:420px;margin:10vh auto;padding:1rem}input,button{width:100%;padding:.7rem;margin:.35rem 0;box-sizing:border-box}.err{color:#a00}</style></head><body><h1>Notify Gateway</h1>{{if .}}<p class="err">{{.}}</p>{{end}}<form method="post" action="/login"><input name="username" placeholder="Benutzer" autocomplete="username"><input type="password" name="password" placeholder="Passwort" autocomplete="current-password"><button>Anmelden</button></form></body></html>`))
func renderLogin(w http.ResponseWriter, msg string) {
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_ = loginTpl.Execute(w, msg)
}
File diff suppressed because one or more lines are too long
+287
View File
@@ -0,0 +1,287 @@
package mailingress
import (
"bytes"
"context"
"crypto/tls"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"sort"
"strings"
"sync"
"time"
"github.com/emersion/go-imap"
"github.com/emersion/go-imap/client"
_ "github.com/emersion/go-message/charset"
messageMail "github.com/emersion/go-message/mail"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/gateway"
"github.com/example/notify-gateway/internal/model"
)
const maxMailBytes = 1 << 20
type checkpoint struct {
Validity uint32 `json:"validity"`
UID uint32 `json:"uid"`
}
type Status struct {
ID string `json:"id"`
At time.Time `json:"at"`
Error string `json:"error,omitempty"`
}
type Poller struct {
Config *config.Store
Queue *gateway.Queue
mu sync.Mutex
statuses map[string]Status
}
func (p *Poller) Statuses() []Status {
p.mu.Lock()
defer p.mu.Unlock()
out := []Status{}
for _, s := range p.statuses {
out = append(out, s)
}
return out
}
func (p *Poller) Run(ctx context.Context) {
next := map[string]time.Time{}
ticker := time.NewTicker(time.Second)
defer ticker.Stop()
for {
for _, m := range p.Config.Get().Ingress.Mail {
if ctx.Err() != nil {
return
}
if !m.Enabled || time.Now().Before(next[m.ID]) {
continue
}
pollCtx, cancel := context.WithTimeout(ctx, 60*time.Second)
err := p.Poll(pollCtx, m)
cancel()
s := Status{ID: m.ID, At: time.Now().UTC()}
if err != nil {
s.Error = "Abruf fehlgeschlagen; Checkpoint bleibt vor der betroffenen Nachricht. Verbindung, MIME-Größe/Text und Zuordnung prüfen."
}
p.mu.Lock()
if p.statuses == nil {
p.statuses = map[string]Status{}
}
p.statuses[m.ID] = s
p.mu.Unlock()
next[m.ID] = time.Now().Add(time.Duration(max(m.PollSeconds, 10)) * time.Second)
}
select {
case <-ctx.Done():
return
case <-ticker.C:
}
}
}
func (p *Poller) Poll(ctx context.Context, m config.MailIngress) error {
host, _, err := net.SplitHostPort(m.Address)
if err != nil {
return err
}
return p.poll(ctx, m, &tls.Config{ServerName: host, MinVersion: tls.VersionTLS12})
}
func (p *Poller) poll(ctx context.Context, m config.MailIngress, tlsConfig *tls.Config) error {
password, err := config.Secret(m.Password)
if err != nil {
return err
}
raw, err := (&net.Dialer{Timeout: 15 * time.Second}).DialContext(ctx, "tcp", m.Address)
if err != nil {
return err
}
defer raw.Close()
stop := context.AfterFunc(ctx, func() { raw.Close() })
defer stop()
if err := raw.SetDeadline(time.Now().Add(30 * time.Second)); err != nil {
return err
}
conn := tls.Client(raw, tlsConfig)
if err := conn.HandshakeContext(ctx); err != nil {
return err
}
c, err := client.New(conn)
if err != nil {
return err
}
defer c.Terminate()
c.Timeout = 30 * time.Second
if err = c.Login(m.Username, password); err != nil {
return err
}
folder := m.Folder
if folder == "" {
folder = "INBOX"
}
mailbox, err := c.Select(folder, true)
if err != nil {
return err
}
if mailbox.UidValidity == 0 || mailbox.UidNext == 0 {
return errors.New("server did not provide stable mailbox UIDs")
}
cpID := gateway.ScopedKey("imap", m.ID, m.Address, m.Username, folder)
saved, err := p.Queue.Store.Checkpoint(ctx, cpID)
if err != nil {
return err
}
cp := checkpoint{}
if saved != "" {
if err := json.Unmarshal([]byte(saved), &cp); err != nil {
return err
}
}
if saved == "" && !m.ImportExisting {
cp = checkpoint{Validity: mailbox.UidValidity, UID: mailbox.UidNext - 1}
return p.save(ctx, cpID, cp)
}
if cp.Validity != mailbox.UidValidity {
cp = checkpoint{Validity: mailbox.UidValidity}
}
if cp.UID == ^uint32(0) {
return nil
}
criteria := imap.NewSearchCriteria()
criteria.Uid = new(imap.SeqSet)
criteria.Uid.AddRange(cp.UID+1, 0)
ids, err := c.UidSearch(criteria)
if err != nil {
return err
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
processed := 0
for _, uid := range ids {
// IMAP n:* may return the largest UID even when it is smaller than n.
if uid <= cp.UID {
continue
}
if processed >= 100 {
break
}
processed++
set := new(imap.SeqSet)
set.AddNum(uid)
section := &imap.BodySectionName{Peek: true, Partial: []int{0, maxMailBytes + 1}}
messages := make(chan *imap.Message, 1)
if err := c.UidFetch(set, []imap.FetchItem{section.FetchItem()}, messages); err != nil {
return err
}
item := <-messages
if item == nil {
return errors.New("message disappeared while fetching")
}
body := item.GetBody(section)
if body == nil {
return errors.New("missing IMAP body")
}
b, err := io.ReadAll(io.LimitReader(body, maxMailBytes+1))
if err != nil {
return err
}
if len(b) > maxMailBytes {
return errors.New("mail exceeds 1 MiB")
}
msg, skip, err := Parse(b, m)
if err != nil {
return err
}
if !skip {
key := gateway.ScopedKey(cpID, fmt.Sprint(cp.Validity), fmt.Sprint(uid))
if _, err := p.Queue.Accept(ctx, msg, key); err != nil {
return err
}
}
cp.UID = uid
if err := p.save(ctx, cpID, cp); err != nil {
return err
}
}
return nil
}
func (p *Poller) save(ctx context.Context, id string, cp checkpoint) error {
b, _ := json.Marshal(cp)
return p.Queue.Store.SetCheckpoint(ctx, id, string(b))
}
// Parse accepts only explicit plain-text MIME bodies; attachments are ignored.
// Sender headers are routing filters, not proof of sender identity.
func Parse(data []byte, c config.MailIngress) (model.InboundMessage, bool, error) {
if len(data) > maxMailBytes {
return model.InboundMessage{}, false, errors.New("mail exceeds 1 MiB")
}
r, err := messageMail.CreateReader(bytes.NewReader(data))
if err != nil {
return model.InboundMessage{}, false, err
}
defer r.Close()
if r.Header.Get("X-Notify-Gateway") != "" || r.Header.Get("Auto-Submitted") != "" && !strings.EqualFold(r.Header.Get("Auto-Submitted"), "no") {
return model.InboundMessage{}, true, nil
}
matches := func(field string, allowed []string) bool {
if len(allowed) == 0 {
return true
}
addresses, err := r.Header.AddressList(field)
if err != nil {
return false
}
for _, a := range addresses {
for _, v := range allowed {
if strings.EqualFold(a.Address, v) {
return true
}
}
}
return false
}
if !matches("From", c.From) || !(len(c.To) == 0 || matches("To", c.To) || matches("Cc", c.To)) {
return model.InboundMessage{}, true, nil
}
subject, err := r.Header.Subject()
if err != nil {
return model.InboundMessage{}, false, err
}
var text strings.Builder
for {
part, err := r.NextPart()
if err == io.EOF {
break
}
if err != nil {
return model.InboundMessage{}, false, err
}
if h, ok := part.Header.(*messageMail.InlineHeader); ok {
typ, _, err := h.ContentType()
if err != nil {
return model.InboundMessage{}, false, err
}
if typ == "text/plain" {
b, err := io.ReadAll(io.LimitReader(part.Body, maxMailBytes+1))
if err != nil {
return model.InboundMessage{}, false, err
}
if text.Len()+len(b) > maxMailBytes {
return model.InboundMessage{}, false, errors.New("decoded mail too large")
}
text.Write(b)
text.WriteByte('\n')
}
}
}
if strings.TrimSpace(text.String()) == "" {
return model.InboundMessage{}, false, errors.New("mail has no plain-text body")
}
return model.InboundMessage{Source: "mail", Channel: c.Channel, Title: subject, Message: strings.TrimSpace(text.String()), Raw: map[string]any{"from": r.Header.Get("From"), "to": r.Header.Get("To"), "message_id": r.Header.Get("Message-ID")}, ReceivedAt: time.Now().UTC()}, false, nil
}
+131
View File
@@ -0,0 +1,131 @@
package mailingress
import (
"context"
"crypto/tls"
"github.com/emersion/go-imap/backend/memory"
"github.com/emersion/go-imap/server"
"github.com/example/notify-gateway/internal/gateway"
"github.com/example/notify-gateway/internal/outbox"
"github.com/example/notify-gateway/internal/testutil"
"io"
"log"
"path/filepath"
"strings"
"testing"
"time"
"github.com/example/notify-gateway/internal/config"
)
func TestIMAPCheckpointsAndDurableAcceptance(t *testing.T) {
ctx := context.Background()
serverTLS, clientTLS := testutil.TLSConfigs()
listener, err := tls.Listen("tcp", "127.0.0.1:0", serverTLS)
if err != nil {
t.Fatal(err)
}
backend := memory.New()
srv := server.New(backend)
srv.ErrorLog = log.New(io.Discard, "", 0)
done := make(chan struct{})
go func() { defer close(done); srv.Serve(listener) }()
defer func() { srv.Close(); <-done }()
dir := t.TempDir()
cfg, _ := config.Open(filepath.Join(dir, "config.json"))
db, err := outbox.Open(filepath.Join(dir, "outbox.db"))
if err != nil {
t.Fatal(err)
}
defer db.Close()
c := cfg.Get()
c.Mappings = []config.Mapping{{ID: "mail", Enabled: true, Source: "mail", Target: "news", TextTemplate: "{{.Message}}"}}
if err := cfg.Replace(c); err != nil {
t.Fatal(err)
}
q := &gateway.Queue{Store: db, Dispatcher: gateway.New(cfg, nil)}
p := &Poller{Config: cfg, Queue: q}
m := config.MailIngress{ID: "mail", Address: listener.Addr().String(), Username: "username", Password: "password", Folder: "INBOX", Channel: "mail"}
// Existing UID 6 is skipped on first setup; mailbox remains read-only.
if err := p.poll(ctx, m, clientTLS); err != nil {
t.Fatal(err)
}
rows, _ := db.List(ctx, 50, 0)
if len(rows) != 0 {
t.Fatal("old mail imported unexpectedly")
}
cpID := gateway.ScopedKey("imap", m.ID, m.Address, m.Username, m.Folder)
old, _ := db.Checkpoint(ctx, cpID)
user, _ := backend.Login(nil, "username", "password")
mailbox, _ := user.GetMailbox("INBOX")
appendMail := func(body string) {
t.Helper()
if err := mailbox.CreateMessage(nil, time.Now(), strings.NewReader("From: sender@example.org\r\nSubject: Test\r\nContent-Type: text/plain\r\n\r\n"+body)); err != nil {
t.Fatal(err)
}
}
appendMail("new mail")
if err := p.poll(ctx, m, clientTLS); err != nil {
t.Fatal(err)
}
rows, _ = db.List(ctx, 50, 0)
if len(rows) != 1 {
t.Fatalf("new mail jobs=%d", len(rows))
}
// Simulate a crash after receipt commit but before checkpoint persistence.
if err := db.SetCheckpoint(ctx, cpID, old); err != nil {
t.Fatal(err)
}
if err := p.poll(ctx, m, clientTLS); err != nil {
t.Fatal(err)
}
rows, _ = db.List(ctx, 50, 0)
if len(rows) != 1 {
t.Fatal("duplicate IMAP delivery")
}
saved, _ := db.Checkpoint(ctx, cpID)
c.Mappings[0].Enabled = false
cfg.Replace(c)
appendMail("pending routing")
if err := p.poll(ctx, m, clientTLS); err == nil {
t.Fatal("missing mapping accepted")
}
after, _ := db.Checkpoint(ctx, cpID)
if after != saved {
t.Fatal("checkpoint advanced past unaccepted mail")
}
c.Mappings[0].Enabled = true
cfg.Replace(c)
if err := p.poll(ctx, m, clientTLS); err != nil {
t.Fatal(err)
}
rows, _ = db.List(ctx, 50, 0)
if len(rows) != 2 {
t.Fatal("mail was lost")
}
if err := p.Poll(ctx, m); err == nil {
t.Fatal("untrusted IMAP certificate accepted")
}
}
func TestMIMEFiltersAndLoopProtection(t *testing.T) {
raw := "From: Sender <sender@example.org>\r\nTo: ops@example.org\r\nSubject: =?UTF-8?Q?Gr=C3=BC=C3=9Fe?=\r\nMIME-Version: 1.0\r\nContent-Type: multipart/mixed; boundary=test\r\n\r\n--test\r\nContent-Type: text/plain; charset=utf-8\r\nContent-Transfer-Encoding: base64\r\n\r\nSGFsbG8=\r\n--test\r\nContent-Type: application/octet-stream\r\nContent-Disposition: attachment; filename=test.txt\r\n\r\nnot forwarded\r\n--test--\r\n"
cfg := config.MailIngress{Channel: "ops", From: []string{"sender@example.org"}, To: []string{"ops@example.org"}}
msg, skip, err := Parse([]byte(raw), cfg)
if err != nil || skip || msg.Title != "Grüße" || msg.Message != "Hallo" || msg.Source != "mail" {
t.Fatalf("%+v %v %v", msg, skip, err)
}
if _, skip, err := Parse([]byte("X-Notify-Gateway: 1\r\n"+raw), cfg); err != nil || !skip {
t.Fatalf("loop: %v %v", skip, err)
}
cfg.From = []string{"other@example.org"}
if _, skip, err := Parse([]byte(raw), cfg); err != nil || !skip {
t.Fatalf("sender filter: %v %v", skip, err)
}
if _, _, err := Parse([]byte(strings.Repeat("x", maxMailBytes+1)), cfg); err == nil {
t.Fatal("oversize accepted")
}
if _, _, err := Parse([]byte("Content-Type: text/html\r\n\r\n<b>hello</b>"), config.MailIngress{}); err == nil {
t.Fatal("HTML-only mail should require attention")
}
}
+25
View File
@@ -0,0 +1,25 @@
package model
import "time"
type InboundMessage struct {
Source string `json:"source"`
Channel string `json:"channel"`
Title string `json:"title"`
Message string `json:"message"`
Address string `json:"address,omitempty"`
Priority int `json:"priority,omitempty"`
Tags []string `json:"tags,omitempty"`
Raw map[string]any `json:"raw,omitempty"`
ReceivedAt time.Time `json:"received_at"`
}
type DeliveryResult struct {
OutboundID string `json:"outbound_id,omitempty"`
Error string `json:"error,omitempty"`
MappingID string `json:"mapping_id"`
MappingName string `json:"mapping_name"`
Target string `json:"target"`
StatusCode int `json:"status_code"`
Response string `json:"response,omitempty"`
}
+128
View File
@@ -0,0 +1,128 @@
// Package outbound implements HTTP delivery independently of ingress protocols.
package outbound
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"strconv"
"time"
"github.com/example/notify-gateway/internal/config"
)
type Response struct {
StatusCode int
Body string
Retryable bool
RetryAfter time.Duration
}
func Send(ctx context.Context, destination config.OutboundConfig, payload map[string]any) (Response, error) {
return SendWithID(ctx, destination, payload, "")
}
func SendWithID(ctx context.Context, destination config.OutboundConfig, payload map[string]any, id string) (Response, error) {
if err := config.ValidateOutbound(destination); err != nil {
return Response{}, err
}
b, err := json.Marshal(payload)
if err != nil {
return Response{}, fmt.Errorf("encode outbound payload: %w", err)
}
if !destination.Live {
preview, _ := json.Marshal(map[string]any{"dry_run": true, "provider": destination.Provider, "body": payload})
return Response{StatusCode: http.StatusOK, Body: string(preview)}, nil
}
if destination.Provider == "smtp" {
return sendSMTP(ctx, destination, payload, id)
}
token, err := config.Secret(destination.BearerToken)
if err != nil {
return Response{}, err
}
if destination.Provider == "ntfy" {
p := map[string]any{"topic": destination.Topic, "title": payload["title"], "message": payload["message"]}
priority := 0
switch v := payload["priority"].(type) {
case int:
priority = v
case float64:
priority = int(v)
}
if priority > 0 {
p["priority"] = min(priority, 5)
}
b, err = json.Marshal(p)
if err != nil {
return Response{}, err
}
}
u, _ := url.Parse(destination.URL)
if destination.Provider == "discord" {
q := u.Query()
q.Set("wait", "true")
u.RawQuery = q.Encode()
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, u.String(), bytes.NewReader(b))
if err != nil {
return Response{}, fmt.Errorf("outbound request could not be created")
}
req.Header.Set("Content-Type", "application/json")
if token != "" {
if destination.Provider == "gotify" {
req.Header.Set("X-Gotify-Key", token)
} else {
req.Header.Set("Authorization", "Bearer "+token)
}
}
if id != "" {
req.Header.Set("Idempotency-Key", id)
req.Header.Set("X-Notify-Gateway-Delivery", id)
}
timeout := destination.TimeoutS
if timeout == 0 {
timeout = 15
}
client := &http.Client{
Timeout: time.Duration(timeout) * time.Second,
// Never forward webhook credentials to a redirect destination.
CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse },
}
resp, err := client.Do(req)
if err != nil {
// url.Error embeds the URL, including Discord's secret token.
if ctx.Err() != nil {
return Response{Retryable: true}, ctx.Err()
}
return Response{Retryable: true}, fmt.Errorf("outbound HTTP request failed (connection or timeout)")
}
defer resp.Body.Close()
_, readErr := io.Copy(io.Discard, io.LimitReader(resp.Body, 64<<10))
out := Response{StatusCode: resp.StatusCode, Retryable: RetryableHTTP(resp.StatusCode), RetryAfter: ParseRetryAfter(resp.Header.Get("Retry-After"), time.Now())}
// Remote bodies may echo credentials or content. Only expose the status.
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return out, fmt.Errorf("outbound HTTP %d", resp.StatusCode)
}
if readErr != nil {
out.Retryable = true
return out, fmt.Errorf("outbound response could not be read")
}
return out, nil
}
func RetryableHTTP(code int) bool {
return code == 0 || code == 408 || code == 425 || code == 429 || code >= 500
}
func ParseRetryAfter(value string, now time.Time) time.Duration {
if n, err := strconv.ParseFloat(value, 64); err == nil && n > 0 && n <= 86400 {
return time.Duration(n * float64(time.Second))
}
if t, err := http.ParseTime(value); err == nil && t.After(now) {
return min(t.Sub(now), 24*time.Hour)
}
return 0
}
+118
View File
@@ -0,0 +1,118 @@
package outbound
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/example/notify-gateway/internal/config"
)
func TestWebhookDeliveryAndDryRun(t *testing.T) {
calls := 0
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
if r.Method != "POST" || r.Header.Get("Authorization") != "Bearer private-token" || r.Header.Get("Content-Type") != "application/json" {
t.Error("incorrect outbound request")
}
var payload map[string]any
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil || payload["message"] != "Grüße" {
t.Error("incorrect payload")
}
w.WriteHeader(http.StatusNoContent)
}))
defer srv.Close()
o := config.OutboundConfig{ID: "test", Provider: "webhook", URL: srv.URL, BearerToken: "private-token"}
r, err := Send(context.Background(), o, map[string]any{"message": "Grüße"})
if err != nil || calls != 0 || !strings.Contains(r.Body, `"dry_run":true`) || strings.Contains(r.Body, "private-token") {
t.Fatalf("dry-run: %+v %v calls=%d", r, err, calls)
}
o.Live = true
r, err = Send(context.Background(), o, map[string]any{"message": "Grüße"})
if err != nil || r.StatusCode != 204 || calls != 1 {
t.Fatalf("live: %+v %v calls=%d", r, err, calls)
}
}
func TestHTTPFailuresAndRedirectsDoNotLeakSecrets(t *testing.T) {
for _, status := range []int{302, 400, 429, 500} {
t.Run(fmt.Sprint(status), func(t *testing.T) {
calls := 0
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
w.Header().Set("Location", "/secret")
w.WriteHeader(status)
fmt.Fprint(w, "private-token")
}))
defer srv.Close()
r, err := Send(context.Background(), config.OutboundConfig{ID: "test", Provider: "webhook", URL: srv.URL + "/private-token", Live: true}, map[string]any{})
if err == nil || r.StatusCode != status || calls != 1 || r.Body != "" || strings.Contains(err.Error(), "private-token") {
t.Fatalf("result=%+v err=%v calls=%d", r, err, calls)
}
})
}
}
type transportFunc func(*http.Request) (*http.Response, error)
func (f transportFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }
func TestDiscordWaitAndTransportErrorRedaction(t *testing.T) {
previous := http.DefaultTransport
t.Cleanup(func() { http.DefaultTransport = previous })
http.DefaultTransport = transportFunc(func(r *http.Request) (*http.Response, error) {
if r.URL.Query().Get("wait") != "true" || r.URL.Query().Get("thread_id") != "123" {
t.Error("Discord query parameters missing")
}
return nil, fmt.Errorf("failure at %s", r.URL)
})
_, err := Send(context.Background(), config.OutboundConfig{ID: "discord", Provider: "discord", URL: "https://discord.com/api/webhooks/123/private-token?thread_id=123&wait=false", Live: true}, map[string]any{"content": "hello"})
if err == nil || strings.Contains(err.Error(), "private-token") {
t.Fatalf("unexpected error: %v", err)
}
}
func TestCanceledContext(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
_, err := Send(ctx, config.OutboundConfig{ID: "test", Provider: "webhook", URL: "http://127.0.0.1:1/private-token", Live: true}, map[string]any{})
if err != context.Canceled {
t.Fatalf("error=%v", err)
}
}
func TestNtfyAndGotifyProtocols(t *testing.T) {
for _, provider := range []string{"ntfy", "gotify"} {
t.Run(provider, func(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var payload map[string]any
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
t.Fatal(err)
}
if payload["title"] != "Test" || payload["message"] != "Hello" {
t.Error(payload)
}
if provider == "ntfy" {
if payload["topic"] != "ops" || payload["priority"] != float64(5) || r.Header.Get("Authorization") != "Bearer secret" {
t.Error("ntfy protocol", payload)
}
} else {
if r.Header.Get("X-Gotify-Key") != "secret" || r.Header.Get("Authorization") != "" {
t.Error("gotify auth")
}
}
w.WriteHeader(200)
}))
defer srv.Close()
t.Setenv("TEST_PROVIDER_TOKEN", "secret")
_, err := Send(context.Background(), config.OutboundConfig{ID: provider, Provider: provider, URL: srv.URL, Topic: "ops", BearerToken: "env:TEST_PROVIDER_TOKEN", Live: true}, map[string]any{"title": "Test", "message": "Hello", "priority": 9})
if err != nil {
t.Fatal(err)
}
})
}
}
+123
View File
@@ -0,0 +1,123 @@
package outbound
import (
"context"
"crypto/tls"
"encoding/base64"
"errors"
"fmt"
"mime"
"net"
"net/smtp"
"net/textproto"
"strconv"
"strings"
"time"
"github.com/example/notify-gateway/internal/config"
)
func mailBytes(o config.OutboundConfig, p map[string]any, id string) []byte {
title, _ := p["title"].(string)
message, _ := p["message"].(string)
title = strings.NewReplacer("\r", " ", "\n", " ").Replace(title)
// Force base64 for non-ASCII bodies and wrap per MIME line-length limits.
encoded := base64.StdEncoding.EncodeToString([]byte(message))
var body strings.Builder
for len(encoded) > 76 {
body.WriteString(encoded[:76] + "\r\n")
encoded = encoded[76:]
}
body.WriteString(encoded + "\r\n")
headers := []string{"From: " + o.From, "To: " + strings.Join(o.To, ", "), "Subject: " + mime.QEncoding.Encode("utf-8", title), "Date: " + time.Now().Format(time.RFC1123Z), "MIME-Version: 1.0", "Content-Type: text/plain; charset=utf-8", "Content-Transfer-Encoding: base64", "Auto-Submitted: auto-generated", "X-Notify-Gateway: 1"}
if id != "" {
headers = append(headers, "Message-ID: <"+id+"@notify-gateway.local>")
}
return []byte(strings.Join(headers, "\r\n") + "\r\n\r\n" + body.String())
}
func sendSMTP(ctx context.Context, o config.OutboundConfig, p map[string]any, id string) (resp Response, err error) {
return sendSMTPWithTLS(ctx, o, p, id, &tls.Config{ServerName: o.SMTPHost, MinVersion: tls.VersionTLS12})
}
func sendSMTPWithTLS(ctx context.Context, o config.OutboundConfig, p map[string]any, id string, tlsConfig *tls.Config) (resp Response, err error) {
defer func() {
if err != nil {
var e *textproto.Error
if errors.As(err, &e) {
resp.StatusCode = e.Code
resp.Retryable = e.Code >= 400 && e.Code < 500
} else {
resp.Retryable = true
}
err = errors.New("SMTP delivery failed")
}
}()
password, err := config.Secret(o.Password)
if err != nil {
return Response{}, err
}
timeout := o.TimeoutS
if timeout == 0 {
timeout = 15
}
dialer := net.Dialer{Timeout: time.Duration(timeout) * time.Second}
address := net.JoinHostPort(o.SMTPHost, strconv.Itoa(o.SMTPPort))
conn, err := dialer.DialContext(ctx, "tcp", address)
if err != nil {
return resp, err
}
defer conn.Close()
rawConn := conn
stop := context.AfterFunc(ctx, func() { rawConn.Close() })
defer stop()
if err = conn.SetDeadline(time.Now().Add(time.Duration(timeout) * time.Second)); err != nil {
return resp, err
}
if o.TLSMode == "tls" {
secure := tls.Client(conn, tlsConfig)
if err = secure.HandshakeContext(ctx); err != nil {
return resp, err
}
conn = secure
}
c, err := smtp.NewClient(conn, o.SMTPHost)
if err != nil {
return resp, err
}
defer c.Close()
if o.TLSMode == "starttls" {
if ok, _ := c.Extension("STARTTLS"); !ok {
return resp, fmt.Errorf("STARTTLS required")
}
if err = c.StartTLS(tlsConfig); err != nil {
return resp, err
}
}
if o.Username != "" {
if err = c.Auth(smtp.PlainAuth("", o.Username, password, o.SMTPHost)); err != nil {
return resp, err
}
}
if err = c.Mail(o.From); err != nil {
return resp, err
}
for _, to := range o.To {
if err = c.Rcpt(to); err != nil {
return resp, err
}
}
w, err := c.Data()
if err != nil {
return resp, err
}
if _, err = w.Write(mailBytes(o, p, id)); err != nil {
return resp, err
}
if err = w.Close(); err != nil {
return resp, err
}
// DATA has been acknowledged. A failed QUIT must not cause another delivery.
_ = c.Quit()
return Response{StatusCode: 250}, nil
}
+135
View File
@@ -0,0 +1,135 @@
package outbound
import (
"context"
"crypto/tls"
"encoding/base64"
"io"
"net"
"net/mail"
"net/textproto"
"strconv"
"strings"
"testing"
"github.com/example/notify-gateway/internal/config"
"github.com/example/notify-gateway/internal/testutil"
)
func smtpFixture(t *testing.T, mode string, rcptCode int) (config.OutboundConfig, *tls.Config, <-chan []byte) {
t.Helper()
serverTLS, clientTLS := testutil.TLSConfigs()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { listener.Close() })
received := make(chan []byte, 1)
done := make(chan struct{})
go func() {
defer close(done)
raw, err := listener.Accept()
if err != nil {
return
}
defer raw.Close()
var conn net.Conn = raw
if mode == "tls" {
conn = tls.Server(raw, serverTLS)
}
tp := textproto.NewConn(conn)
tp.PrintfLine("220 test SMTP")
for {
line, err := tp.ReadLine()
if err != nil {
return
}
verb := strings.Split(line, " ")[0]
switch verb {
case "EHLO":
if mode == "starttls" {
tp.PrintfLine("250-test\r\n250-STARTTLS\r\n250 AUTH PLAIN")
} else {
tp.PrintfLine("250-test\r\n250 AUTH PLAIN")
}
case "STARTTLS":
tp.PrintfLine("220 upgrade")
conn = tls.Server(raw, serverTLS)
tp = textproto.NewConn(conn)
case "AUTH":
data, _ := base64.StdEncoding.DecodeString(strings.TrimPrefix(line, "AUTH PLAIN "))
if string(data) != "\x00user\x00password" {
t.Error("unexpected SMTP auth")
}
tp.PrintfLine("235 OK")
case "MAIL":
tp.PrintfLine("250 OK")
case "RCPT":
tp.PrintfLine("%d recipient", rcptCode)
case "DATA":
tp.PrintfLine("354 data")
b, err := tp.ReadDotBytes()
if err != nil {
return
}
received <- b
tp.PrintfLine("250 accepted")
case "QUIT":
return // A lost QUIT response must still count as success.
default:
tp.PrintfLine("500 unsupported")
}
}
}()
t.Cleanup(func() { listener.Close(); <-done })
host, port, _ := net.SplitHostPort(listener.Addr().String())
p, _ := strconv.Atoi(port)
return config.OutboundConfig{ID: "mail", Provider: "smtp", SMTPHost: host, SMTPPort: p, TLSMode: mode, Username: "user", Password: "password", From: "sender@example.org", To: []string{"to@example.org"}, Live: true, TimeoutS: 2}, clientTLS, received
}
func TestSMTPEncryptedDeliveryAndMIME(t *testing.T) {
for _, mode := range []string{"tls", "starttls"} {
t.Run(mode, func(t *testing.T) {
o, tlsConfig, received := smtpFixture(t, mode, 250)
r, err := sendSMTPWithTLS(context.Background(), o, map[string]any{"title": "Grüße\r\nBcc: bad@example.org", "message": "Hallo Welt äöü"}, "stable-id", tlsConfig)
if err != nil || r.StatusCode != 250 {
t.Fatalf("%+v %v", r, err)
}
msg, err := mail.ReadMessage(strings.NewReader(string(<-received)))
if err != nil {
t.Fatal(err)
}
if msg.Header.Get("Bcc") != "" || msg.Header.Get("Message-ID") != "<stable-id@notify-gateway.local>" || msg.Header.Get("Auto-Submitted") != "auto-generated" {
t.Fatal("unsafe/missing headers")
}
body, err := io.ReadAll(base64.NewDecoder(base64.StdEncoding, msg.Body))
if err != nil || string(body) != "Hallo Welt äöü" {
t.Fatalf("body=%s %v", body, err)
}
})
}
}
func TestSMTPFailures(t *testing.T) {
for _, code := range []int{450, 550} {
t.Run(strconv.Itoa(code), func(t *testing.T) {
o, tlsConfig, _ := smtpFixture(t, "tls", code)
r, err := sendSMTPWithTLS(context.Background(), o, map[string]any{"message": "test"}, "id", tlsConfig)
if err == nil || r.StatusCode != code || r.Retryable != (code == 450) {
t.Fatalf("%+v %v", r, err)
}
})
}
t.Run("untrusted certificate", func(t *testing.T) {
o, _, _ := smtpFixture(t, "tls", 250)
if _, err := sendSMTP(context.Background(), o, map[string]any{"message": "test"}, "id"); err == nil {
t.Fatal("untrusted certificate accepted")
}
})
t.Run("no STARTTLS", func(t *testing.T) {
o, tlsConfig, _ := smtpFixture(t, "plain", 250)
o.TLSMode = "starttls"
if _, err := sendSMTPWithTLS(context.Background(), o, map[string]any{"message": "test"}, "id", tlsConfig); err == nil {
t.Fatal("plaintext SMTP accepted")
}
})
}
+281
View File
@@ -0,0 +1,281 @@
package outbox
import (
"context"
"crypto/rand"
"database/sql"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"time"
_ "modernc.org/sqlite"
)
var ErrConflict = errors.New("idempotency key already used for different content")
var ErrNotRetryable = errors.New("only dead deliveries can be retried")
type Store struct{ db *sql.DB }
type Job struct {
ID string `json:"id"`
ReceiptID string `json:"receipt_id"`
MappingID string `json:"mapping_id"`
Target string `json:"target"`
State string `json:"state"`
Attempts int `json:"attempts"`
Created int64 `json:"created_at"`
Next int64 `json:"next_attempt_at"`
Status int `json:"status_code"`
Error string `json:"error,omitempty"`
Data json.RawMessage `json:"-"`
Lease string `json:"-"`
}
type Receipt struct {
ID string `json:"id"`
Duplicate bool `json:"duplicate"`
}
type Attempt struct {
At int64 `json:"at"`
Status int `json:"status_code"`
Error string `json:"error,omitempty"`
State string `json:"state"`
}
func ID() string {
b := make([]byte, 16)
if _, err := rand.Read(b); err != nil {
panic(err)
}
return hex.EncodeToString(b)
}
func Open(path string) (*Store, error) {
if err := os.MkdirAll(filepath.Dir(path), 0700); err != nil {
return nil, err
}
// Create with restrictive permissions before SQLite opens the file.
f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR, 0600)
if err != nil {
return nil, err
}
f.Close()
if err := os.Chmod(path, 0600); err != nil {
return nil, err
}
db, err := sql.Open("sqlite", path)
if err != nil {
return nil, err
}
db.SetMaxOpenConns(1)
s := &Store{db: db}
fail := func(err error) (*Store, error) { db.Close(); return nil, err }
for _, q := range []string{"PRAGMA busy_timeout=2000", "PRAGMA journal_mode=WAL", "PRAGMA synchronous=FULL", "PRAGMA foreign_keys=ON"} {
if _, err := db.Exec(q); err != nil {
return fail(err)
}
}
var version int
if err := db.QueryRow("PRAGMA user_version").Scan(&version); err != nil {
return fail(err)
}
if version > 1 {
return fail(fmt.Errorf("unsupported outbox schema %d", version))
}
if version == 0 {
tx, err := db.Begin()
if err != nil {
return fail(err)
}
_, err = tx.Exec(`CREATE TABLE receipts(id TEXT PRIMARY KEY,scope TEXT NOT NULL,key TEXT NOT NULL,digest TEXT NOT NULL,created INTEGER NOT NULL,UNIQUE(scope,key));
CREATE TABLE deliveries(id TEXT PRIMARY KEY,receipt_id TEXT NOT NULL REFERENCES receipts(id),mapping_id TEXT NOT NULL,target TEXT NOT NULL,data BLOB NOT NULL,state TEXT NOT NULL,attempts INTEGER NOT NULL DEFAULT 0,created INTEGER NOT NULL,next INTEGER NOT NULL,status INTEGER NOT NULL DEFAULT 0,error TEXT NOT NULL DEFAULT '',lease TEXT NOT NULL DEFAULT '');
CREATE INDEX due_deliveries ON deliveries(state,next);
CREATE TABLE attempts(id INTEGER PRIMARY KEY,delivery_id TEXT NOT NULL REFERENCES deliveries(id),at INTEGER NOT NULL,status INTEGER NOT NULL,error TEXT NOT NULL,state TEXT NOT NULL);
CREATE TABLE checkpoints(id TEXT PRIMARY KEY,value TEXT NOT NULL);
PRAGMA user_version=1;`)
if err != nil {
tx.Rollback()
return fail(err)
}
if err := tx.Commit(); err != nil {
return fail(err)
}
}
return s, nil
}
func (s *Store) Close() error { return s.db.Close() }
func (s *Store) Ping(ctx context.Context) error { return s.db.PingContext(ctx) }
func (s *Store) Lookup(ctx context.Context, scope, key, digest string) (Receipt, error) {
var r Receipt
var old string
err := s.db.QueryRowContext(ctx, "SELECT id,digest FROM receipts WHERE scope=? AND key=?", scope, key).Scan(&r.ID, &old)
if errors.Is(err, sql.ErrNoRows) {
return Receipt{}, nil
}
if err != nil {
return r, err
}
if old != digest {
return Receipt{}, ErrConflict
}
r.Duplicate = true
return r, nil
}
// Enqueue commits the receipt and every destination together, before acceptance.
func (s *Store) Enqueue(ctx context.Context, scope, key, digest string, jobs []Job) (Receipt, error) {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return Receipt{}, err
}
defer tx.Rollback()
r := Receipt{ID: ID()}
now := time.Now().Unix()
result, err := tx.ExecContext(ctx, "INSERT INTO receipts(id,scope,key,digest,created) VALUES(?,?,?,?,?) ON CONFLICT(scope,key) DO NOTHING", r.ID, scope, key, digest, now)
if err != nil {
return Receipt{}, err
}
n, _ := result.RowsAffected()
if n == 0 {
var old string
if err := tx.QueryRowContext(ctx, "SELECT id,digest FROM receipts WHERE scope=? AND key=?", scope, key).Scan(&r.ID, &old); err != nil {
return Receipt{}, err
}
if old != digest {
return Receipt{}, ErrConflict
}
r.Duplicate = true
} else {
for _, j := range jobs {
if _, err := tx.ExecContext(ctx, "INSERT INTO deliveries(id,receipt_id,mapping_id,target,data,state,created,next) VALUES(?,?,?,?,?,'pending',?,?)", ID(), r.ID, j.MappingID, j.Target, []byte(j.Data), now, now); err != nil {
return Receipt{}, err
}
}
}
return r, tx.Commit()
}
func (s *Store) Claim(ctx context.Context, now time.Time) (*Job, error) {
j := new(Job)
lease := ID()
err := s.db.QueryRowContext(ctx, `UPDATE deliveries SET state='sending',attempts=attempts+1,next=?,lease=? WHERE id=(SELECT id FROM deliveries WHERE state IN ('pending','sending') AND next<=? ORDER BY next,created LIMIT 1) RETURNING id,receipt_id,mapping_id,target,data,state,attempts,created,next,status,error,lease`, now.Add(10*time.Minute).Unix(), lease, now.Unix()).Scan(&j.ID, &j.ReceiptID, &j.MappingID, &j.Target, &j.Data, &j.State, &j.Attempts, &j.Created, &j.Next, &j.Status, &j.Error, &j.Lease)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return j, err
}
func (s *Store) Finish(ctx context.Context, j Job, state string, status int, message string, next time.Time) error {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
r, err := tx.ExecContext(ctx, "UPDATE deliveries SET state=?,status=?,error=?,next=?,lease='' WHERE id=? AND lease=? AND state='sending'", state, status, message, next.Unix(), j.ID, j.Lease)
if err != nil {
return err
}
n, _ := r.RowsAffected()
if n != 1 {
return errors.New("delivery lease lost")
}
_, err = tx.ExecContext(ctx, "INSERT INTO attempts(delivery_id,at,status,error,state) VALUES(?,?,?,?,?)", j.ID, time.Now().Unix(), status, message, state)
if err != nil {
return err
}
return tx.Commit()
}
func (s *Store) Retry(ctx context.Context, id string) error {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
r, err := tx.ExecContext(ctx, "UPDATE deliveries SET state='pending',attempts=0,next=?,error='' WHERE id=? AND state='dead'", time.Now().Unix(), id)
if err != nil {
return err
}
n, _ := r.RowsAffected()
if n != 1 {
return ErrNotRetryable
}
_, err = tx.ExecContext(ctx, "INSERT INTO attempts(delivery_id,at,status,error,state) VALUES(?,?,0,'','manual_retry')", id, time.Now().Unix())
if err != nil {
return err
}
return tx.Commit()
}
func (s *Store) List(ctx context.Context, limit, offset int) ([]Job, error) {
if limit < 1 || limit > 200 {
limit = 50
}
if offset < 0 {
offset = 0
}
rows, err := s.db.QueryContext(ctx, "SELECT id,receipt_id,mapping_id,target,state,attempts,created,next,status,error FROM deliveries ORDER BY created DESC,rowid DESC LIMIT ? OFFSET ?", limit, offset)
if err != nil {
return nil, err
}
defer rows.Close()
out := []Job{}
for rows.Next() {
var j Job
if err := rows.Scan(&j.ID, &j.ReceiptID, &j.MappingID, &j.Target, &j.State, &j.Attempts, &j.Created, &j.Next, &j.Status, &j.Error); err != nil {
return nil, err
}
out = append(out, j)
}
return out, rows.Err()
}
func (s *Store) History(ctx context.Context, id string) ([]Attempt, error) {
rows, err := s.db.QueryContext(ctx, "SELECT at,status,error,state FROM attempts WHERE delivery_id=? ORDER BY id DESC LIMIT 200", id)
if err != nil {
return nil, err
}
defer rows.Close()
out := []Attempt{}
for rows.Next() {
var a Attempt
if err := rows.Scan(&a.At, &a.Status, &a.Error, &a.State); err != nil {
return nil, err
}
out = append(out, a)
}
return out, rows.Err()
}
func (s *Store) Counts(ctx context.Context) (map[string]int64, error) {
out := map[string]int64{"pending": 0, "sending": 0, "succeeded": 0, "dry_run": 0, "dead": 0}
rows, err := s.db.QueryContext(ctx, "SELECT state,count(*) FROM deliveries GROUP BY state")
if err != nil {
return nil, err
}
for rows.Next() {
var state string
var n int64
if err := rows.Scan(&state, &n); err != nil {
rows.Close()
return nil, err
}
out[state] = n
}
err = rows.Err()
rows.Close()
return out, err
}
func (s *Store) Checkpoint(ctx context.Context, id string) (string, error) {
var v string
err := s.db.QueryRowContext(ctx, "SELECT value FROM checkpoints WHERE id=?", id).Scan(&v)
if errors.Is(err, sql.ErrNoRows) {
return "", nil
}
return v, err
}
func (s *Store) SetCheckpoint(ctx context.Context, id, value string) error {
_, err := s.db.ExecContext(ctx, "INSERT INTO checkpoints(id,value) VALUES(?,?) ON CONFLICT(id) DO UPDATE SET value=excluded.value", id, value)
return err
}
+100
View File
@@ -0,0 +1,100 @@
package outbox
import (
"context"
"encoding/json"
"errors"
"path/filepath"
"sync"
"testing"
"time"
)
func TestPersistenceDeduplicationAndLeaseRecovery(t *testing.T) {
ctx := context.Background()
path := filepath.Join(t.TempDir(), "outbox.db")
s, err := Open(path)
if err != nil {
t.Fatal(err)
}
jobs := []Job{{MappingID: "one", Target: "webhook", Data: json.RawMessage(`{"payload":"test"}`)}, {MappingID: "two", Target: "smtp", Data: json.RawMessage(`{}`)}}
receipt, err := s.Enqueue(ctx, "source", "key", "digest", jobs)
if err != nil {
t.Fatal(err)
}
duplicate, err := s.Enqueue(ctx, "source", "key", "digest", jobs)
if err != nil || !duplicate.Duplicate || duplicate.ID != receipt.ID {
t.Fatalf("duplicate=%+v err=%v", duplicate, err)
}
if _, err := s.Enqueue(ctx, "source", "key", "different", jobs); !errors.Is(err, ErrConflict) {
t.Fatalf("conflict=%v", err)
}
j, err := s.Claim(ctx, time.Now())
if err != nil || j == nil {
t.Fatalf("claim: %v", err)
}
if err := s.Finish(ctx, *j, "succeeded", 204, "", time.Now()); err != nil {
t.Fatal(err)
}
if err := s.Retry(ctx, j.ID); !errors.Is(err, ErrNotRetryable) {
t.Fatal("successful delivery was retryable")
}
orphan, err := s.Claim(ctx, time.Now())
if err != nil || orphan == nil {
t.Fatal(err)
}
s.Close()
s, err = Open(path)
if err != nil {
t.Fatal(err)
}
defer s.Close()
if j, err := s.Claim(ctx, time.Now()); err != nil || j != nil {
t.Fatalf("active lease should remain protected: %+v %v", j, err)
}
recovered, err := s.Claim(ctx, time.Now().Add(11*time.Minute))
if err != nil || recovered == nil || recovered.ID != orphan.ID || recovered.Attempts != 2 {
t.Fatalf("recovery=%+v %v", recovered, err)
}
if err := s.Finish(ctx, *orphan, "succeeded", 200, "", time.Now()); err == nil {
t.Fatal("stale worker completed a new lease")
}
if err := s.Finish(ctx, *recovered, "dead", 400, "bad request", time.Now()); err != nil {
t.Fatal(err)
}
if err := s.Retry(ctx, recovered.ID); err != nil {
t.Fatal(err)
}
history, err := s.History(ctx, recovered.ID)
if err != nil || len(history) != 2 || history[0].State != "manual_retry" {
t.Fatalf("history=%+v %v", history, err)
}
all, err := s.List(ctx, 50, 0)
if err != nil || len(all) != 2 {
t.Fatalf("jobs=%+v %v", all, err)
}
}
func TestConcurrentAcceptanceIsAtomic(t *testing.T) {
s, err := Open(filepath.Join(t.TempDir(), "outbox.db"))
if err != nil {
t.Fatal(err)
}
defer s.Close()
var wg sync.WaitGroup
for i := 0; i < 12; i++ {
wg.Add(1)
go func() {
defer wg.Done()
_, err := s.Enqueue(context.Background(), "scope", "one", "digest", []Job{{MappingID: "a", Target: "webhook", Data: json.RawMessage(`{}`)}})
if err != nil {
t.Error(err)
}
}()
}
wg.Wait()
rows, err := s.List(context.Background(), 50, 0)
if err != nil || len(rows) != 1 {
t.Fatalf("rows=%+v %v", rows, err)
}
}
+17
View File
@@ -0,0 +1,17 @@
// Package testutil provides local fixtures for integration tests only.
package testutil
import (
"crypto/tls"
"crypto/x509"
"net/http"
"net/http/httptest"
)
func TLSConfigs() (*tls.Config, *tls.Config) {
srv := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {}))
defer srv.Close()
roots := x509.NewCertPool()
roots.AddCert(srv.Certificate())
return &tls.Config{Certificates: srv.TLS.Certificates, MinVersion: tls.VersionTLS12}, &tls.Config{RootCAs: roots, ServerName: "127.0.0.1", MinVersion: tls.VersionTLS12}
}