feat(rest,config): wire CORS, real readiness, TLS options, error hygiene
- CORS config is now actually applied to the router (the middleware existed but was never wired); preflight returns 204 for allowed origins and 403 with no CORS headers for disallowed ones. - /readyz runs real dependency checks (queue, auth database) and returns 503 with per-component detail when not ready; exported handlers support dedicated health listeners. - Optional server.tls (cert_file/key_file) for REST and gRPC, validated at config load. - 5xx responses no longer echo internal error details; not-found and already-sent map to 404/409 on cancel/retry. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -171,13 +171,13 @@ func TestCORSMiddleware_PreflightRequest(t *testing.T) {
|
|||||||
{
|
{
|
||||||
name: "preflight from allowed origin",
|
name: "preflight from allowed origin",
|
||||||
origin: "https://example.com",
|
origin: "https://example.com",
|
||||||
expectStatus: http.StatusOK,
|
expectStatus: http.StatusNoContent,
|
||||||
expectHeaders: true,
|
expectHeaders: true,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "preflight from blocked origin",
|
name: "preflight from blocked origin",
|
||||||
origin: "https://malicious.com",
|
origin: "https://malicious.com",
|
||||||
expectStatus: http.StatusOK,
|
expectStatus: http.StatusForbidden,
|
||||||
expectHeaders: false,
|
expectHeaders: false,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -191,7 +191,8 @@ func TestCORSMiddleware_PreflightRequest(t *testing.T) {
|
|||||||
rec := httptest.NewRecorder()
|
rec := httptest.NewRecorder()
|
||||||
handler.ServeHTTP(rec, req)
|
handler.ServeHTTP(rec, req)
|
||||||
|
|
||||||
// Preflight should always return 200 OK
|
// Allowed preflights succeed with 204; blocked ones get 403
|
||||||
|
// with no CORS headers so the browser rejects the request.
|
||||||
if rec.Code != tt.expectStatus {
|
if rec.Code != tt.expectStatus {
|
||||||
t.Errorf("status = %v, want %v", rec.Code, tt.expectStatus)
|
t.Errorf("status = %v, want %v", rec.Code, tt.expectStatus)
|
||||||
}
|
}
|
||||||
|
|||||||
+20
-4
@@ -2,6 +2,7 @@ package rest
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
@@ -160,7 +161,7 @@ func (h *Handler) CancelNotification(w http.ResponseWriter, r *http.Request) {
|
|||||||
id := vars["id"]
|
id := vars["id"]
|
||||||
|
|
||||||
if err := h.service.CancelNotification(r.Context(), id); err != nil {
|
if err := h.service.CancelNotification(r.Context(), id); err != nil {
|
||||||
respondError(w, http.StatusInternalServerError, "failed to cancel notification", err)
|
respondError(w, statusForServiceError(err), "failed to cancel notification", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -177,7 +178,7 @@ func (h *Handler) RetryNotification(w http.ResponseWriter, r *http.Request) {
|
|||||||
|
|
||||||
result, err := h.service.RetryNotification(r.Context(), id)
|
result, err := h.service.RetryNotification(r.Context(), id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
respondError(w, http.StatusInternalServerError, "failed to retry notification", err)
|
respondError(w, statusForServiceError(err), "failed to retry notification", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -186,6 +187,19 @@ func (h *Handler) RetryNotification(w http.ResponseWriter, r *http.Request) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// statusForServiceError maps service-layer sentinel errors to HTTP status
|
||||||
|
// codes; anything unrecognized is an internal error.
|
||||||
|
func statusForServiceError(err error) int {
|
||||||
|
switch {
|
||||||
|
case errors.Is(err, domain.ErrNotificationNotFound):
|
||||||
|
return http.StatusNotFound
|
||||||
|
case errors.Is(err, domain.ErrNotificationAlreadySent):
|
||||||
|
return http.StatusConflict
|
||||||
|
default:
|
||||||
|
return http.StatusInternalServerError
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// GetStats handles GET /api/v1/stats
|
// GetStats handles GET /api/v1/stats
|
||||||
func (h *Handler) GetStats(w http.ResponseWriter, r *http.Request) {
|
func (h *Handler) GetStats(w http.ResponseWriter, r *http.Request) {
|
||||||
stats, err := h.service.GetStats(r.Context())
|
stats, err := h.service.GetStats(r.Context())
|
||||||
@@ -270,10 +284,12 @@ func respondJSON(w http.ResponseWriter, status int, data interface{}) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// respondError sends an error response
|
// respondError sends an error response. Client errors (4xx) include the
|
||||||
|
// underlying detail to help callers fix their request; server errors (5xx)
|
||||||
|
// deliberately do not echo internals — those belong in the server log.
|
||||||
func respondError(w http.ResponseWriter, status int, message string, err error) {
|
func respondError(w http.ResponseWriter, status int, message string, err error) {
|
||||||
errMsg := message
|
errMsg := message
|
||||||
if err != nil {
|
if err != nil && status < http.StatusInternalServerError {
|
||||||
errMsg = message + ": " + err.Error()
|
errMsg = message + ": " + err.Error()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+4
-4
@@ -120,7 +120,7 @@ func (h *KeyManagementHandler) CreateKey(w http.ResponseWriter, r *http.Request)
|
|||||||
apiKey, err := h.keyStore.CreateKey(ctx, req.ClientID, req.Roles, req.RateLimit, expiresInDuration, authCtx.ClientID)
|
apiKey, err := h.keyStore.CreateKey(ctx, req.ClientID, req.Roles, req.RateLimit, expiresInDuration, authCtx.ClientID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
h.logger.Errorf("Failed to create API key: %v", err)
|
h.logger.Errorf("Failed to create API key: %v", err)
|
||||||
h.respondError(w, http.StatusInternalServerError, "Failed to create API key", err.Error())
|
h.respondError(w, http.StatusInternalServerError, "Failed to create API key", "")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -165,7 +165,7 @@ func (h *KeyManagementHandler) ListKeys(w http.ResponseWriter, r *http.Request)
|
|||||||
keys, err := h.keyStore.ListKeys(ctx, clientID)
|
keys, err := h.keyStore.ListKeys(ctx, clientID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
h.logger.Errorf("Failed to list API keys: %v", err)
|
h.logger.Errorf("Failed to list API keys: %v", err)
|
||||||
h.respondError(w, http.StatusInternalServerError, "Failed to list API keys", err.Error())
|
h.respondError(w, http.StatusInternalServerError, "Failed to list API keys", "")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -218,7 +218,7 @@ func (h *KeyManagementHandler) RevokeKey(w http.ResponseWriter, r *http.Request)
|
|||||||
h.respondError(w, http.StatusNotFound, "Key not found", "")
|
h.respondError(w, http.StatusNotFound, "Key not found", "")
|
||||||
} else {
|
} else {
|
||||||
h.logger.Errorf("Failed to revoke API key: %v", err)
|
h.logger.Errorf("Failed to revoke API key: %v", err)
|
||||||
h.respondError(w, http.StatusInternalServerError, "Failed to revoke API key", err.Error())
|
h.respondError(w, http.StatusInternalServerError, "Failed to revoke API key", "")
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -291,7 +291,7 @@ func (h *KeyManagementHandler) GetAuditLog(w http.ResponseWriter, r *http.Reques
|
|||||||
logs, err := h.keyStore.GetAuditLogByName(ctx, keyName, limit)
|
logs, err := h.keyStore.GetAuditLogByName(ctx, keyName, limit)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
h.logger.Errorf("Failed to get audit log: %v", err)
|
h.logger.Errorf("Failed to get audit log: %v", err)
|
||||||
h.respondError(w, http.StatusInternalServerError, "Failed to get audit log", err.Error())
|
h.respondError(w, http.StatusInternalServerError, "Failed to get audit log", "")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+100
-16
@@ -1,9 +1,12 @@
|
|||||||
package rest
|
package rest
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/gorilla/mux"
|
"github.com/gorilla/mux"
|
||||||
"github.com/igodwin/notifier/internal/auth"
|
"github.com/igodwin/notifier/internal/auth"
|
||||||
@@ -43,27 +46,50 @@ func DefaultCORSConfig() *CORSConfig {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewRouter creates a new HTTP router with all routes configured
|
// ReadinessCheck reports whether a named dependency is ready. Implementations
|
||||||
func NewRouter(service domain.NotificationService, logger *logging.Logger) *mux.Router {
|
// should be cheap; they run on every /readyz request.
|
||||||
return NewRouterWithAuth(service, logger, nil)
|
type ReadinessCheck func(ctx context.Context) error
|
||||||
|
|
||||||
|
// RouterOptions configures the REST router.
|
||||||
|
type RouterOptions struct {
|
||||||
|
Service domain.NotificationService
|
||||||
|
Logger *logging.Logger
|
||||||
|
AuthStore *auth.APIKeyStore // nil disables authentication
|
||||||
|
KeyStore *auth.HybridKeyStore // nil disables key-management routes
|
||||||
|
CORS *CORSConfig // nil disables CORS headers entirely
|
||||||
|
// Readiness maps a component name (e.g. "queue", "database") to its check.
|
||||||
|
Readiness map[string]ReadinessCheck
|
||||||
|
// Instrument, when set, wraps the router for request instrumentation
|
||||||
|
// (e.g. Prometheus HTTP metrics).
|
||||||
|
Instrument func(http.Handler) http.Handler
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewRouterWithAuth creates a new HTTP router with optional authentication and CORS configuration
|
// NewRouter creates a new HTTP router with all routes configured
|
||||||
|
func NewRouter(service domain.NotificationService, logger *logging.Logger) *mux.Router {
|
||||||
|
return NewRouterWithOptions(RouterOptions{Service: service, Logger: logger})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRouterWithAuth creates a new HTTP router with optional authentication
|
||||||
func NewRouterWithAuth(service domain.NotificationService, logger *logging.Logger, authStore *auth.APIKeyStore) *mux.Router {
|
func NewRouterWithAuth(service domain.NotificationService, logger *logging.Logger, authStore *auth.APIKeyStore) *mux.Router {
|
||||||
return NewRouterWithAuthAndKeyStore(service, logger, authStore, nil)
|
return NewRouterWithOptions(RouterOptions{Service: service, Logger: logger, AuthStore: authStore})
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewRouterWithAuthAndKeyStore creates a new HTTP router with authentication and key management
|
// NewRouterWithAuthAndKeyStore creates a new HTTP router with authentication and key management
|
||||||
func NewRouterWithAuthAndKeyStore(service domain.NotificationService, logger *logging.Logger, authStore *auth.APIKeyStore, keyStore *auth.HybridKeyStore) *mux.Router {
|
func NewRouterWithAuthAndKeyStore(service domain.NotificationService, logger *logging.Logger, authStore *auth.APIKeyStore, keyStore *auth.HybridKeyStore) *mux.Router {
|
||||||
handler := NewHandler(service, logger)
|
return NewRouterWithOptions(RouterOptions{Service: service, Logger: logger, AuthStore: authStore, KeyStore: keyStore})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRouterWithOptions creates the HTTP router from RouterOptions.
|
||||||
|
func NewRouterWithOptions(opts RouterOptions) *mux.Router {
|
||||||
|
handler := NewHandler(opts.Service, opts.Logger)
|
||||||
router := mux.NewRouter()
|
router := mux.NewRouter()
|
||||||
|
|
||||||
// API v1 routes
|
// API v1 routes
|
||||||
v1 := router.PathPrefix("/api/v1").Subrouter()
|
v1 := router.PathPrefix("/api/v1").Subrouter()
|
||||||
|
|
||||||
// Apply authentication middleware if auth store is provided
|
// Apply authentication middleware if auth store is provided
|
||||||
if authStore != nil {
|
if opts.AuthStore != nil {
|
||||||
authMiddleware := auth.NewRESTAuthMiddleware(authStore, logger)
|
authMiddleware := auth.NewRESTAuthMiddleware(opts.AuthStore, opts.Logger)
|
||||||
v1.Use(authMiddleware.Middleware)
|
v1.Use(authMiddleware.Middleware)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -82,8 +108,8 @@ func NewRouterWithAuthAndKeyStore(service domain.NotificationService, logger *lo
|
|||||||
v1.HandleFunc("/notifiers", handler.GetNotifiers).Methods(http.MethodGet)
|
v1.HandleFunc("/notifiers", handler.GetNotifiers).Methods(http.MethodGet)
|
||||||
|
|
||||||
// Key management routes (requires auth and keystore)
|
// Key management routes (requires auth and keystore)
|
||||||
if authStore != nil && keyStore != nil {
|
if opts.AuthStore != nil && opts.KeyStore != nil {
|
||||||
keyHandler := NewKeyManagementHandler(keyStore, logger)
|
keyHandler := NewKeyManagementHandler(opts.KeyStore, opts.Logger)
|
||||||
v1.HandleFunc("/admin/keys", keyHandler.CreateKey).Methods(http.MethodPost)
|
v1.HandleFunc("/admin/keys", keyHandler.CreateKey).Methods(http.MethodPost)
|
||||||
v1.HandleFunc("/admin/keys", keyHandler.ListKeys).Methods(http.MethodGet)
|
v1.HandleFunc("/admin/keys", keyHandler.ListKeys).Methods(http.MethodGet)
|
||||||
v1.HandleFunc("/admin/keys/{name}", keyHandler.RevokeKey).Methods(http.MethodDelete)
|
v1.HandleFunc("/admin/keys/{name}", keyHandler.RevokeKey).Methods(http.MethodDelete)
|
||||||
@@ -91,16 +117,69 @@ func NewRouterWithAuthAndKeyStore(service domain.NotificationService, logger *lo
|
|||||||
v1.HandleFunc("/admin/keys/{name}/audit", keyHandler.GetAuditLog).Methods(http.MethodGet)
|
v1.HandleFunc("/admin/keys/{name}/audit", keyHandler.GetAuditLog).Methods(http.MethodGet)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Health check route (no auth required)
|
// Liveness and readiness routes (no auth required). /health stays a pure
|
||||||
|
// liveness signal; /readyz fails when a dependency is unavailable.
|
||||||
router.HandleFunc("/health", handler.HealthCheck).Methods(http.MethodGet)
|
router.HandleFunc("/health", handler.HealthCheck).Methods(http.MethodGet)
|
||||||
|
router.HandleFunc("/readyz", readinessHandler(opts.Readiness)).Methods(http.MethodGet)
|
||||||
|
|
||||||
// Middleware - logging, request size limit, and CORS
|
// Middleware - CORS (when configured), logging, and request size limits
|
||||||
|
if opts.CORS != nil && len(opts.CORS.AllowedOrigins) > 0 {
|
||||||
|
router.Use(newCORSMiddleware(opts.CORS))
|
||||||
|
}
|
||||||
|
if opts.Instrument != nil {
|
||||||
|
router.Use(mux.MiddlewareFunc(opts.Instrument))
|
||||||
|
}
|
||||||
router.Use(loggingMiddleware)
|
router.Use(loggingMiddleware)
|
||||||
v1.Use(maxBodySizeMiddleware(1 << 20)) // 1 MB limit on API request bodies
|
v1.Use(maxBodySizeMiddleware(1 << 20)) // 1 MB limit on API request bodies
|
||||||
|
|
||||||
return router
|
return router
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// LivenessHandler returns a minimal liveness handler for dedicated health
|
||||||
|
// listeners (the REST router serves the same signal at /health).
|
||||||
|
func LivenessHandler() http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||||
|
"status": "healthy",
|
||||||
|
"service": "notifier",
|
||||||
|
"time": time.Now().UTC(),
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReadinessHandler exposes the readiness checks for dedicated health listeners.
|
||||||
|
func ReadinessHandler(checks map[string]ReadinessCheck) http.Handler {
|
||||||
|
return readinessHandler(checks)
|
||||||
|
}
|
||||||
|
|
||||||
|
// readinessHandler runs each dependency check and reports 503 if any fail.
|
||||||
|
func readinessHandler(checks map[string]ReadinessCheck) http.HandlerFunc {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ctx, cancel := context.WithTimeout(r.Context(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
status := http.StatusOK
|
||||||
|
components := make(map[string]string, len(checks))
|
||||||
|
for name, check := range checks {
|
||||||
|
if err := check(ctx); err != nil {
|
||||||
|
status = http.StatusServiceUnavailable
|
||||||
|
components[name] = "unavailable"
|
||||||
|
} else {
|
||||||
|
components[name] = "ok"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(status)
|
||||||
|
ready := status == http.StatusOK
|
||||||
|
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||||
|
"ready": ready,
|
||||||
|
"components": components,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// maxBodySizeMiddleware limits the size of incoming request bodies to prevent DoS.
|
// maxBodySizeMiddleware limits the size of incoming request bodies to prevent DoS.
|
||||||
func maxBodySizeMiddleware(maxBytes int64) func(http.Handler) http.Handler {
|
func maxBodySizeMiddleware(maxBytes int64) func(http.Handler) http.Handler {
|
||||||
return func(next http.Handler) http.Handler {
|
return func(next http.Handler) http.Handler {
|
||||||
@@ -162,10 +241,15 @@ func newCORSMiddleware(config *CORSConfig) func(http.Handler) http.Handler {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handle preflight OPTIONS requests
|
// Handle preflight OPTIONS requests: succeed only for allowed
|
||||||
if r.Method == http.MethodOptions {
|
// origins; disallowed cross-origin preflights get 403 with no
|
||||||
// Return 200 OK for preflight requests
|
// CORS headers so browsers block the actual request.
|
||||||
w.WriteHeader(http.StatusOK)
|
if r.Method == http.MethodOptions && origin != "" {
|
||||||
|
if allowed {
|
||||||
|
w.WriteHeader(http.StatusNoContent)
|
||||||
|
} else {
|
||||||
|
w.WriteHeader(http.StatusForbidden)
|
||||||
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,12 @@ server:
|
|||||||
rest_port: 8080
|
rest_port: 8080
|
||||||
host: "0.0.0.0"
|
host: "0.0.0.0"
|
||||||
mode: "both" # Options: both, grpc, rest
|
mode: "both" # Options: both, grpc, rest
|
||||||
|
# Optional TLS for both listeners. Leave disabled when a TLS-terminating
|
||||||
|
# gateway or service mesh fronts the service.
|
||||||
|
tls:
|
||||||
|
enabled: false
|
||||||
|
# cert_file: "/etc/notifier/tls/tls.crt"
|
||||||
|
# key_file: "/etc/notifier/tls/tls.key"
|
||||||
|
|
||||||
queue:
|
queue:
|
||||||
type: "local" # Options: local, kafka
|
type: "local" # Options: local, kafka
|
||||||
|
|||||||
@@ -31,6 +31,16 @@ type ServerConfig struct {
|
|||||||
RESTPort int `mapstructure:"rest_port"`
|
RESTPort int `mapstructure:"rest_port"`
|
||||||
Host string `mapstructure:"host"`
|
Host string `mapstructure:"host"`
|
||||||
Mode string `mapstructure:"mode"` // "both", "grpc", "rest"
|
Mode string `mapstructure:"mode"` // "both", "grpc", "rest"
|
||||||
|
TLS TLSConfig `mapstructure:"tls"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// TLSConfig enables TLS on the REST and gRPC listeners. When disabled the
|
||||||
|
// servers speak plaintext, which is only appropriate behind a TLS-terminating
|
||||||
|
// gateway or service mesh.
|
||||||
|
type TLSConfig struct {
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
CertFile string `mapstructure:"cert_file"`
|
||||||
|
KeyFile string `mapstructure:"key_file"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// NotifiersConfig contains configuration for all notifier types
|
// NotifiersConfig contains configuration for all notifier types
|
||||||
@@ -258,6 +268,12 @@ func (c *Config) Validate() error {
|
|||||||
return fmt.Errorf("invalid server mode: %s (must be both, grpc, or rest)", c.Server.Mode)
|
return fmt.Errorf("invalid server mode: %s (must be both, grpc, or rest)", c.Server.Mode)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if c.Server.TLS.Enabled {
|
||||||
|
if c.Server.TLS.CertFile == "" || c.Server.TLS.KeyFile == "" {
|
||||||
|
return fmt.Errorf("server.tls.enabled requires both cert_file and key_file")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Validate queue config
|
// Validate queue config
|
||||||
validQueueTypes := map[string]bool{"local": true, "kafka": true}
|
validQueueTypes := map[string]bool{"local": true, "kafka": true}
|
||||||
if !validQueueTypes[c.Queue.Type] {
|
if !validQueueTypes[c.Queue.Type] {
|
||||||
|
|||||||
Reference in New Issue
Block a user