feat(canary): event + notify domain — contracts, senders, services
Lays the domain layer for Phase 10 notification dispatch.
- event/contract.go: NotifyInfo (token-shape DTO that breaks the would-be
event→notify→token cycle), Notifier interface, TokenIncrementer interface,
Store interface (Insert+UpdateNotifyStatus+PruneToLimit).
- token.Token.NotifyInfo() helper bridges *token.Token → event.NotifyInfo
for the main.go adapter.
- notify/types.go: Sender + StatusWriter interfaces.
- notify.Service: per-channel sender registry, async fire-and-forget
Notify with bounded sendTimeout, WaitGroup for graceful drain, status
writeback (sent / failed) on the configured StatusWriter.
- telegram.Sender: POST /bot{TOKEN}/sendMessage with MarkdownV2-escaped
body, cenkalti/backoff/v5 retry (3 tries / 30s window), permanent on
4xx, 10s overall + 5s connect timeouts. Note: spec body said
parse_mode=Markdown but escape table is the V2 set; using MarkdownV2
keeps the escape rules consistent.
- webhook.Sender: POST user URL with versioned JSON envelope (§10.4),
optional HMAC-SHA256 signing via X-Canary-Signature, URL validation
at the call site (rejects non-http(s), missing host, userinfo). Same
retry/timeout policy as telegram.
- event.Service: Record(ctx, info, evt) inserts → IncrementTriggerCount
→ Redis SetNX dedup gate (15m TTL, fail-open on Redis error). First
trigger calls notifier; duplicates INCR + UpdateNotifyStatus(deduped).
- event.Service.RunRetentionLoop(ctx, interval, limit) tickered prune.
Adds cenkalti/backoff/v5 + miniredis/v2 to go.mod for backoff API and
unit-test Redis fakes.
This commit is contained in:
parent
1fe61345ef
commit
0757d9f196
|
|
@ -31,7 +31,9 @@ require (
|
|||
dario.cat/mergo v1.0.2 // indirect
|
||||
github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c // indirect
|
||||
github.com/Microsoft/go-winio v0.6.2 // indirect
|
||||
github.com/alicebob/miniredis/v2 v2.38.0 // indirect
|
||||
github.com/cenkalti/backoff/v4 v4.3.0 // indirect
|
||||
github.com/cenkalti/backoff/v5 v5.0.3 // indirect
|
||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||
github.com/clipperhouse/uax29/v2 v2.7.0 // indirect
|
||||
github.com/containerd/errdefs v1.0.0 // indirect
|
||||
|
|
@ -89,6 +91,7 @@ require (
|
|||
github.com/sirupsen/logrus v1.9.4 // indirect
|
||||
github.com/tklauser/go-sysconf v0.3.16 // indirect
|
||||
github.com/tklauser/numcpus v0.11.0 // indirect
|
||||
github.com/yuin/gopher-lua v1.1.1 // indirect
|
||||
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
||||
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.68.0 // indirect
|
||||
|
|
|
|||
|
|
@ -9,12 +9,16 @@ github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c h1:udKWzYgxTojEK
|
|||
github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
|
||||
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
||||
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
|
||||
github.com/alicebob/miniredis/v2 v2.38.0 h1:nZAzCR+Lj+Vxk4ZXzm2NuKq2O33RXj1XxJ2e2uP9jiw=
|
||||
github.com/alicebob/miniredis/v2 v2.38.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM=
|
||||
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
||||
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
||||
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
||||
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
|
||||
github.com/cenkalti/backoff/v4 v4.3.0 h1:MyRJ/UdXutAwSAT+s3wNd7MfTIcy71VQueUuFK343L8=
|
||||
github.com/cenkalti/backoff/v4 v4.3.0/go.mod h1:Y3VNntkOUPxTVeUxJ/G5vcM//AlwfmyYozVcomhLiZE=
|
||||
github.com/cenkalti/backoff/v5 v5.0.3 h1:ZN+IMa753KfX5hd8vVaMixjnqRZ3y8CuJKRKj1xcsSM=
|
||||
github.com/cenkalti/backoff/v5 v5.0.3/go.mod h1:rkhZdG3JZukswDf7f0cwqPNk4K0sa+F97BxZthm/crw=
|
||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/clipperhouse/uax29/v2 v2.7.0 h1:+gs4oBZ2gPfVrKPthwbMzWZDaAFPGYK72F0NJv2v7Vk=
|
||||
|
|
@ -200,6 +204,8 @@ github.com/tklauser/go-sysconf v0.3.16 h1:frioLaCQSsF5Cy1jgRBrzr6t502KIIwQ0MArYI
|
|||
github.com/tklauser/go-sysconf v0.3.16/go.mod h1:/qNL9xxDhc7tx3HSRsLWNnuzbVfh3e7gh/BmM179nYI=
|
||||
github.com/tklauser/numcpus v0.11.0 h1:nSTwhKH5e1dMNsCdVBukSZrURJRoHbSEQjdEbY+9RXw=
|
||||
github.com/tklauser/numcpus v0.11.0/go.mod h1:z+LwcLq54uWZTX0u/bGobaV34u6V7KNlTZejzM6/3MQ=
|
||||
github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M=
|
||||
github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw=
|
||||
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
||||
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
||||
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
|
||||
|
|
|
|||
|
|
@ -0,0 +1,39 @@
|
|||
// ©AngelaMos | 2026
|
||||
// contract.go
|
||||
|
||||
package event
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
type NotifyInfo struct {
|
||||
TokenID string
|
||||
ManageID string
|
||||
Type string
|
||||
Memo string
|
||||
AlertChannel string
|
||||
TelegramBot string
|
||||
TelegramChat string
|
||||
WebhookURL string
|
||||
}
|
||||
|
||||
type Notifier interface {
|
||||
Notify(info NotifyInfo, evt *Event)
|
||||
}
|
||||
|
||||
type TokenIncrementer interface {
|
||||
IncrementTriggerCount(ctx context.Context, id string) error
|
||||
}
|
||||
|
||||
type Store interface {
|
||||
Insert(ctx context.Context, e *Event) error
|
||||
UpdateNotifyStatus(
|
||||
ctx context.Context,
|
||||
eventID int64,
|
||||
status NotifyStatus,
|
||||
sentAt *time.Time,
|
||||
) error
|
||||
PruneToLimit(ctx context.Context, perTokenLimit int) (int64, error)
|
||||
}
|
||||
|
|
@ -0,0 +1,155 @@
|
|||
// ©AngelaMos | 2026
|
||||
// service.go
|
||||
|
||||
package event
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
const (
|
||||
dedupKeyPrefix = "dedup:trigger:"
|
||||
defaultDedupTTL = 15 * time.Minute
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
repo Store
|
||||
tokens TokenIncrementer
|
||||
rdb *redis.Client
|
||||
notifier Notifier
|
||||
dedupTTL time.Duration
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
type ServiceConfig struct {
|
||||
DedupTTL time.Duration
|
||||
Logger *slog.Logger
|
||||
}
|
||||
|
||||
func NewService(
|
||||
repo Store,
|
||||
tokens TokenIncrementer,
|
||||
rdb *redis.Client,
|
||||
notifier Notifier,
|
||||
cfg ServiceConfig,
|
||||
) *Service {
|
||||
if cfg.DedupTTL <= 0 {
|
||||
cfg.DedupTTL = defaultDedupTTL
|
||||
}
|
||||
if cfg.Logger == nil {
|
||||
cfg.Logger = slog.Default()
|
||||
}
|
||||
return &Service{
|
||||
repo: repo,
|
||||
tokens: tokens,
|
||||
rdb: rdb,
|
||||
notifier: notifier,
|
||||
dedupTTL: cfg.DedupTTL,
|
||||
logger: cfg.Logger,
|
||||
}
|
||||
}
|
||||
|
||||
func DedupKey(tokenID, sourceIP string) string {
|
||||
return dedupKeyPrefix + tokenID + ":" + sourceIP
|
||||
}
|
||||
|
||||
func (s *Service) Record(
|
||||
ctx context.Context,
|
||||
info NotifyInfo,
|
||||
evt *Event,
|
||||
) error {
|
||||
if err := s.repo.Insert(ctx, evt); err != nil {
|
||||
return fmt.Errorf("insert event: %w", err)
|
||||
}
|
||||
|
||||
if s.tokens != nil {
|
||||
if err := s.tokens.IncrementTriggerCount(
|
||||
ctx,
|
||||
info.TokenID,
|
||||
); err != nil {
|
||||
s.logger.WarnContext(ctx, "increment trigger count",
|
||||
"error", err, "token_id", info.TokenID)
|
||||
}
|
||||
}
|
||||
|
||||
first := s.dedupGate(ctx, info.TokenID, evt.SourceIP)
|
||||
if !first {
|
||||
if err := s.repo.UpdateNotifyStatus(
|
||||
ctx, evt.ID, NotifyDeduped, nil,
|
||||
); err != nil {
|
||||
s.logger.WarnContext(ctx, "update notify status deduped",
|
||||
"error", err, "event_id", evt.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if s.notifier != nil {
|
||||
s.notifier.Notify(info, evt)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) dedupGate(
|
||||
ctx context.Context,
|
||||
tokenID, sourceIP string,
|
||||
) bool {
|
||||
if s.rdb == nil {
|
||||
return true
|
||||
}
|
||||
key := DedupKey(tokenID, sourceIP)
|
||||
set, err := s.rdb.SetNX(ctx, key, 1, s.dedupTTL).Result()
|
||||
if err != nil {
|
||||
s.logger.WarnContext(ctx, "dedup setnx failed (fail-open)",
|
||||
"error", err, "key", key)
|
||||
return true
|
||||
}
|
||||
if set {
|
||||
return true
|
||||
}
|
||||
if _, iErr := s.rdb.Incr(ctx, key).Result(); iErr != nil {
|
||||
s.logger.WarnContext(ctx, "dedup incr failed",
|
||||
"error", iErr, "key", key)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *Service) RunRetentionLoop(
|
||||
ctx context.Context,
|
||||
interval time.Duration,
|
||||
perTokenLimit int,
|
||||
) {
|
||||
if interval <= 0 || perTokenLimit <= 0 {
|
||||
s.logger.WarnContext(ctx, "retention loop disabled (invalid config)",
|
||||
"interval", interval, "limit", perTokenLimit)
|
||||
return
|
||||
}
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
s.logger.InfoContext(ctx, "retention loop started",
|
||||
"interval", interval, "per_token_limit", perTokenLimit)
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
s.logger.InfoContext(ctx, "retention loop stopped")
|
||||
return
|
||||
case <-ticker.C:
|
||||
n, err := s.repo.PruneToLimit(ctx, perTokenLimit)
|
||||
if err != nil {
|
||||
s.logger.WarnContext(ctx, "retention loop: prune failed",
|
||||
"error", err, "per_token_limit", perTokenLimit)
|
||||
continue
|
||||
}
|
||||
if n > 0 {
|
||||
s.logger.InfoContext(ctx, "retention loop: pruned events",
|
||||
"deleted", n, "per_token_limit", perTokenLimit)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,576 @@
|
|||
// ©AngelaMos | 2026
|
||||
// service_test.go
|
||||
|
||||
package event_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strconv"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||
)
|
||||
|
||||
const testTokenID = "tokevtsvc001"
|
||||
|
||||
type fakeStore struct {
|
||||
mu sync.Mutex
|
||||
inserted []*event.Event
|
||||
insertErr error
|
||||
statusUpdates []statusUpdate
|
||||
statusErr error
|
||||
pruneCount int64
|
||||
pruneErr error
|
||||
pruneLastLimit int
|
||||
}
|
||||
|
||||
type statusUpdate struct {
|
||||
id int64
|
||||
status event.NotifyStatus
|
||||
sentAt *time.Time
|
||||
}
|
||||
|
||||
func (f *fakeStore) Insert(_ context.Context, e *event.Event) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
if f.insertErr != nil {
|
||||
return f.insertErr
|
||||
}
|
||||
e.ID = int64(len(f.inserted) + 1)
|
||||
if e.TriggeredAt.IsZero() {
|
||||
e.TriggeredAt = time.Now().UTC()
|
||||
}
|
||||
if e.NotifyStatus == "" {
|
||||
e.NotifyStatus = event.NotifyPending
|
||||
}
|
||||
f.inserted = append(f.inserted, e)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeStore) UpdateNotifyStatus(
|
||||
_ context.Context,
|
||||
id int64,
|
||||
status event.NotifyStatus,
|
||||
sentAt *time.Time,
|
||||
) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
if f.statusErr != nil {
|
||||
return f.statusErr
|
||||
}
|
||||
f.statusUpdates = append(f.statusUpdates, statusUpdate{id, status, sentAt})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeStore) PruneToLimit(
|
||||
_ context.Context,
|
||||
perTokenLimit int,
|
||||
) (int64, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.pruneLastLimit = perTokenLimit
|
||||
if f.pruneErr != nil {
|
||||
return 0, f.pruneErr
|
||||
}
|
||||
return f.pruneCount, nil
|
||||
}
|
||||
|
||||
func (f *fakeStore) snapshot() ([]*event.Event, []statusUpdate) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
ev := make([]*event.Event, len(f.inserted))
|
||||
copy(ev, f.inserted)
|
||||
su := make([]statusUpdate, len(f.statusUpdates))
|
||||
copy(su, f.statusUpdates)
|
||||
return ev, su
|
||||
}
|
||||
|
||||
type fakeIncrementer struct {
|
||||
mu sync.Mutex
|
||||
calls []string
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *fakeIncrementer) IncrementTriggerCount(
|
||||
_ context.Context,
|
||||
id string,
|
||||
) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.calls = append(f.calls, id)
|
||||
return f.err
|
||||
}
|
||||
|
||||
func (f *fakeIncrementer) callCount() int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return len(f.calls)
|
||||
}
|
||||
|
||||
type fakeNotifier struct {
|
||||
mu sync.Mutex
|
||||
calls []notifyCall
|
||||
}
|
||||
|
||||
type notifyCall struct {
|
||||
info event.NotifyInfo
|
||||
evt *event.Event
|
||||
}
|
||||
|
||||
func (f *fakeNotifier) Notify(info event.NotifyInfo, evt *event.Event) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.calls = append(f.calls, notifyCall{info, evt})
|
||||
}
|
||||
|
||||
func (f *fakeNotifier) callCount() int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return len(f.calls)
|
||||
}
|
||||
|
||||
func setupRedis(t *testing.T) (*redis.Client, *miniredis.Miniredis) {
|
||||
t.Helper()
|
||||
mr, err := miniredis.Run()
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(mr.Close)
|
||||
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
||||
t.Cleanup(func() {
|
||||
if cErr := rdb.Close(); cErr != nil {
|
||||
t.Logf("redis close: %v", cErr)
|
||||
}
|
||||
})
|
||||
return rdb, mr
|
||||
}
|
||||
|
||||
func sampleInfo() event.NotifyInfo {
|
||||
return event.NotifyInfo{
|
||||
TokenID: testTokenID,
|
||||
ManageID: "abcd-1234",
|
||||
Type: "webbug",
|
||||
Memo: "test",
|
||||
AlertChannel: "telegram",
|
||||
TelegramBot: "bot",
|
||||
TelegramChat: "chat",
|
||||
}
|
||||
}
|
||||
|
||||
func sampleEvent(ip string) *event.Event {
|
||||
return &event.Event{TokenID: testTokenID, SourceIP: ip}
|
||||
}
|
||||
|
||||
func newSvc(
|
||||
t *testing.T,
|
||||
store event.Store,
|
||||
tokens event.TokenIncrementer,
|
||||
rdb *redis.Client,
|
||||
notifier event.Notifier,
|
||||
) *event.Service {
|
||||
t.Helper()
|
||||
return event.NewService(store, tokens, rdb, notifier, event.ServiceConfig{
|
||||
DedupTTL: 15 * time.Minute,
|
||||
Logger: slog.New(slog.NewTextHandler(testWriter{t}, nil)),
|
||||
})
|
||||
}
|
||||
|
||||
type testWriter struct{ t *testing.T }
|
||||
|
||||
func (w testWriter) Write(
|
||||
p []byte,
|
||||
) (int, error) {
|
||||
w.t.Log(string(p))
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func TestService_Record_InsertsEvent(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, _ := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := newSvc(t, store, inc, rdb, notifier)
|
||||
|
||||
evt := sampleEvent("203.0.113.1")
|
||||
require.NoError(t, svc.Record(context.Background(), sampleInfo(), evt))
|
||||
|
||||
inserted, _ := store.snapshot()
|
||||
require.Len(t, inserted, 1)
|
||||
require.Equal(t, "203.0.113.1", inserted[0].SourceIP)
|
||||
require.NotZero(t, evt.ID, "Insert assigns ID")
|
||||
}
|
||||
|
||||
func TestService_Record_IncrementsTriggerCount(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, _ := setupRedis(t)
|
||||
svc := newSvc(t, store, inc, rdb, nil)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
)
|
||||
require.Equal(t, 1, inc.callCount())
|
||||
}
|
||||
|
||||
func TestService_Record_FirstTriggerNotifies(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, _ := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := newSvc(t, store, inc, rdb, notifier)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
)
|
||||
require.Equal(t, 1, notifier.callCount())
|
||||
|
||||
_, statusUpdates := store.snapshot()
|
||||
require.Empty(
|
||||
t,
|
||||
statusUpdates,
|
||||
"first trigger should not write 'deduped' status; notify.Service handles sent/failed writeback async",
|
||||
)
|
||||
}
|
||||
|
||||
func TestService_Record_DuplicateMarksDeduped(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, mr := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := newSvc(t, store, inc, rdb, notifier)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
)
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
)
|
||||
|
||||
require.Equal(t, 1, notifier.callCount(), "duplicate must not notify")
|
||||
|
||||
inserted, statusUpdates := store.snapshot()
|
||||
require.Len(t, inserted, 2, "both events still recorded")
|
||||
require.Len(t, statusUpdates, 1, "duplicate writes deduped status")
|
||||
require.Equal(t, event.NotifyDeduped, statusUpdates[0].status)
|
||||
require.Equal(t, inserted[1].ID, statusUpdates[0].id)
|
||||
|
||||
dedupKey := "dedup:trigger:" + testTokenID + ":203.0.113.1"
|
||||
val, err := mr.Get(dedupKey)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "2", val, "INCR bumps counter")
|
||||
}
|
||||
|
||||
func TestService_Record_DifferentIPsBothNotify(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, _ := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := newSvc(t, store, inc, rdb, notifier)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
)
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.2"),
|
||||
),
|
||||
)
|
||||
|
||||
require.Equal(t, 2, notifier.callCount(), "different IPs each notify")
|
||||
}
|
||||
|
||||
func TestService_Record_DedupTTLExpiry(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, mr := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := event.NewService(store, inc, rdb, notifier, event.ServiceConfig{
|
||||
DedupTTL: 1 * time.Second,
|
||||
Logger: slog.New(slog.NewTextHandler(testWriter{t}, nil)),
|
||||
})
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
)
|
||||
mr.FastForward(2 * time.Second)
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
)
|
||||
|
||||
require.Equal(
|
||||
t,
|
||||
2,
|
||||
notifier.callCount(),
|
||||
"after TTL expiry second trigger notifies again",
|
||||
)
|
||||
}
|
||||
|
||||
func TestService_Record_DedupKeyShape(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, mr := setupRedis(t)
|
||||
svc := newSvc(t, store, inc, rdb, nil)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.99"),
|
||||
),
|
||||
)
|
||||
|
||||
keys := mr.Keys()
|
||||
require.Contains(t, keys, "dedup:trigger:"+testTokenID+":203.0.113.99")
|
||||
}
|
||||
|
||||
func TestService_Record_RedisDownFailsOpen(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, mr := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := newSvc(t, store, inc, rdb, notifier)
|
||||
|
||||
mr.Close()
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
)
|
||||
require.Equal(t, 1, notifier.callCount(),
|
||||
"redis down → fail open → still notify so we don't miss alerts")
|
||||
}
|
||||
|
||||
func TestService_Record_InsertErrorReturns(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{insertErr: errors.New("db down")}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, _ := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := newSvc(t, store, inc, rdb, notifier)
|
||||
|
||||
err := svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.Equal(
|
||||
t,
|
||||
0,
|
||||
notifier.callCount(),
|
||||
"insert failure prevents notify — we don't have an event id to write back to",
|
||||
)
|
||||
}
|
||||
|
||||
func TestService_Record_IncrementErrorDoesNotPropagate(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{err: errors.New("update failed")}
|
||||
rdb, _ := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := newSvc(t, store, inc, rdb, notifier)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
),
|
||||
"increment failure is best-effort; record should still succeed",
|
||||
)
|
||||
require.Equal(t, 1, notifier.callCount())
|
||||
}
|
||||
|
||||
func TestService_Record_NilNotifierNoCrash(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, _ := setupRedis(t)
|
||||
svc := newSvc(t, store, inc, rdb, nil)
|
||||
|
||||
require.NotPanics(t, func() {
|
||||
if err := svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent("203.0.113.1"),
|
||||
); err != nil {
|
||||
t.Logf("record: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_Record_ConcurrentSafe(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
inc := &fakeIncrementer{}
|
||||
rdb, _ := setupRedis(t)
|
||||
notifier := &fakeNotifier{}
|
||||
svc := newSvc(t, store, inc, rdb, notifier)
|
||||
|
||||
const n = 20
|
||||
var wg sync.WaitGroup
|
||||
var notified atomic.Int32
|
||||
for i := range n {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
ip := "203.0.113." + strconv.Itoa(i+1)
|
||||
if err := svc.Record(
|
||||
context.Background(),
|
||||
sampleInfo(),
|
||||
sampleEvent(ip),
|
||||
); err == nil {
|
||||
notified.Add(1)
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
require.Equal(t, int32(n), notified.Load())
|
||||
require.Equal(t, n, notifier.callCount())
|
||||
}
|
||||
|
||||
func TestService_RunRetentionLoop_PrunesAtInterval(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{pruneCount: 5}
|
||||
rdb, _ := setupRedis(t)
|
||||
svc := newSvc(t, store, &fakeIncrementer{}, rdb, nil)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
svc.RunRetentionLoop(ctx, 25*time.Millisecond, 100)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
require.Eventually(t, func() bool {
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
return store.pruneLastLimit == 100
|
||||
}, 1*time.Second, 5*time.Millisecond)
|
||||
|
||||
cancel()
|
||||
<-done
|
||||
}
|
||||
|
||||
func TestService_RunRetentionLoop_StopsOnContextCancel(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
rdb, _ := setupRedis(t)
|
||||
svc := newSvc(t, store, &fakeIncrementer{}, rdb, nil)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
svc.RunRetentionLoop(ctx, 10*time.Millisecond, 50)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(500 * time.Millisecond):
|
||||
t.Fatal("retention loop did not stop on cancel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_RunRetentionLoop_ContinuesOnPruneError(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{pruneErr: errors.New("db down")}
|
||||
rdb, _ := setupRedis(t)
|
||||
svc := newSvc(t, store, &fakeIncrementer{}, rdb, nil)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
svc.RunRetentionLoop(ctx, 10*time.Millisecond, 50)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
time.Sleep(60 * time.Millisecond)
|
||||
cancel()
|
||||
<-done
|
||||
}
|
||||
|
||||
func TestService_RunRetentionLoop_DisabledOnInvalidConfig(t *testing.T) {
|
||||
t.Parallel()
|
||||
store := &fakeStore{}
|
||||
rdb, _ := setupRedis(t)
|
||||
svc := newSvc(t, store, &fakeIncrementer{}, rdb, nil)
|
||||
|
||||
ctx := context.Background()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
svc.RunRetentionLoop(ctx, 0, 100)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(500 * time.Millisecond):
|
||||
t.Fatal(
|
||||
"retention loop should have returned immediately on invalid interval",
|
||||
)
|
||||
}
|
||||
|
||||
require.Equal(t, 0, store.pruneLastLimit)
|
||||
}
|
||||
|
|
@ -0,0 +1,123 @@
|
|||
// ©AngelaMos | 2026
|
||||
// service.go
|
||||
|
||||
package notify
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||
)
|
||||
|
||||
const defaultSendTimeout = 30 * time.Second
|
||||
|
||||
type Service struct {
|
||||
senders map[string]Sender
|
||||
status StatusWriter
|
||||
logger *slog.Logger
|
||||
sendTimeout time.Duration
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
type Option func(*Service)
|
||||
|
||||
func WithLogger(l *slog.Logger) Option {
|
||||
return func(s *Service) { s.logger = l }
|
||||
}
|
||||
|
||||
func WithSendTimeout(d time.Duration) Option {
|
||||
return func(s *Service) { s.sendTimeout = d }
|
||||
}
|
||||
|
||||
func NewService(status StatusWriter, opts ...Option) *Service {
|
||||
s := &Service{
|
||||
senders: make(map[string]Sender),
|
||||
status: status,
|
||||
logger: slog.Default(),
|
||||
sendTimeout: defaultSendTimeout,
|
||||
}
|
||||
for _, o := range opts {
|
||||
o(s)
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *Service) Register(senders ...Sender) {
|
||||
for _, sender := range senders {
|
||||
if sender == nil {
|
||||
continue
|
||||
}
|
||||
s.senders[sender.Channel()] = sender
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) Notify(info event.NotifyInfo, evt *event.Event) {
|
||||
s.wg.Add(1)
|
||||
go func() {
|
||||
defer s.wg.Done()
|
||||
s.dispatch(info, evt)
|
||||
}()
|
||||
}
|
||||
|
||||
func (s *Service) Wait() {
|
||||
s.wg.Wait()
|
||||
}
|
||||
|
||||
func (s *Service) dispatch(info event.NotifyInfo, evt *event.Event) {
|
||||
ctx, cancel := context.WithTimeout(
|
||||
context.Background(),
|
||||
s.sendTimeout,
|
||||
)
|
||||
defer cancel()
|
||||
|
||||
sender, ok := s.senders[info.AlertChannel]
|
||||
if !ok {
|
||||
s.logger.WarnContext(ctx, "notify: no sender registered",
|
||||
"channel", info.AlertChannel,
|
||||
"event_id", evt.ID,
|
||||
"token_id", info.TokenID,
|
||||
)
|
||||
s.markStatus(ctx, evt.ID, event.NotifyFailed, nil)
|
||||
return
|
||||
}
|
||||
|
||||
if err := sender.Send(ctx, info, evt); err != nil {
|
||||
s.logger.WarnContext(ctx, "notify: send failed",
|
||||
"channel", info.AlertChannel,
|
||||
"event_id", evt.ID,
|
||||
"token_id", info.TokenID,
|
||||
"error", err,
|
||||
)
|
||||
s.markStatus(ctx, evt.ID, event.NotifyFailed, nil)
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UTC()
|
||||
s.markStatus(ctx, evt.ID, event.NotifySent, &now)
|
||||
}
|
||||
|
||||
func (s *Service) markStatus(
|
||||
ctx context.Context,
|
||||
eventID int64,
|
||||
status event.NotifyStatus,
|
||||
sentAt *time.Time,
|
||||
) {
|
||||
if s.status == nil {
|
||||
return
|
||||
}
|
||||
if err := s.status.UpdateNotifyStatus(
|
||||
ctx,
|
||||
eventID,
|
||||
status,
|
||||
sentAt,
|
||||
); err != nil {
|
||||
s.logger.WarnContext(ctx, "notify: status writeback failed",
|
||||
"event_id", eventID,
|
||||
"status", status,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,256 @@
|
|||
// ©AngelaMos | 2026
|
||||
// service_test.go
|
||||
|
||||
package notify_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/notify"
|
||||
)
|
||||
|
||||
type fakeSender struct {
|
||||
channel string
|
||||
calls atomic.Int32
|
||||
returnErr error
|
||||
lastInfo atomic.Value
|
||||
delay time.Duration
|
||||
respectCtx bool
|
||||
}
|
||||
|
||||
func (f *fakeSender) Channel() string { return f.channel }
|
||||
|
||||
func (f *fakeSender) Send(
|
||||
ctx context.Context,
|
||||
info event.NotifyInfo,
|
||||
_ *event.Event,
|
||||
) error {
|
||||
f.calls.Add(1)
|
||||
f.lastInfo.Store(info)
|
||||
if f.delay > 0 {
|
||||
if f.respectCtx {
|
||||
select {
|
||||
case <-time.After(f.delay):
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
} else {
|
||||
time.Sleep(f.delay)
|
||||
}
|
||||
}
|
||||
return f.returnErr
|
||||
}
|
||||
|
||||
type fakeStatusWriter struct {
|
||||
mu sync.Mutex
|
||||
updates []statusUpdate
|
||||
err error
|
||||
}
|
||||
|
||||
type statusUpdate struct {
|
||||
eventID int64
|
||||
status event.NotifyStatus
|
||||
sentAt *time.Time
|
||||
}
|
||||
|
||||
func (f *fakeStatusWriter) UpdateNotifyStatus(
|
||||
_ context.Context,
|
||||
id int64,
|
||||
status event.NotifyStatus,
|
||||
sentAt *time.Time,
|
||||
) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.updates = append(
|
||||
f.updates,
|
||||
statusUpdate{eventID: id, status: status, sentAt: sentAt},
|
||||
)
|
||||
return f.err
|
||||
}
|
||||
|
||||
func (f *fakeStatusWriter) snapshot() []statusUpdate {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
out := make([]statusUpdate, len(f.updates))
|
||||
copy(out, f.updates)
|
||||
return out
|
||||
}
|
||||
|
||||
func sampleEvent(id int64) *event.Event {
|
||||
return &event.Event{
|
||||
ID: id,
|
||||
TokenID: "tokfoo000001",
|
||||
TriggeredAt: time.Now().UTC(),
|
||||
SourceIP: "203.0.113.1",
|
||||
}
|
||||
}
|
||||
|
||||
func sampleInfo(channel string) event.NotifyInfo {
|
||||
return event.NotifyInfo{
|
||||
TokenID: "tokfoo000001",
|
||||
ManageID: "abcd",
|
||||
Type: "webbug",
|
||||
Memo: "test",
|
||||
AlertChannel: channel,
|
||||
TelegramBot: "bot",
|
||||
TelegramChat: "chat",
|
||||
WebhookURL: "https://example.com/h",
|
||||
}
|
||||
}
|
||||
|
||||
func newService(
|
||||
t *testing.T,
|
||||
status notify.StatusWriter,
|
||||
senders ...notify.Sender,
|
||||
) *notify.Service {
|
||||
t.Helper()
|
||||
svc := notify.NewService(status,
|
||||
notify.WithLogger(slog.New(slog.NewTextHandler(testWriter{t}, nil))),
|
||||
notify.WithSendTimeout(2*time.Second),
|
||||
)
|
||||
svc.Register(senders...)
|
||||
return svc
|
||||
}
|
||||
|
||||
type testWriter struct{ t *testing.T }
|
||||
|
||||
func (w testWriter) Write(
|
||||
p []byte,
|
||||
) (int, error) {
|
||||
w.t.Log(string(p))
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func TestService_RoutesByChannel(t *testing.T) {
|
||||
t.Parallel()
|
||||
tg := &fakeSender{channel: "telegram"}
|
||||
wh := &fakeSender{channel: "webhook"}
|
||||
status := &fakeStatusWriter{}
|
||||
svc := newService(t, status, tg, wh)
|
||||
|
||||
svc.Notify(sampleInfo("telegram"), sampleEvent(1))
|
||||
svc.Notify(sampleInfo("webhook"), sampleEvent(2))
|
||||
svc.Wait()
|
||||
|
||||
require.Equal(t, int32(1), tg.calls.Load())
|
||||
require.Equal(t, int32(1), wh.calls.Load())
|
||||
}
|
||||
|
||||
func TestService_MarksSentOnSuccess(t *testing.T) {
|
||||
t.Parallel()
|
||||
tg := &fakeSender{channel: "telegram"}
|
||||
status := &fakeStatusWriter{}
|
||||
svc := newService(t, status, tg)
|
||||
|
||||
svc.Notify(sampleInfo("telegram"), sampleEvent(1))
|
||||
svc.Wait()
|
||||
|
||||
updates := status.snapshot()
|
||||
require.Len(t, updates, 1)
|
||||
require.Equal(t, int64(1), updates[0].eventID)
|
||||
require.Equal(t, event.NotifySent, updates[0].status)
|
||||
require.NotNil(t, updates[0].sentAt)
|
||||
}
|
||||
|
||||
func TestService_MarksFailedOnSenderError(t *testing.T) {
|
||||
t.Parallel()
|
||||
tg := &fakeSender{channel: "telegram", returnErr: errors.New("api blew up")}
|
||||
status := &fakeStatusWriter{}
|
||||
svc := newService(t, status, tg)
|
||||
|
||||
svc.Notify(sampleInfo("telegram"), sampleEvent(7))
|
||||
svc.Wait()
|
||||
|
||||
updates := status.snapshot()
|
||||
require.Len(t, updates, 1)
|
||||
require.Equal(t, int64(7), updates[0].eventID)
|
||||
require.Equal(t, event.NotifyFailed, updates[0].status)
|
||||
require.Nil(t, updates[0].sentAt)
|
||||
}
|
||||
|
||||
func TestService_MarksFailedWhenChannelUnknown(t *testing.T) {
|
||||
t.Parallel()
|
||||
tg := &fakeSender{channel: "telegram"}
|
||||
status := &fakeStatusWriter{}
|
||||
svc := newService(t, status, tg)
|
||||
|
||||
svc.Notify(sampleInfo("smoke-signal"), sampleEvent(99))
|
||||
svc.Wait()
|
||||
|
||||
require.Equal(t, int32(0), tg.calls.Load())
|
||||
updates := status.snapshot()
|
||||
require.Len(t, updates, 1)
|
||||
require.Equal(t, event.NotifyFailed, updates[0].status)
|
||||
}
|
||||
|
||||
func TestService_NotifyIsAsync(t *testing.T) {
|
||||
t.Parallel()
|
||||
tg := &fakeSender{channel: "telegram", delay: 100 * time.Millisecond}
|
||||
status := &fakeStatusWriter{}
|
||||
svc := newService(t, status, tg)
|
||||
|
||||
start := time.Now()
|
||||
svc.Notify(sampleInfo("telegram"), sampleEvent(1))
|
||||
require.Less(t, time.Since(start), 50*time.Millisecond,
|
||||
"Notify should return immediately, not wait for the send")
|
||||
svc.Wait()
|
||||
require.Equal(t, int32(1), tg.calls.Load())
|
||||
}
|
||||
|
||||
func TestService_DispatchTimeoutBoundsSender(t *testing.T) {
|
||||
t.Parallel()
|
||||
tg := &fakeSender{
|
||||
channel: "telegram",
|
||||
delay: 2 * time.Second,
|
||||
respectCtx: true,
|
||||
}
|
||||
status := &fakeStatusWriter{}
|
||||
svc := notify.NewService(status,
|
||||
notify.WithSendTimeout(50*time.Millisecond),
|
||||
)
|
||||
svc.Register(tg)
|
||||
|
||||
svc.Notify(sampleInfo("telegram"), sampleEvent(1))
|
||||
svc.Wait()
|
||||
|
||||
updates := status.snapshot()
|
||||
require.Len(t, updates, 1)
|
||||
require.Equal(t, event.NotifyFailed, updates[0].status,
|
||||
"timeout should mark event failed")
|
||||
}
|
||||
|
||||
func TestService_StatusWriterErrorDoesNotPanic(t *testing.T) {
|
||||
t.Parallel()
|
||||
tg := &fakeSender{channel: "telegram"}
|
||||
status := &fakeStatusWriter{err: errors.New("db down")}
|
||||
svc := newService(t, status, tg)
|
||||
|
||||
require.NotPanics(t, func() {
|
||||
svc.Notify(sampleInfo("telegram"), sampleEvent(1))
|
||||
svc.Wait()
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_ConcurrentNotifyAllComplete(t *testing.T) {
|
||||
t.Parallel()
|
||||
tg := &fakeSender{channel: "telegram", delay: 10 * time.Millisecond}
|
||||
status := &fakeStatusWriter{}
|
||||
svc := newService(t, status, tg)
|
||||
|
||||
const n = 50
|
||||
for i := range n {
|
||||
svc.Notify(sampleInfo("telegram"), sampleEvent(int64(i+1)))
|
||||
}
|
||||
svc.Wait()
|
||||
require.Equal(t, int32(n), tg.calls.Load())
|
||||
require.Len(t, status.snapshot(), n)
|
||||
}
|
||||
|
|
@ -0,0 +1,298 @@
|
|||
// ©AngelaMos | 2026
|
||||
// sender.go
|
||||
|
||||
package telegram
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/cenkalti/backoff/v5"
|
||||
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||
)
|
||||
|
||||
const (
|
||||
Channel = "telegram"
|
||||
|
||||
defaultAPIBase = "https://api.telegram.org"
|
||||
defaultMaxTries = 3
|
||||
defaultMaxElapsed = 30 * time.Second
|
||||
defaultInitialInterval = 500 * time.Millisecond
|
||||
defaultOverallTimeout = 10 * time.Second
|
||||
defaultDialTimeout = 5 * time.Second
|
||||
|
||||
uaTruncateRunes = 80
|
||||
|
||||
parseModeMarkdownV2 = "MarkdownV2"
|
||||
contentTypeJSON = "application/json"
|
||||
|
||||
v2SpecialChars = "_*[]()~`>#+-=|{}.!"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrChannelNotConfigured = errors.New(
|
||||
"telegram: bot token or chat id not configured",
|
||||
)
|
||||
ErrTelegramAPI = errors.New("telegram: api error")
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
APIBase string
|
||||
ManageURL string
|
||||
HTTPClient *http.Client
|
||||
MaxTries uint
|
||||
MaxElapsed time.Duration
|
||||
InitialInterval time.Duration
|
||||
}
|
||||
|
||||
type Option func(*Config)
|
||||
|
||||
func WithMaxTries(n uint) Option {
|
||||
return func(c *Config) { c.MaxTries = n }
|
||||
}
|
||||
|
||||
func WithMaxElapsed(d time.Duration) Option {
|
||||
return func(c *Config) { c.MaxElapsed = d }
|
||||
}
|
||||
|
||||
func WithInitialInterval(d time.Duration) Option {
|
||||
return func(c *Config) { c.InitialInterval = d }
|
||||
}
|
||||
|
||||
func WithHTTPClient(client *http.Client) Option {
|
||||
return func(c *Config) { c.HTTPClient = client }
|
||||
}
|
||||
|
||||
type Sender struct {
|
||||
apiBase string
|
||||
manageURL string
|
||||
httpClient *http.Client
|
||||
maxTries uint
|
||||
maxElapsed time.Duration
|
||||
initialInterval time.Duration
|
||||
}
|
||||
|
||||
func NewSender(cfg Config, opts ...Option) *Sender {
|
||||
for _, o := range opts {
|
||||
o(&cfg)
|
||||
}
|
||||
if cfg.APIBase == "" {
|
||||
cfg.APIBase = defaultAPIBase
|
||||
}
|
||||
if cfg.HTTPClient == nil {
|
||||
cfg.HTTPClient = defaultHTTPClient()
|
||||
}
|
||||
if cfg.MaxTries == 0 {
|
||||
cfg.MaxTries = defaultMaxTries
|
||||
}
|
||||
if cfg.MaxElapsed == 0 {
|
||||
cfg.MaxElapsed = defaultMaxElapsed
|
||||
}
|
||||
if cfg.InitialInterval == 0 {
|
||||
cfg.InitialInterval = defaultInitialInterval
|
||||
}
|
||||
return &Sender{
|
||||
apiBase: strings.TrimRight(cfg.APIBase, "/"),
|
||||
manageURL: strings.TrimRight(cfg.ManageURL, "/"),
|
||||
httpClient: cfg.HTTPClient,
|
||||
maxTries: cfg.MaxTries,
|
||||
maxElapsed: cfg.MaxElapsed,
|
||||
initialInterval: cfg.InitialInterval,
|
||||
}
|
||||
}
|
||||
|
||||
func defaultHTTPClient() *http.Client {
|
||||
dialer := &net.Dialer{Timeout: defaultDialTimeout}
|
||||
return &http.Client{
|
||||
Timeout: defaultOverallTimeout,
|
||||
Transport: &http.Transport{
|
||||
DialContext: dialer.DialContext,
|
||||
TLSHandshakeTimeout: defaultDialTimeout,
|
||||
ResponseHeaderTimeout: defaultOverallTimeout,
|
||||
ExpectContinueTimeout: time.Second,
|
||||
IdleConnTimeout: 30 * time.Second,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Sender) Channel() string { return Channel }
|
||||
|
||||
func (s *Sender) Send(
|
||||
ctx context.Context,
|
||||
info event.NotifyInfo,
|
||||
evt *event.Event,
|
||||
) error {
|
||||
if info.TelegramBot == "" || info.TelegramChat == "" {
|
||||
return ErrChannelNotConfigured
|
||||
}
|
||||
endpoint := s.apiBase + "/bot" + info.TelegramBot + "/sendMessage"
|
||||
body, err := json.Marshal(map[string]string{
|
||||
"chat_id": info.TelegramChat,
|
||||
"text": buildMessage(info, evt, s.manageURL),
|
||||
"parse_mode": parseModeMarkdownV2,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("telegram: marshal body: %w", err)
|
||||
}
|
||||
|
||||
expBackoff := backoff.NewExponentialBackOff()
|
||||
expBackoff.InitialInterval = s.initialInterval
|
||||
expBackoff.MaxInterval = 5 * time.Second
|
||||
|
||||
_, err = backoff.Retry(
|
||||
ctx,
|
||||
func() (struct{}, error) {
|
||||
return struct{}{}, s.doRequest(ctx, endpoint, body)
|
||||
},
|
||||
backoff.WithBackOff(expBackoff),
|
||||
backoff.WithMaxTries(s.maxTries),
|
||||
backoff.WithMaxElapsedTime(s.maxElapsed),
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Sender) doRequest(
|
||||
ctx context.Context,
|
||||
endpoint string,
|
||||
body []byte,
|
||||
) error {
|
||||
req, err := http.NewRequestWithContext(
|
||||
ctx,
|
||||
http.MethodPost,
|
||||
endpoint,
|
||||
bytes.NewReader(body),
|
||||
)
|
||||
if err != nil {
|
||||
return backoff.Permanent(
|
||||
fmt.Errorf("telegram: build request: %w", err),
|
||||
)
|
||||
}
|
||||
req.Header.Set("Content-Type", contentTypeJSON)
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("telegram: do request: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if cErr := resp.Body.Close(); cErr != nil {
|
||||
slog.WarnContext(ctx, "telegram: close body",
|
||||
"error", cErr)
|
||||
}
|
||||
}()
|
||||
|
||||
respBody, rErr := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
if rErr != nil {
|
||||
slog.WarnContext(ctx, "telegram: read body", "error", rErr)
|
||||
}
|
||||
|
||||
switch {
|
||||
case resp.StatusCode >= 200 && resp.StatusCode < 300:
|
||||
return nil
|
||||
case resp.StatusCode >= 400 && resp.StatusCode < 500:
|
||||
return backoff.Permanent(fmt.Errorf(
|
||||
"%w: status=%d body=%s",
|
||||
ErrTelegramAPI, resp.StatusCode, string(respBody),
|
||||
))
|
||||
default:
|
||||
return fmt.Errorf(
|
||||
"%w: status=%d body=%s",
|
||||
ErrTelegramAPI, resp.StatusCode, string(respBody),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func buildMessage(
|
||||
info event.NotifyInfo,
|
||||
evt *event.Event,
|
||||
manageURL string,
|
||||
) string {
|
||||
var b strings.Builder
|
||||
b.WriteString("🚨 *Canary triggered:* ")
|
||||
b.WriteString(EscapeMD(info.Memo))
|
||||
b.WriteString("\n\n*Type:* ")
|
||||
b.WriteString(EscapeMD(info.Type))
|
||||
b.WriteString("\n*From:* ")
|
||||
b.WriteString(EscapeMD(evt.SourceIP))
|
||||
if loc := formatGeo(evt); loc != "" {
|
||||
b.WriteString(" ")
|
||||
b.WriteString(loc)
|
||||
}
|
||||
b.WriteString("\n*Time:* ")
|
||||
b.WriteString(EscapeMD(evt.TriggeredAt.UTC().Format(time.RFC3339)))
|
||||
if evt.UserAgent != nil && *evt.UserAgent != "" {
|
||||
b.WriteString("\n*UA:* ")
|
||||
b.WriteString(EscapeMD(truncateRunes(*evt.UserAgent, uaTruncateRunes)))
|
||||
}
|
||||
if manageURL != "" && info.ManageID != "" {
|
||||
b.WriteString("\n\n[View full event timeline](")
|
||||
b.WriteString(manageURL + "/m/" + info.ManageID)
|
||||
b.WriteString(")")
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func formatGeo(evt *event.Event) string {
|
||||
city := derefStr(evt.GeoCity)
|
||||
country := derefStr(evt.GeoCountry)
|
||||
asnOrg := derefStr(evt.GeoASNOrg)
|
||||
|
||||
var parens string
|
||||
switch {
|
||||
case city != "" && country != "":
|
||||
parens = "(" + EscapeMD(city) + ", " + EscapeMD(country) + ")"
|
||||
case country != "":
|
||||
parens = "(" + EscapeMD(country) + ")"
|
||||
case city != "":
|
||||
parens = "(" + EscapeMD(city) + ")"
|
||||
}
|
||||
if parens == "" && asnOrg == "" {
|
||||
return ""
|
||||
}
|
||||
if asnOrg == "" {
|
||||
return parens
|
||||
}
|
||||
if parens == "" {
|
||||
return "— " + EscapeMD(asnOrg)
|
||||
}
|
||||
return parens + " — " + EscapeMD(asnOrg)
|
||||
}
|
||||
|
||||
func derefStr(s *string) string {
|
||||
if s == nil {
|
||||
return ""
|
||||
}
|
||||
return *s
|
||||
}
|
||||
|
||||
func truncateRunes(s string, max int) string {
|
||||
if max <= 0 {
|
||||
return ""
|
||||
}
|
||||
r := []rune(s)
|
||||
if len(r) <= max {
|
||||
return s
|
||||
}
|
||||
return string(r[:max])
|
||||
}
|
||||
|
||||
func EscapeMD(s string) string {
|
||||
var b strings.Builder
|
||||
b.Grow(len(s) + 8)
|
||||
for _, r := range s {
|
||||
if strings.ContainsRune(v2SpecialChars, r) {
|
||||
b.WriteRune('\\')
|
||||
}
|
||||
b.WriteRune(r)
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
|
@ -0,0 +1,388 @@
|
|||
// ©AngelaMos | 2026
|
||||
// sender_test.go
|
||||
|
||||
package telegram_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/notify/telegram"
|
||||
)
|
||||
|
||||
const (
|
||||
testBotToken = "111222:ABCDEFG"
|
||||
testChatID = "98765"
|
||||
testMemo = "prod-db-creds"
|
||||
testTokenID = "tokabc123def"
|
||||
testManageID = "11111111-2222-3333-4444-555555555555"
|
||||
)
|
||||
|
||||
func sampleInfo() event.NotifyInfo {
|
||||
return event.NotifyInfo{
|
||||
TokenID: testTokenID,
|
||||
ManageID: testManageID,
|
||||
Type: "envfile",
|
||||
Memo: testMemo,
|
||||
AlertChannel: "telegram",
|
||||
TelegramBot: testBotToken,
|
||||
TelegramChat: testChatID,
|
||||
}
|
||||
}
|
||||
|
||||
func sampleEvent() *event.Event {
|
||||
ua := "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||||
city := "Toronto"
|
||||
country := "CA"
|
||||
asnOrg := "Cloudflare, Inc."
|
||||
return &event.Event{
|
||||
ID: 42,
|
||||
TokenID: testTokenID,
|
||||
TriggeredAt: time.Date(2026, 5, 14, 12, 30, 0, 0, time.UTC),
|
||||
SourceIP: "203.0.113.45",
|
||||
UserAgent: &ua,
|
||||
GeoCity: &city,
|
||||
GeoCountry: &country,
|
||||
GeoASNOrg: &asnOrg,
|
||||
}
|
||||
}
|
||||
|
||||
type capture struct {
|
||||
calls atomic.Int32
|
||||
lastURL atomic.Value
|
||||
lastBody atomic.Value
|
||||
}
|
||||
|
||||
func writeOK(t *testing.T, w http.ResponseWriter) {
|
||||
t.Helper()
|
||||
w.WriteHeader(http.StatusOK)
|
||||
if _, err := w.Write([]byte(`{"ok":true}`)); err != nil {
|
||||
t.Logf("write: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func newCaptureServer(
|
||||
t *testing.T,
|
||||
handler http.HandlerFunc,
|
||||
) (*httptest.Server, *capture) {
|
||||
t.Helper()
|
||||
c := &capture{}
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
c.calls.Add(1)
|
||||
c.lastURL.Store(r.URL.String())
|
||||
body, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
t.Logf("read body: %v", err)
|
||||
return
|
||||
}
|
||||
c.lastBody.Store(body)
|
||||
handler(w, r)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
return srv, c
|
||||
}
|
||||
|
||||
func loadBody(t *testing.T, c *capture) []byte {
|
||||
t.Helper()
|
||||
raw, ok := c.lastBody.Load().([]byte)
|
||||
require.True(t, ok, "no body captured")
|
||||
return raw
|
||||
}
|
||||
|
||||
func newSender(
|
||||
t *testing.T,
|
||||
apiBase string,
|
||||
opts ...telegram.Option,
|
||||
) *telegram.Sender {
|
||||
t.Helper()
|
||||
cfg := telegram.Config{
|
||||
APIBase: apiBase,
|
||||
ManageURL: "https://canary.example.com",
|
||||
}
|
||||
for _, o := range opts {
|
||||
o(&cfg)
|
||||
}
|
||||
return telegram.NewSender(cfg)
|
||||
}
|
||||
|
||||
func TestSender_Channel(t *testing.T) {
|
||||
t.Parallel()
|
||||
s := newSender(t, "https://api.telegram.org")
|
||||
require.Equal(t, "telegram", s.Channel())
|
||||
}
|
||||
|
||||
func TestSender_Send_PostsToCorrectURL(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) { writeOK(t, w) },
|
||||
)
|
||||
s := newSender(t, srv.URL)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(), sampleEvent()),
|
||||
)
|
||||
require.Equal(t, int32(1), cap.calls.Load())
|
||||
require.Equal(t, "/bot"+testBotToken+"/sendMessage", cap.lastURL.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_BodyShape(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) { writeOK(t, w) },
|
||||
)
|
||||
s := newSender(t, srv.URL)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(), sampleEvent()),
|
||||
)
|
||||
|
||||
var body map[string]string
|
||||
require.NoError(t, json.Unmarshal(loadBody(t, cap), &body))
|
||||
require.Equal(t, testChatID, body["chat_id"])
|
||||
require.Equal(t, "MarkdownV2", body["parse_mode"])
|
||||
require.NotEmpty(t, body["text"])
|
||||
}
|
||||
|
||||
func TestSender_Send_MessageContainsKeyFields(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) { writeOK(t, w) },
|
||||
)
|
||||
s := newSender(t, srv.URL)
|
||||
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(), sampleEvent()),
|
||||
)
|
||||
|
||||
var body map[string]string
|
||||
require.NoError(t, json.Unmarshal(loadBody(t, cap), &body))
|
||||
text := body["text"]
|
||||
|
||||
require.Contains(t, text, "Canary triggered")
|
||||
require.Contains(
|
||||
t,
|
||||
text,
|
||||
`prod\-db\-creds`,
|
||||
"memo escaped (- is V2 special)",
|
||||
)
|
||||
require.Contains(t, text, "envfile")
|
||||
require.Contains(t, text, `203\.0\.113\.45`, "IP dots escaped")
|
||||
require.Contains(t, text, "Toronto")
|
||||
require.Contains(t, text, "CA")
|
||||
require.Contains(t, text, `Cloudflare, Inc\.`, "asn_org . escaped")
|
||||
require.Contains(t, text, "View full event timeline", "manage link present")
|
||||
require.Contains(t,
|
||||
text,
|
||||
"https://canary.example.com/m/"+testManageID,
|
||||
"manage URL present in link",
|
||||
)
|
||||
}
|
||||
|
||||
func TestSender_Send_TruncatesUserAgentTo80Chars(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) { writeOK(t, w) },
|
||||
)
|
||||
s := newSender(t, srv.URL)
|
||||
|
||||
longUA := strings.Repeat("X", 200)
|
||||
evt := sampleEvent()
|
||||
evt.UserAgent = &longUA
|
||||
|
||||
require.NoError(t, s.Send(context.Background(), sampleInfo(), evt))
|
||||
|
||||
var body map[string]string
|
||||
require.NoError(t, json.Unmarshal(loadBody(t, cap), &body))
|
||||
require.Contains(t, body["text"], strings.Repeat("X", 80))
|
||||
require.NotContains(t, body["text"], strings.Repeat("X", 81),
|
||||
"user agent should be truncated to 80 chars")
|
||||
}
|
||||
|
||||
func TestSender_Send_HandlesNilGeoAndUA(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) { writeOK(t, w) },
|
||||
)
|
||||
s := newSender(t, srv.URL)
|
||||
|
||||
evt := &event.Event{
|
||||
ID: 1,
|
||||
TokenID: testTokenID,
|
||||
TriggeredAt: time.Date(2026, 5, 14, 12, 30, 0, 0, time.UTC),
|
||||
SourceIP: "203.0.113.45",
|
||||
}
|
||||
require.NoError(t, s.Send(context.Background(), sampleInfo(), evt))
|
||||
|
||||
var body map[string]string
|
||||
require.NoError(t, json.Unmarshal(loadBody(t, cap), &body))
|
||||
require.Contains(t, body["text"], `203\.0\.113\.45`)
|
||||
}
|
||||
|
||||
func TestSender_Send_EmptyBotReturnsConfigErr(t *testing.T) {
|
||||
t.Parallel()
|
||||
s := newSender(t, "http://unused.test")
|
||||
info := sampleInfo()
|
||||
info.TelegramBot = ""
|
||||
err := s.Send(context.Background(), info, sampleEvent())
|
||||
require.ErrorIs(t, err, telegram.ErrChannelNotConfigured)
|
||||
}
|
||||
|
||||
func TestSender_Send_EmptyChatReturnsConfigErr(t *testing.T) {
|
||||
t.Parallel()
|
||||
s := newSender(t, "http://unused.test")
|
||||
info := sampleInfo()
|
||||
info.TelegramChat = ""
|
||||
err := s.Send(context.Background(), info, sampleEvent())
|
||||
require.ErrorIs(t, err, telegram.ErrChannelNotConfigured)
|
||||
}
|
||||
|
||||
func TestSender_Send_RetriesOn5xxThenSucceeds(t *testing.T) {
|
||||
t.Parallel()
|
||||
var attempts atomic.Int32
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
n := attempts.Add(1)
|
||||
if n == 1 {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
writeOK(t, w)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
s := newSender(t, srv.URL,
|
||||
telegram.WithMaxTries(3),
|
||||
telegram.WithMaxElapsed(2*time.Second),
|
||||
telegram.WithInitialInterval(5*time.Millisecond),
|
||||
)
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(), sampleEvent()),
|
||||
)
|
||||
require.Equal(t, int32(2), attempts.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_PermanentOn4xx(t *testing.T) {
|
||||
t.Parallel()
|
||||
var attempts atomic.Int32
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
attempts.Add(1)
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
if _, err := w.Write(
|
||||
[]byte(`{"ok":false,"description":"chat not found"}`),
|
||||
); err != nil {
|
||||
t.Logf("write: %v", err)
|
||||
}
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
s := newSender(t, srv.URL,
|
||||
telegram.WithMaxTries(5),
|
||||
telegram.WithInitialInterval(5*time.Millisecond),
|
||||
)
|
||||
err := s.Send(context.Background(), sampleInfo(), sampleEvent())
|
||||
require.Error(t, err)
|
||||
require.Equal(t, int32(1), attempts.Load(), "no retry on 4xx")
|
||||
}
|
||||
|
||||
func TestSender_Send_AbortsAfterMaxTries(t *testing.T) {
|
||||
t.Parallel()
|
||||
var attempts atomic.Int32
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
attempts.Add(1)
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
s := newSender(t, srv.URL,
|
||||
telegram.WithMaxTries(3),
|
||||
telegram.WithMaxElapsed(5*time.Second),
|
||||
telegram.WithInitialInterval(2*time.Millisecond),
|
||||
)
|
||||
err := s.Send(context.Background(), sampleInfo(), sampleEvent())
|
||||
require.Error(t, err)
|
||||
require.Equal(t, int32(3), attempts.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_RespectsContextCancel(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {
|
||||
time.Sleep(2 * time.Second)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
s := newSender(t, srv.URL,
|
||||
telegram.WithMaxTries(3),
|
||||
telegram.WithInitialInterval(5*time.Millisecond),
|
||||
)
|
||||
ctx, cancel := context.WithTimeout(
|
||||
context.Background(),
|
||||
100*time.Millisecond,
|
||||
)
|
||||
defer cancel()
|
||||
|
||||
err := s.Send(ctx, sampleInfo(), sampleEvent())
|
||||
require.Error(t, err)
|
||||
require.True(t,
|
||||
errors.Is(err, context.DeadlineExceeded) ||
|
||||
errors.Is(err, context.Canceled),
|
||||
"expected context error, got %v", err,
|
||||
)
|
||||
}
|
||||
|
||||
func TestEscapeMD(t *testing.T) {
|
||||
t.Parallel()
|
||||
cases := []struct {
|
||||
in, want string
|
||||
}{
|
||||
{"plain", "plain"},
|
||||
{"a.b", `a\.b`},
|
||||
{"a-b", `a\-b`},
|
||||
{"file.txt", `file\.txt`},
|
||||
{"203.0.113.45", `203\.0\.113\.45`},
|
||||
{"hello!", `hello\!`},
|
||||
{"a_b*c", `a\_b\*c`},
|
||||
{"(parens)", `\(parens\)`},
|
||||
{"[bracket]", `\[bracket\]`},
|
||||
{"~tilde~", `\~tilde\~`},
|
||||
{"`code`", "\\`code\\`"},
|
||||
{"a>b<c", `a\>b<c`},
|
||||
{"a#b", `a\#b`},
|
||||
{"a+b=c", `a\+b\=c`},
|
||||
{"a|b{c}d", `a\|b\{c\}d`},
|
||||
{"unicode—em", "unicode—em"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.in, func(t *testing.T) {
|
||||
require.Equal(t, tc.want, telegram.EscapeMD(tc.in))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,29 @@
|
|||
// ©AngelaMos | 2026
|
||||
// types.go
|
||||
|
||||
package notify
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||
)
|
||||
|
||||
type Sender interface {
|
||||
Channel() string
|
||||
Send(
|
||||
ctx context.Context,
|
||||
info event.NotifyInfo,
|
||||
evt *event.Event,
|
||||
) error
|
||||
}
|
||||
|
||||
type StatusWriter interface {
|
||||
UpdateNotifyStatus(
|
||||
ctx context.Context,
|
||||
eventID int64,
|
||||
status event.NotifyStatus,
|
||||
sentAt *time.Time,
|
||||
) error
|
||||
}
|
||||
|
|
@ -0,0 +1,317 @@
|
|||
// ©AngelaMos | 2026
|
||||
// sender.go
|
||||
|
||||
package webhook
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/cenkalti/backoff/v5"
|
||||
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||
)
|
||||
|
||||
const (
|
||||
Channel = "webhook"
|
||||
|
||||
envelopeVersion = "1"
|
||||
envelopeEvent = "canary.triggered"
|
||||
|
||||
defaultMaxTries = 3
|
||||
defaultMaxElapsed = 30 * time.Second
|
||||
defaultInitialInterval = 500 * time.Millisecond
|
||||
defaultOverallTimeout = 10 * time.Second
|
||||
defaultDialTimeout = 5 * time.Second
|
||||
|
||||
contentTypeJSON = "application/json"
|
||||
signatureHeaderName = "X-Canary-Signature"
|
||||
signaturePrefix = "sha256="
|
||||
)
|
||||
|
||||
var (
|
||||
ErrChannelNotConfigured = errors.New(
|
||||
"webhook: webhook URL not configured",
|
||||
)
|
||||
ErrInvalidWebhookURL = errors.New("webhook: invalid url")
|
||||
ErrWebhookAPI = errors.New("webhook: api error")
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
ManageURL string
|
||||
HMACSecret string
|
||||
HTTPClient *http.Client
|
||||
MaxTries uint
|
||||
MaxElapsed time.Duration
|
||||
InitialInterval time.Duration
|
||||
}
|
||||
|
||||
type Option func(*Config)
|
||||
|
||||
func WithMaxTries(n uint) Option { return func(c *Config) { c.MaxTries = n } }
|
||||
|
||||
func WithMaxElapsed(d time.Duration) Option {
|
||||
return func(c *Config) { c.MaxElapsed = d }
|
||||
}
|
||||
|
||||
func WithInitialInterval(d time.Duration) Option {
|
||||
return func(c *Config) { c.InitialInterval = d }
|
||||
}
|
||||
|
||||
func WithHTTPClient(client *http.Client) Option {
|
||||
return func(c *Config) { c.HTTPClient = client }
|
||||
}
|
||||
|
||||
type Sender struct {
|
||||
manageURL string
|
||||
hmacSecret string
|
||||
httpClient *http.Client
|
||||
maxTries uint
|
||||
maxElapsed time.Duration
|
||||
initialInterval time.Duration
|
||||
}
|
||||
|
||||
func NewSender(cfg Config, opts ...Option) *Sender {
|
||||
for _, o := range opts {
|
||||
o(&cfg)
|
||||
}
|
||||
if cfg.HTTPClient == nil {
|
||||
cfg.HTTPClient = defaultHTTPClient()
|
||||
}
|
||||
if cfg.MaxTries == 0 {
|
||||
cfg.MaxTries = defaultMaxTries
|
||||
}
|
||||
if cfg.MaxElapsed == 0 {
|
||||
cfg.MaxElapsed = defaultMaxElapsed
|
||||
}
|
||||
if cfg.InitialInterval == 0 {
|
||||
cfg.InitialInterval = defaultInitialInterval
|
||||
}
|
||||
return &Sender{
|
||||
manageURL: strings.TrimRight(cfg.ManageURL, "/"),
|
||||
hmacSecret: cfg.HMACSecret,
|
||||
httpClient: cfg.HTTPClient,
|
||||
maxTries: cfg.MaxTries,
|
||||
maxElapsed: cfg.MaxElapsed,
|
||||
initialInterval: cfg.InitialInterval,
|
||||
}
|
||||
}
|
||||
|
||||
func defaultHTTPClient() *http.Client {
|
||||
dialer := &net.Dialer{Timeout: defaultDialTimeout}
|
||||
return &http.Client{
|
||||
Timeout: defaultOverallTimeout,
|
||||
Transport: &http.Transport{
|
||||
DialContext: dialer.DialContext,
|
||||
TLSHandshakeTimeout: defaultDialTimeout,
|
||||
ResponseHeaderTimeout: defaultOverallTimeout,
|
||||
ExpectContinueTimeout: time.Second,
|
||||
IdleConnTimeout: 30 * time.Second,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Sender) Channel() string { return Channel }
|
||||
|
||||
func (s *Sender) Send(
|
||||
ctx context.Context,
|
||||
info event.NotifyInfo,
|
||||
evt *event.Event,
|
||||
) error {
|
||||
if strings.TrimSpace(info.WebhookURL) == "" {
|
||||
return ErrChannelNotConfigured
|
||||
}
|
||||
if err := validateURL(info.WebhookURL); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
body, err := json.Marshal(buildEnvelope(info, evt, s.manageURL))
|
||||
if err != nil {
|
||||
return fmt.Errorf("webhook: marshal envelope: %w", err)
|
||||
}
|
||||
|
||||
expBackoff := backoff.NewExponentialBackOff()
|
||||
expBackoff.InitialInterval = s.initialInterval
|
||||
expBackoff.MaxInterval = 5 * time.Second
|
||||
|
||||
_, err = backoff.Retry(
|
||||
ctx,
|
||||
func() (struct{}, error) {
|
||||
return struct{}{}, s.doRequest(ctx, info.WebhookURL, body)
|
||||
},
|
||||
backoff.WithBackOff(expBackoff),
|
||||
backoff.WithMaxTries(s.maxTries),
|
||||
backoff.WithMaxElapsedTime(s.maxElapsed),
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Sender) doRequest(
|
||||
ctx context.Context,
|
||||
endpoint string,
|
||||
body []byte,
|
||||
) error {
|
||||
req, err := http.NewRequestWithContext(
|
||||
ctx,
|
||||
http.MethodPost,
|
||||
endpoint,
|
||||
bytes.NewReader(body),
|
||||
)
|
||||
if err != nil {
|
||||
return backoff.Permanent(
|
||||
fmt.Errorf("webhook: build request: %w", err),
|
||||
)
|
||||
}
|
||||
req.Header.Set("Content-Type", contentTypeJSON)
|
||||
if s.hmacSecret != "" {
|
||||
req.Header.Set(
|
||||
signatureHeaderName,
|
||||
computeSignature(s.hmacSecret, body),
|
||||
)
|
||||
}
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("webhook: do request: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if cErr := resp.Body.Close(); cErr != nil {
|
||||
slog.WarnContext(ctx, "webhook: close body",
|
||||
"error", cErr)
|
||||
}
|
||||
}()
|
||||
|
||||
respBody, rErr := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
if rErr != nil {
|
||||
slog.WarnContext(ctx, "webhook: read body", "error", rErr)
|
||||
}
|
||||
|
||||
switch {
|
||||
case resp.StatusCode >= 200 && resp.StatusCode < 300:
|
||||
return nil
|
||||
case resp.StatusCode >= 400 && resp.StatusCode < 500:
|
||||
return backoff.Permanent(fmt.Errorf(
|
||||
"%w: status=%d body=%s",
|
||||
ErrWebhookAPI, resp.StatusCode, string(respBody),
|
||||
))
|
||||
default:
|
||||
return fmt.Errorf(
|
||||
"%w: status=%d body=%s",
|
||||
ErrWebhookAPI, resp.StatusCode, string(respBody),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func validateURL(raw string) error {
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: parse: %w", ErrInvalidWebhookURL, err)
|
||||
}
|
||||
scheme := strings.ToLower(u.Scheme)
|
||||
if scheme != "http" && scheme != "https" {
|
||||
return fmt.Errorf(
|
||||
"%w: scheme must be http or https, got %q",
|
||||
ErrInvalidWebhookURL, u.Scheme,
|
||||
)
|
||||
}
|
||||
if u.Host == "" {
|
||||
return fmt.Errorf("%w: missing host", ErrInvalidWebhookURL)
|
||||
}
|
||||
if u.User != nil {
|
||||
return fmt.Errorf("%w: userinfo not allowed", ErrInvalidWebhookURL)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type envelope struct {
|
||||
Version string `json:"version"`
|
||||
Event string `json:"event"`
|
||||
Token tokenSection `json:"token"`
|
||||
Trigger triggerSection `json:"trigger"`
|
||||
}
|
||||
|
||||
type tokenSection struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
Memo string `json:"memo"`
|
||||
ManageURL string `json:"manage_url"`
|
||||
}
|
||||
|
||||
type triggerSection struct {
|
||||
TriggeredAt time.Time `json:"triggered_at"`
|
||||
SourceIP string `json:"source_ip"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
Geo geoSection `json:"geo"`
|
||||
Extra json.RawMessage `json:"extra"`
|
||||
}
|
||||
|
||||
type geoSection struct {
|
||||
Country string `json:"country"`
|
||||
City string `json:"city"`
|
||||
ASNOrg string `json:"asn_org"`
|
||||
}
|
||||
|
||||
func buildEnvelope(
|
||||
info event.NotifyInfo,
|
||||
evt *event.Event,
|
||||
manageURL string,
|
||||
) envelope {
|
||||
extra := evt.Extra
|
||||
if len(extra) == 0 {
|
||||
extra = json.RawMessage(`{}`)
|
||||
}
|
||||
return envelope{
|
||||
Version: envelopeVersion,
|
||||
Event: envelopeEvent,
|
||||
Token: tokenSection{
|
||||
ID: info.TokenID,
|
||||
Type: info.Type,
|
||||
Memo: info.Memo,
|
||||
ManageURL: buildManageURL(manageURL, info.ManageID),
|
||||
},
|
||||
Trigger: triggerSection{
|
||||
TriggeredAt: evt.TriggeredAt.UTC(),
|
||||
SourceIP: evt.SourceIP,
|
||||
UserAgent: derefStr(evt.UserAgent),
|
||||
Geo: geoSection{
|
||||
Country: derefStr(evt.GeoCountry),
|
||||
City: derefStr(evt.GeoCity),
|
||||
ASNOrg: derefStr(evt.GeoASNOrg),
|
||||
},
|
||||
Extra: extra,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildManageURL(base, id string) string {
|
||||
if base == "" || id == "" {
|
||||
return ""
|
||||
}
|
||||
return base + "/m/" + id
|
||||
}
|
||||
|
||||
func derefStr(s *string) string {
|
||||
if s == nil {
|
||||
return ""
|
||||
}
|
||||
return *s
|
||||
}
|
||||
|
||||
func computeSignature(secret string, body []byte) string {
|
||||
mac := hmac.New(sha256.New, []byte(secret))
|
||||
mac.Write(body)
|
||||
return signaturePrefix + hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
|
@ -0,0 +1,355 @@
|
|||
// ©AngelaMos | 2026
|
||||
// sender_test.go
|
||||
|
||||
package webhook_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/notify/webhook"
|
||||
)
|
||||
|
||||
const (
|
||||
testTokenID = "tokwh01abcde"
|
||||
testManageID = "abcd1111-2222-3333-4444-555555555555"
|
||||
)
|
||||
|
||||
func sampleEvent() *event.Event {
|
||||
ua := "TestUA/1.0"
|
||||
city := "Toronto"
|
||||
country := "CA"
|
||||
asnOrg := "Test, Inc."
|
||||
asn := 12345
|
||||
return &event.Event{
|
||||
ID: 7,
|
||||
TokenID: testTokenID,
|
||||
TriggeredAt: time.Date(2026, 5, 14, 12, 30, 0, 0, time.UTC),
|
||||
SourceIP: "203.0.113.45",
|
||||
UserAgent: &ua,
|
||||
GeoCity: &city,
|
||||
GeoCountry: &country,
|
||||
GeoASN: &asn,
|
||||
GeoASNOrg: &asnOrg,
|
||||
Extra: json.RawMessage(`{"custom":"value"}`),
|
||||
}
|
||||
}
|
||||
|
||||
func sampleInfo(webhookURL string) event.NotifyInfo {
|
||||
return event.NotifyInfo{
|
||||
TokenID: testTokenID,
|
||||
ManageID: testManageID,
|
||||
Type: "envfile",
|
||||
Memo: "prod-creds",
|
||||
AlertChannel: "webhook",
|
||||
WebhookURL: webhookURL,
|
||||
}
|
||||
}
|
||||
|
||||
type capture struct {
|
||||
calls atomic.Int32
|
||||
lastBody atomic.Value
|
||||
lastSig atomic.Value
|
||||
}
|
||||
|
||||
func newCaptureServer(
|
||||
t *testing.T,
|
||||
handler http.HandlerFunc,
|
||||
) (*httptest.Server, *capture) {
|
||||
t.Helper()
|
||||
c := &capture{}
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
c.calls.Add(1)
|
||||
body, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
t.Logf("read body: %v", err)
|
||||
return
|
||||
}
|
||||
c.lastBody.Store(body)
|
||||
c.lastSig.Store(r.Header.Get("X-Canary-Signature"))
|
||||
handler(w, r)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
return srv, c
|
||||
}
|
||||
|
||||
func loadBody(t *testing.T, c *capture) []byte {
|
||||
t.Helper()
|
||||
raw, ok := c.lastBody.Load().([]byte)
|
||||
require.True(t, ok, "no body captured")
|
||||
return raw
|
||||
}
|
||||
|
||||
func newSender(t *testing.T, opts ...webhook.Option) *webhook.Sender {
|
||||
t.Helper()
|
||||
return webhook.NewSender(webhook.Config{
|
||||
ManageURL: "https://canary.example.com",
|
||||
}, opts...)
|
||||
}
|
||||
|
||||
func TestSender_Channel(t *testing.T) {
|
||||
t.Parallel()
|
||||
require.Equal(t, "webhook", newSender(t).Channel())
|
||||
}
|
||||
|
||||
func TestSender_Send_PostsToProvidedURL(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
},
|
||||
)
|
||||
s := newSender(t)
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(srv.URL), sampleEvent()),
|
||||
)
|
||||
require.Equal(t, int32(1), cap.calls.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_BodyEnvelopeShape(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
},
|
||||
)
|
||||
s := newSender(t)
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(srv.URL), sampleEvent()),
|
||||
)
|
||||
|
||||
var env map[string]any
|
||||
require.NoError(t, json.Unmarshal(loadBody(t, cap), &env))
|
||||
|
||||
require.Equal(t, "1", env["version"])
|
||||
require.Equal(t, "canary.triggered", env["event"])
|
||||
|
||||
tok, ok := env["token"].(map[string]any)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, testTokenID, tok["id"])
|
||||
require.Equal(t, "envfile", tok["type"])
|
||||
require.Equal(t, "prod-creds", tok["memo"])
|
||||
require.Equal(
|
||||
t,
|
||||
"https://canary.example.com/m/"+testManageID,
|
||||
tok["manage_url"],
|
||||
)
|
||||
|
||||
trig, ok := env["trigger"].(map[string]any)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "203.0.113.45", trig["source_ip"])
|
||||
require.Equal(t, "TestUA/1.0", trig["user_agent"])
|
||||
require.NotEmpty(t, trig["triggered_at"])
|
||||
|
||||
geo, ok := trig["geo"].(map[string]any)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "CA", geo["country"])
|
||||
require.Equal(t, "Toronto", geo["city"])
|
||||
require.Equal(t, "Test, Inc.", geo["asn_org"])
|
||||
|
||||
extra, ok := trig["extra"].(map[string]any)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "value", extra["custom"])
|
||||
}
|
||||
|
||||
func TestSender_Send_NoSignatureWithoutSecret(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
},
|
||||
)
|
||||
s := newSender(t)
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(srv.URL), sampleEvent()),
|
||||
)
|
||||
require.Empty(t, cap.lastSig.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_HMACSignatureWhenSecretSet(t *testing.T) {
|
||||
t.Parallel()
|
||||
const secret = "topsecret"
|
||||
srv, cap := newCaptureServer(
|
||||
t,
|
||||
func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
},
|
||||
)
|
||||
s := webhook.NewSender(webhook.Config{
|
||||
ManageURL: "https://canary.example.com",
|
||||
HMACSecret: secret,
|
||||
})
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(srv.URL), sampleEvent()),
|
||||
)
|
||||
|
||||
body := loadBody(t, cap)
|
||||
mac := hmac.New(sha256.New, []byte(secret))
|
||||
if _, err := mac.Write(body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := "sha256=" + hex.EncodeToString(mac.Sum(nil))
|
||||
|
||||
require.Equal(t, want, cap.lastSig.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_EmptyURLReturnsConfigErr(t *testing.T) {
|
||||
t.Parallel()
|
||||
s := newSender(t)
|
||||
err := s.Send(context.Background(), sampleInfo(""), sampleEvent())
|
||||
require.ErrorIs(t, err, webhook.ErrChannelNotConfigured)
|
||||
}
|
||||
|
||||
func TestSender_Send_RejectsNonHTTPScheme(t *testing.T) {
|
||||
t.Parallel()
|
||||
cases := []string{
|
||||
"ftp://example.com/hook",
|
||||
"file:///etc/passwd",
|
||||
"javascript:alert(1)",
|
||||
"not-a-url",
|
||||
}
|
||||
for _, u := range cases {
|
||||
t.Run(u, func(t *testing.T) {
|
||||
s := newSender(t)
|
||||
err := s.Send(context.Background(), sampleInfo(u), sampleEvent())
|
||||
require.ErrorIs(t, err, webhook.ErrInvalidWebhookURL)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSender_Send_RejectsURLWithoutHost(t *testing.T) {
|
||||
t.Parallel()
|
||||
s := newSender(t)
|
||||
err := s.Send(
|
||||
context.Background(),
|
||||
sampleInfo("http:///nohost"),
|
||||
sampleEvent(),
|
||||
)
|
||||
require.ErrorIs(t, err, webhook.ErrInvalidWebhookURL)
|
||||
}
|
||||
|
||||
func TestSender_Send_RejectsURLWithUserInfo(t *testing.T) {
|
||||
t.Parallel()
|
||||
s := newSender(t)
|
||||
err := s.Send(
|
||||
context.Background(),
|
||||
sampleInfo("https://user:pass@example.com/h"),
|
||||
sampleEvent(),
|
||||
)
|
||||
require.ErrorIs(t, err, webhook.ErrInvalidWebhookURL)
|
||||
}
|
||||
|
||||
func TestSender_Send_RetriesOn5xxThenSucceeds(t *testing.T) {
|
||||
t.Parallel()
|
||||
var attempts atomic.Int32
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
n := attempts.Add(1)
|
||||
if n == 1 {
|
||||
w.WriteHeader(http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
s := newSender(t,
|
||||
webhook.WithMaxTries(3),
|
||||
webhook.WithMaxElapsed(2*time.Second),
|
||||
webhook.WithInitialInterval(5*time.Millisecond),
|
||||
)
|
||||
require.NoError(
|
||||
t,
|
||||
s.Send(context.Background(), sampleInfo(srv.URL), sampleEvent()),
|
||||
)
|
||||
require.Equal(t, int32(2), attempts.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_PermanentOn4xx(t *testing.T) {
|
||||
t.Parallel()
|
||||
var attempts atomic.Int32
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
attempts.Add(1)
|
||||
w.WriteHeader(http.StatusForbidden)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
s := newSender(t,
|
||||
webhook.WithMaxTries(5),
|
||||
webhook.WithInitialInterval(5*time.Millisecond),
|
||||
)
|
||||
err := s.Send(context.Background(), sampleInfo(srv.URL), sampleEvent())
|
||||
require.Error(t, err)
|
||||
require.Equal(t, int32(1), attempts.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_AbortsAfterMaxTries(t *testing.T) {
|
||||
t.Parallel()
|
||||
var attempts atomic.Int32
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
attempts.Add(1)
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
s := newSender(t,
|
||||
webhook.WithMaxTries(3),
|
||||
webhook.WithMaxElapsed(5*time.Second),
|
||||
webhook.WithInitialInterval(2*time.Millisecond),
|
||||
)
|
||||
err := s.Send(context.Background(), sampleInfo(srv.URL), sampleEvent())
|
||||
require.Error(t, err)
|
||||
require.Equal(t, int32(3), attempts.Load())
|
||||
}
|
||||
|
||||
func TestSender_Send_RespectsContextCancel(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv := httptest.NewServer(
|
||||
http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {
|
||||
time.Sleep(2 * time.Second)
|
||||
}),
|
||||
)
|
||||
t.Cleanup(srv.Close)
|
||||
s := newSender(t,
|
||||
webhook.WithMaxTries(3),
|
||||
webhook.WithInitialInterval(5*time.Millisecond),
|
||||
)
|
||||
ctx, cancel := context.WithTimeout(
|
||||
context.Background(),
|
||||
100*time.Millisecond,
|
||||
)
|
||||
defer cancel()
|
||||
err := s.Send(ctx, sampleInfo(srv.URL), sampleEvent())
|
||||
require.Error(t, err)
|
||||
require.True(
|
||||
t,
|
||||
errors.Is(err, context.DeadlineExceeded) ||
|
||||
errors.Is(err, context.Canceled),
|
||||
"expected context error, got %v",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
|
@ -50,3 +50,23 @@ type Generator interface {
|
|||
r *http.Request,
|
||||
) (*event.Event, *TriggerResponse, error)
|
||||
}
|
||||
|
||||
func (t *Token) NotifyInfo() event.NotifyInfo {
|
||||
return event.NotifyInfo{
|
||||
TokenID: t.ID,
|
||||
ManageID: t.ManageID,
|
||||
Type: string(t.Type),
|
||||
Memo: t.Memo,
|
||||
AlertChannel: string(t.AlertChannel),
|
||||
TelegramBot: derefString(t.TelegramBot),
|
||||
TelegramChat: derefString(t.TelegramChat),
|
||||
WebhookURL: derefString(t.WebhookURL),
|
||||
}
|
||||
}
|
||||
|
||||
func derefString(s *string) string {
|
||||
if s == nil {
|
||||
return ""
|
||||
}
|
||||
return *s
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in New Issue