diff --git a/events/event_stream.go b/events/event_stream.go index ef71959..891382c 100644 --- a/events/event_stream.go +++ b/events/event_stream.go @@ -83,6 +83,8 @@ func eventStreamProcessor(ctx context.Context, eventChan <-chan mtglib.Event, ob observer.EventIPBlocklisted(typedEvt) case mtglib.EventConcurrencyLimited: observer.EventConcurrencyLimited(typedEvt) + case mtglib.EventReplayAttack: + observer.EventReplayAttack(typedEvt) } } } diff --git a/events/event_stream_test.go b/events/event_stream_test.go index 83e6b77..179e29f 100644 --- a/events/event_stream_test.go +++ b/events/event_stream_test.go @@ -206,6 +206,28 @@ func (suite *EventStreamTestSuite) TestEventIPBlocklisted() { time.Sleep(100 * time.Millisecond) } +func (suite *EventStreamTestSuite) TestEventReplayAttack() { + evt := mtglib.EventReplayAttack{ + CreatedAt: time.Now(), + ConnID: "CONNID", + } + + for _, v := range []*ObserverMock{suite.observerMock1, suite.observerMock2} { + v. + On("EventReplayAttack", mock.Anything). + Once(). + Run(func(args mock.Arguments) { + caught := args.Get(0).(mtglib.EventReplayAttack) + + suite.Equal(evt.CreatedAt, caught.CreatedAt) + suite.Equal(evt.StreamID(), caught.StreamID()) + }) + } + + suite.stream.Send(suite.ctx, evt) + time.Sleep(100 * time.Millisecond) +} + func (suite *EventStreamTestSuite) TearDownTest() { suite.stream.Shutdown() suite.ctxCancel() diff --git a/events/init.go b/events/init.go index b6dbd22..1204652 100644 --- a/events/init.go +++ b/events/init.go @@ -10,6 +10,7 @@ type Observer interface { EventTraffic(mtglib.EventTraffic) EventConcurrencyLimited(mtglib.EventConcurrencyLimited) EventIPBlocklisted(mtglib.EventIPBlocklisted) + EventReplayAttack(mtglib.EventReplayAttack) Shutdown() } diff --git a/events/init_test.go b/events/init_test.go index 5fccfad..1b2e0cd 100644 --- a/events/init_test.go +++ b/events/init_test.go @@ -37,6 +37,10 @@ func (o *ObserverMock) EventIPBlocklisted(evt mtglib.EventIPBlocklisted) { o.Called(evt) } +func (o *ObserverMock) EventReplayAttack(evt mtglib.EventReplayAttack) { + o.Called(evt) +} + func (o *ObserverMock) Shutdown() { o.Called() } diff --git a/events/multi_observer.go b/events/multi_observer.go index cd74567..eeeacae 100644 --- a/events/multi_observer.go +++ b/events/multi_observer.go @@ -115,6 +115,21 @@ func (m multiObserver) EventIPBlocklisted(evt mtglib.EventIPBlocklisted) { wg.Wait() } +func (m multiObserver) EventReplayAttack(evt mtglib.EventReplayAttack) { + wg := &sync.WaitGroup{} + wg.Add(len(m.observers)) + + for _, v := range m.observers { + go func(obs Observer) { + defer wg.Done() + + obs.EventReplayAttack(evt) + }(v) + } + + wg.Wait() +} + func (m multiObserver) Shutdown() { for _, v := range m.observers { v.Shutdown() diff --git a/events/noop.go b/events/noop.go index 27594d8..afc2c66 100644 --- a/events/noop.go +++ b/events/noop.go @@ -24,6 +24,7 @@ func (n noopObserver) EventTraffic(_ mtglib.EventTraffic) func (n noopObserver) EventFinish(_ mtglib.EventFinish) {} func (n noopObserver) EventConcurrencyLimited(_ mtglib.EventConcurrencyLimited) {} func (n noopObserver) EventIPBlocklisted(_ mtglib.EventIPBlocklisted) {} +func (n noopObserver) EventReplayAttack(_ mtglib.EventReplayAttack) {} func (n noopObserver) Shutdown() {} func NewNoopObserver() Observer { diff --git a/events/noop_test.go b/events/noop_test.go index 9a1d2bd..024febc 100644 --- a/events/noop_test.go +++ b/events/noop_test.go @@ -45,11 +45,17 @@ func (suite *NoopTestSuite) SetupSuite() { CreatedAt: time.Now(), ConnID: "connID", }, - "concurrency-limited": mtglib.EventConcurrencyLimited{}, + "concurrency-limited": mtglib.EventConcurrencyLimited{ + CreatedAt: time.Now(), + }, "ip-blacklisted": mtglib.EventIPBlocklisted{ RemoteIP: net.ParseIP("10.0.0.10"), CreatedAt: time.Now(), }, + "replay-attack": mtglib.EventReplayAttack{ + CreatedAt: time.Now(), + ConnID: "connID", + }, } suite.ctx = context.Background() } @@ -88,6 +94,8 @@ func (suite *NoopTestSuite) TestObserver() { observer.EventConcurrencyLimited(typedEvt) case mtglib.EventIPBlocklisted: observer.EventIPBlocklisted(typedEvt) + case mtglib.EventReplayAttack: + observer.EventReplayAttack(typedEvt) } }) } diff --git a/mtglib/events.go b/mtglib/events.go index e7a5eae..e4c57d1 100644 --- a/mtglib/events.go +++ b/mtglib/events.go @@ -99,3 +99,16 @@ func (e EventIPBlocklisted) StreamID() string { func (e EventIPBlocklisted) Timestamp() time.Time { return e.CreatedAt } + +type EventReplayAttack struct { + CreatedAt time.Time + ConnID string +} + +func (e EventReplayAttack) StreamID() string { + return e.ConnID +} + +func (e EventReplayAttack) Timestamp() time.Time { + return e.CreatedAt +} diff --git a/mtglib/events_test.go b/mtglib/events_test.go index c3a2fb9..1f4814e 100644 --- a/mtglib/events_test.go +++ b/mtglib/events_test.go @@ -87,6 +87,16 @@ func (suite *EventsTestSuite) TestEventIPBlocklisted() { suite.WithinDuration(time.Now(), evt.Timestamp(), 10*time.Millisecond) } +func (suite *EventsTestSuite) TestEventReplayAttack() { + evt := mtglib.EventReplayAttack{ + CreatedAt: time.Now(), + ConnID: "CONNID", + } + + suite.Equal("CONNID", evt.StreamID()) + suite.WithinDuration(time.Now(), evt.Timestamp(), 10*time.Millisecond) +} + func TestEvents(t *testing.T) { t.Parallel() suite.Run(t, &EventsTestSuite{}) diff --git a/mtglib/proxy.go b/mtglib/proxy.go index 731c68c..e0d5974 100644 --- a/mtglib/proxy.go +++ b/mtglib/proxy.go @@ -166,6 +166,10 @@ func (p *Proxy) doFakeTLSHandshake(ctx *streamContext) bool { if p.antiReplayCache.SeenBefore(hello.SessionID) { p.logger.Warning("replay attack has been detected!") + p.eventStream.Send(p.ctx, EventReplayAttack{ + CreatedAt: time.Now(), + ConnID: ctx.connID, + }) p.doDomainFronting(ctx, rewind) return false diff --git a/stats/prometheus.go b/stats/prometheus.go index 2e7fe10..aa2933d 100644 --- a/stats/prometheus.go +++ b/stats/prometheus.go @@ -110,10 +110,14 @@ func (p prometheusProcessor) EventConcurrencyLimited(_ mtglib.EventConcurrencyLi p.factory.metricConcurrencyLimited.Inc() } -func (p prometheusProcessor) EventIPBlocklisted(evt mtglib.EventIPBlocklisted) { +func (p prometheusProcessor) EventIPBlocklisted(_ mtglib.EventIPBlocklisted) { p.factory.metricIPBlocklisted.Inc() } +func (p prometheusProcessor) EventReplayAttack(_ mtglib.EventReplayAttack) { + p.factory.metricReplayAttacks.Inc() +} + func (p prometheusProcessor) Shutdown() { for _, v := range p.streams { releaseStreamInfo(v) diff --git a/stats/prometheus_test.go b/stats/prometheus_test.go index 242fe75..abb3bec 100644 --- a/stats/prometheus_test.go +++ b/stats/prometheus_test.go @@ -198,6 +198,19 @@ func (suite *PrometheusTestSuite) TestEventIPBlocklisted() { suite.Contains(data, `mtg_ip_blocklisted 1`) } +func (suite *PrometheusTestSuite) TestEventReplayAttack() { + suite.prometheus.EventReplayAttack(mtglib.EventReplayAttack{ + CreatedAt: time.Now(), + ConnID: "connID", + }) + + time.Sleep(100 * time.Millisecond) + + data, err := suite.Get() + suite.NoError(err) + suite.Contains(data, `mtg_replay_attacks 1`) +} + func TestPrometheus(t *testing.T) { t.Parallel() suite.Run(t, &PrometheusTestSuite{}) diff --git a/stats/statsd.go b/stats/statsd.go index 5bb2a38..a1f0ca4 100644 --- a/stats/statsd.go +++ b/stats/statsd.go @@ -114,10 +114,14 @@ func (s statsdProcessor) EventConcurrencyLimited(_ mtglib.EventConcurrencyLimite s.client.Incr(MetricConcurrencyLimited, 1) } -func (s statsdProcessor) EventIPBlocklisted(evt mtglib.EventIPBlocklisted) { +func (s statsdProcessor) EventIPBlocklisted(_ mtglib.EventIPBlocklisted) { s.client.Incr(MetricIPBlocklisted, 1) } +func (s statsdProcessor) EventReplayAttack(_ mtglib.EventReplayAttack) { + s.client.Incr(MetricReplayAttacks, 1) +} + func (s statsdProcessor) Shutdown() { now := time.Now() events := make([]mtglib.EventFinish, 0, len(s.streams)) diff --git a/stats/statsd_test.go b/stats/statsd_test.go index ee55ad8..02b34ec 100644 --- a/stats/statsd_test.go +++ b/stats/statsd_test.go @@ -228,6 +228,16 @@ func (suite *StatsdTestSuite) TestEventIPBlocklisted() { suite.Equal("mtg.ip_blocklisted:1|c", suite.statsdServer.String()) } +func (suite *StatsdTestSuite) TestEventReplayAttack() { + suite.statsd.EventReplayAttack(mtglib.EventReplayAttack{ + CreatedAt: time.Now(), + ConnID: "connID", + }) + + time.Sleep(statsdSleepTime) + suite.Equal("mtg.replay_attacks:1|c", suite.statsdServer.String()) +} + func TestStatsd(t *testing.T) { t.Parallel() suite.Run(t, &StatsdTestSuite{})