Skip to content

Commit a0cbe9e

Browse files
authored
Merge pull request #4 from LumeWeb/feat/tls-get-certificate-throttle
feat: add tls_get_certificate event with throttled status webhooks
2 parents 20f34f1 + d85eff7 commit a0cbe9e

9 files changed

Lines changed: 723 additions & 30 deletions

File tree

config.go

Lines changed: 24 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2,18 +2,21 @@ package certwebhook
22

33
import (
44
"os"
5+
"time"
56
)
67

78
const (
89
GatewaySecretHeader = "X-Gateway-Secret"
910

10-
EnvGatewaySecret = "GATEWAY_SECRET"
11-
EnvPortalURL = "PORTAL_URL"
11+
EnvGatewaySecret = "GATEWAY_SECRET"
12+
EnvPortalURL = "PORTAL_URL"
13+
EnvThrottleInterval = "THROTTLE_INTERVAL"
1214
)
1315

1416
type Config struct {
15-
PortalURL string
16-
GatewaySecret string
17+
PortalURL string
18+
GatewaySecret string
19+
ThrottleInterval string
1720
}
1821

1922
func (c *Config) Provision() {
@@ -23,6 +26,23 @@ func (c *Config) Provision() {
2326
if c.GatewaySecret == "" {
2427
c.GatewaySecret = os.Getenv(EnvGatewaySecret)
2528
}
29+
if c.ThrottleInterval == "" {
30+
c.ThrottleInterval = os.Getenv(EnvThrottleInterval)
31+
}
32+
}
33+
34+
func (c *Config) throttleInterval() time.Duration {
35+
if c.ThrottleInterval == "" {
36+
return defaultThrottleInterval
37+
}
38+
d, err := time.ParseDuration(c.ThrottleInterval)
39+
if err != nil {
40+
return defaultThrottleInterval
41+
}
42+
if d <= 0 {
43+
return defaultThrottleInterval
44+
}
45+
return d
2646
}
2747

2848
func (c *Config) Validate() error {

config_test.go

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,54 @@
1+
package certwebhook
2+
3+
import (
4+
"testing"
5+
"time"
6+
7+
"github.com/stretchr/testify/assert"
8+
)
9+
10+
func TestConfigThrottleInterval_Default(t *testing.T) {
11+
c := &Config{}
12+
assert.Equal(t, defaultThrottleInterval, c.throttleInterval())
13+
}
14+
15+
func TestConfigThrottleInterval_ValidDuration(t *testing.T) {
16+
c := &Config{ThrottleInterval: "10m"}
17+
assert.Equal(t, 10*time.Minute, c.throttleInterval())
18+
}
19+
20+
func TestConfigThrottleInterval_Seconds(t *testing.T) {
21+
c := &Config{ThrottleInterval: "30s"}
22+
assert.Equal(t, 30*time.Second, c.throttleInterval())
23+
}
24+
25+
func TestConfigThrottleInterval_InvalidString(t *testing.T) {
26+
c := &Config{ThrottleInterval: "not a duration"}
27+
assert.Equal(t, defaultThrottleInterval, c.throttleInterval())
28+
}
29+
30+
func TestConfigThrottleInterval_Zero(t *testing.T) {
31+
c := &Config{ThrottleInterval: "0"}
32+
assert.Equal(t, defaultThrottleInterval, c.throttleInterval())
33+
}
34+
35+
func TestConfigThrottleInterval_Negative(t *testing.T) {
36+
c := &Config{ThrottleInterval: "-5m"}
37+
assert.Equal(t, defaultThrottleInterval, c.throttleInterval())
38+
}
39+
40+
func TestConfigThrottleInterval_EnvFallback(t *testing.T) {
41+
t.Setenv(EnvThrottleInterval, "15m")
42+
c := &Config{}
43+
c.Provision()
44+
assert.Equal(t, "15m", c.ThrottleInterval)
45+
assert.Equal(t, 15*time.Minute, c.throttleInterval())
46+
}
47+
48+
func TestConfigThrottleInterval_ExplicitOverEnv(t *testing.T) {
49+
t.Setenv(EnvThrottleInterval, "15m")
50+
c := &Config{ThrottleInterval: "2m"}
51+
c.Provision()
52+
assert.Equal(t, "2m", c.ThrottleInterval)
53+
assert.Equal(t, 2*time.Minute, c.throttleInterval())
54+
}

events.go

Lines changed: 69 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -11,54 +11,44 @@ import (
1111
"go.uber.org/zap"
1212
)
1313

14-
// Caddy event names for certificate lifecycle
1514
const (
16-
// EventCertObtained is fired when a certificate is first obtained
1715
EventCertObtained = "cert_obtained"
1816

19-
// EventCertRenewed is fired when a certificate is renewed
2017
EventCertRenewed = "cert_renewed"
2118

22-
// EventCertExpired is fired when a certificate expires
2319
EventCertExpired = "cert_expired"
20+
21+
EventTLSGetCertificate = "tls_get_certificate"
2422
)
2523

26-
// Log message constants
2724
const (
2825
LogMsgEventsAppNotAvailable = "events app not available"
2926
LogMsgEventsAppNotExpectedType = "events app is not of expected type"
3027
LogMsgSubscribedToCertObtained = "subscribed to cert_obtained event"
3128
LogMsgSubscribedToCertRenewed = "subscribed to cert_renewed event"
3229
LogMsgSubscribedToCertExpired = "subscribed to cert_expired event"
30+
LogMsgSubscribedToTLSGetCert = "subscribed to tls_get_certificate event"
3331
LogMsgUnknownEventType = "unknown event type"
3432
LogMsgFailedToExtractEventData = "failed to extract event data"
3533
LogMsgFailedToMapEventToStatus = "failed to map event to status"
3634
LogMsgCertificateEventProcessed = "certificate event processed"
35+
LogMsgTLSGetCertThrottled = "tls_get_certificate event throttled"
36+
LogMsgTLSGetCertProcessed = "tls_get_certificate event processed"
3737
)
3838

39-
// EventData represents certificate event data from Caddy
4039
type EventData struct {
41-
// Domain is the certificate domain name
4240
Domain string `json:"domain"`
4341

44-
// Timestamp is when the event occurred (ISO 8601 format)
4542
Timestamp string `json:"timestamp"`
4643

47-
// Error is any error that occurred during certificate operation
4844
Error string `json:"error,omitempty"`
4945

50-
// Raw is the raw event data for debugging
5146
Raw map[string]any `json:"-"`
5247

53-
// EventType is the type of event (cert_obtained, cert_renewed, cert_expired)
5448
EventType string `json:"event_type"`
5549
}
5650

57-
// Event handlers for certificate lifecycle events
58-
59-
// subscribeToEvents registers handlers for certificate events
6051
func (h *CertWebhookApp) subscribeToEvents(ctx caddy.Context) error {
61-
// Get the events app from context
6252
eventsAppIface, err := ctx.App("events")
6353
if err != nil {
6454
h.logger.Warn(LogMsgEventsAppNotAvailable, zap.Error(err))
@@ -71,10 +61,8 @@ func (h *CertWebhookApp) subscribeToEvents(ctx caddy.Context) error {
7161
return nil
7262
}
7363

74-
// Store reference for cleanup
7564
h.eventsApp = eventsApp
7665

77-
// Subscribe to cert_obtained event
7866
if err := eventsApp.On(EventCertObtained, h); err != nil {
7967
return fmt.Errorf("failed to subscribe to cert_obtained event: %w", err)
8068
}
@@ -90,23 +78,28 @@ func (h *CertWebhookApp) subscribeToEvents(ctx caddy.Context) error {
9078
}
9179
h.logger.Info(LogMsgSubscribedToCertExpired)
9280

81+
if err := eventsApp.On(EventTLSGetCertificate, h); err != nil {
82+
return fmt.Errorf("failed to subscribe to tls_get_certificate event: %w", err)
83+
}
84+
h.logger.Info(LogMsgSubscribedToTLSGetCert)
85+
9386
return nil
9487
}
9588

96-
// Handle implements caddyevents.Handler interface to process certificate events
9789
func (h *CertWebhookApp) Handle(ctx context.Context, data caddy.Event) error {
9890
eventName := data.Name()
9991

10092
switch eventName {
10193
case EventCertObtained, EventCertRenewed, EventCertExpired:
10294
return h.handleCertEvent(eventName, data)
95+
case EventTLSGetCertificate:
96+
return h.handleTLSGetCertificateEvent(data)
10397
default:
10498
h.logger.Warn(LogMsgUnknownEventType, zap.String("event", eventName))
10599
return nil
106100
}
107101
}
108102

109-
// handleCertEvent processes certificate lifecycle events (obtained, renewed, expired)
110103
func (h *CertWebhookApp) handleCertEvent(eventType string, data caddy.Event) error {
111104
h.logger.Debug("handling cert event", zap.String("event_type", eventType))
112105

@@ -126,6 +119,14 @@ func (h *CertWebhookApp) handleCertEvent(eventType string, data caddy.Event) err
126119
return err
127120
}
128121

122+
if !h.shouldSend(eventData.Domain, status) {
123+
h.logger.Debug("cert event throttled",
124+
zap.String("event_type", eventType),
125+
zap.String("domain", eventData.Domain),
126+
zap.String("status", string(status)))
127+
return nil
128+
}
129+
129130
h.logger.Info(LogMsgCertificateEventProcessed,
130131
zap.String("event_type", eventType),
131132
zap.String("domain", eventData.Domain),
@@ -135,6 +136,55 @@ func (h *CertWebhookApp) handleCertEvent(eventType string, data caddy.Event) err
135136
return h.sendWebhook(eventData.Domain, status, eventData.Error, eventData.Timestamp)
136137
}
137138

139+
func (h *CertWebhookApp) handleTLSGetCertificateEvent(data caddy.Event) error {
140+
domain, err := extractDomainFromClientHello(data)
141+
if err != nil {
142+
h.logger.Debug("failed to extract domain from tls_get_certificate", zap.Error(err))
143+
return nil
144+
}
145+
146+
status := h.certStatusFn(domain)
147+
148+
if !h.shouldSend(domain, status) {
149+
h.logger.Debug(LogMsgTLSGetCertThrottled,
150+
zap.String("domain", domain),
151+
zap.String("status", string(status)))
152+
return nil
153+
}
154+
155+
timestamp := time.Now().UTC().Format(time.RFC3339)
156+
157+
h.logger.Info(LogMsgTLSGetCertProcessed,
158+
zap.String("domain", domain),
159+
zap.String("status", string(status)),
160+
zap.String("timestamp", timestamp))
161+
162+
return h.sendWebhook(domain, status, "", timestamp)
163+
}
164+
165+
func extractDomainFromClientHello(event caddy.Event) (string, error) {
166+
if len(event.Data) == 0 {
167+
return "", fmt.Errorf("event data is empty")
168+
}
169+
170+
raw := make(map[string]any)
171+
if err := decodeJSON(event.Data, &raw); err != nil {
172+
return "", fmt.Errorf("failed to unmarshal event data: %w", err)
173+
}
174+
175+
clientHello, ok := raw["client_hello"].(map[string]any)
176+
if !ok {
177+
return "", fmt.Errorf("client_hello not found in event data")
178+
}
179+
180+
serverName, ok := clientHello["ServerName"].(string)
181+
if !ok || serverName == "" {
182+
return "", fmt.Errorf("ServerName not found in client_hello")
183+
}
184+
185+
return serverName, nil
186+
}
187+
138188
// extractEventData extracts domain, timestamp, and error information from Caddy event data
139189
func (h *CertWebhookApp) extractEventData(eventType string, event caddy.Event) (*EventData, error) {
140190
data := &EventData{

events_test.go

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
1+
package certwebhook
2+
3+
import (
4+
"testing"
5+
6+
"github.com/caddyserver/caddy/v2"
7+
)
8+
9+
func TestExtractDomainFromClientHello(t *testing.T) {
10+
tests := []struct {
11+
name string
12+
data map[string]any
13+
want string
14+
wantErr bool
15+
}{
16+
{
17+
name: "valid client_hello with ServerName",
18+
data: map[string]any{
19+
"client_hello": map[string]any{
20+
"ServerName": "example.com",
21+
},
22+
},
23+
want: "example.com",
24+
wantErr: false,
25+
},
26+
{
27+
name: "empty ServerName",
28+
data: map[string]any{
29+
"client_hello": map[string]any{
30+
"ServerName": "",
31+
},
32+
},
33+
want: "",
34+
wantErr: true,
35+
},
36+
{
37+
name: "empty event data",
38+
data: map[string]any{},
39+
want: "",
40+
wantErr: true,
41+
},
42+
{
43+
name: "client_hello not a map",
44+
data: map[string]any{
45+
"client_hello": "not a map",
46+
},
47+
want: "",
48+
wantErr: true,
49+
},
50+
{
51+
name: "ServerName not a string",
52+
data: map[string]any{
53+
"client_hello": map[string]any{
54+
"ServerName": 123,
55+
},
56+
},
57+
want: "",
58+
wantErr: true,
59+
},
60+
}
61+
62+
for _, tt := range tests {
63+
t.Run(tt.name, func(t *testing.T) {
64+
event := caddy.Event{Data: tt.data}
65+
got, err := extractDomainFromClientHello(event)
66+
if (err != nil) != tt.wantErr {
67+
t.Errorf("extractDomainFromClientHello() error = %v, wantErr %v", err, tt.wantErr)
68+
return
69+
}
70+
if got != tt.want {
71+
t.Errorf("extractDomainFromClientHello() = %v, want %v", got, tt.want)
72+
}
73+
})
74+
}
75+
}

0 commit comments

Comments
 (0)