| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323 |
- // HTTP POST handler for ingestd. Implements the seven protection
- // layers from SPEC §22 in order:
- //
- // 1. payload-size cap
- // 2. (per-IP) — M1: deferred to M5 with the WS path
- // 3. per-source token bucket
- // 4. per-company token bucket
- // 5. schema validate
- // 6. (broker CB) — M9
- // 7. (quarantine) — M9
- //
- // Auth: Stripe-style HMAC-SHA256 in header `X-BA-Signature: t=<ts>,v1=<hex>`.
- // The source's secret is keyed by (company_id, source_id) — for M0
- // the secret is fetched from the source table; for M0 the simplest
- // impl is a single env var per source, but we go straight to the
- // DB lookup so M2 doesn't have to change this path.
- package main
- import (
- "crypto/hmac"
- "crypto/sha256"
- "crypto/subtle"
- "encoding/hex"
- "encoding/json"
- "errors"
- "fmt"
- "io"
- "log/slog"
- "net/http"
- "os"
- "strconv"
- "strings"
- "time"
- "git3.techno-world.net/lrosales/broad-announce/internal/alert"
- "git3.techno-world.net/lrosales/broad-announce/internal/broker"
- "git3.techno-world.net/lrosales/broad-announce/internal/config"
- "git3.techno-world.net/lrosales/broad-announce/internal/dedupe"
- "git3.techno-world.net/lrosales/broad-announce/internal/observability"
- "git3.techno-world.net/lrosales/broad-announce/internal/ratelimit"
- "github.com/nats-io/nats.go"
- )
- // httpDeps is what the handler needs. Injected so tests can swap.
- type httpDeps struct {
- Logger *slog.Logger
- Metrics *observability.IngestdMetrics
- Limiter *ratelimit.Limiter
- Deduper *dedupe.Deduper
- JetStream natsPublisher
- MaxBytes int
- // For M0 we skip the DB lookup and read sources from a small
- // static map. M2 replaces this with a real store.
- Sources map[string]SourceConfig
- // Now is overridable in tests.
- Now func() time.Time
- }
- // SourceConfig is what we need to know about a source to authenticate
- // + rate-limit it. The full Sources row has more fields; this is the
- // hot-path subset.
- type SourceConfig struct {
- CompanyID string
- HMACSecret []byte
- RateLimitPerSec int
- AllowedTargets []string // M2
- }
- // natsPublisher is the minimal NATS interface the handler uses.
- type natsPublisher interface {
- PublishAsync(subj string, data []byte) error
- }
- // AcceptResponse is the JSON body returned on 202.
- type AcceptResponse struct {
- AlertID string `json:"alert_id"`
- DedupeCount uint32 `json:"dedupe_count"`
- ReceivedAt string `json:"received_at"`
- }
- // RegisterRoutes wires the ingestd HTTP routes onto the given mux.
- func RegisterRoutes(mux *http.ServeMux, d *httpDeps) {
- mux.HandleFunc("POST /v1/ingest", d.handleIngest)
- }
- // handleIngest is the M0 HTTP POST endpoint. Auth, validate, dedupe,
- // publish.
- func (d *httpDeps) handleIngest(w http.ResponseWriter, r *http.Request) {
- ctx := r.Context()
- now := d.now()
- // 1. Payload-size cap. We use MaxBytesReader so a streaming
- // client can't lie about Content-Length and try to OOM us.
- r.Body = http.MaxBytesReader(w, r.Body, int64(d.MaxBytes))
- body, err := io.ReadAll(r.Body)
- if err != nil {
- var maxErr *http.MaxBytesError
- if errors.As(err, &maxErr) {
- d.Metrics.AlertsReceived.WithLabelValues("payload_too_large").Inc()
- writeErr(w, http.StatusRequestEntityTooLarge, "payload_too_large", "")
- return
- }
- d.Metrics.AlertsReceived.WithLabelValues("invalid").Inc()
- writeErr(w, http.StatusBadRequest, "bad_request", err.Error())
- return
- }
- _ = r.Body.Close()
- // 5. Parse + validate. We treat any parse failure as invalid.
- var a alert.Alert
- if err := json.Unmarshal(body, &a); err != nil {
- d.Metrics.AlertsReceived.WithLabelValues("invalid").Inc()
- writeErr(w, http.StatusBadRequest, "invalid_json", err.Error())
- return
- }
- if err := a.Validate(); err != nil {
- d.Metrics.AlertsReceived.WithLabelValues("invalid").Inc()
- writeErr(w, http.StatusBadRequest, "invalid", err.Error())
- return
- }
- // Look up source. M0: in-memory map. M2: DB.
- src, ok := d.Sources[a.CompanyID+":"+a.SourceID]
- if !ok {
- d.Metrics.AlertsReceived.WithLabelValues("invalid").Inc()
- writeErr(w, http.StatusUnauthorized, "unknown_source",
- fmt.Sprintf("no such source %s/%s", a.CompanyID, a.SourceID))
- return
- }
- // Auth. Stripe-style: X-BA-Signature: t=<unix>,v1=<hex>
- if !verifyHMAC(r.Header.Get("X-BA-Signature"), src.HMACSecret, body, now) {
- d.Metrics.AlertsReceived.WithLabelValues("invalid").Inc()
- writeErr(w, http.StatusUnauthorized, "bad_signature", "")
- return
- }
- // 3. Per-source rate limit.
- if ok, ttl, err := d.Limiter.Allow(ctx, "source:"+a.CompanyID+":"+a.SourceID, src.RateLimitPerSec); err != nil {
- // Fail open on Redis errors — we don't want a Redis blip
- // to take down ingestion. Log loud, count it.
- d.Logger.Warn("ratelimit redis error (failing open)", "err", err, "scope", "source")
- } else if !ok {
- d.Metrics.AlertsReceived.WithLabelValues("rate_limited").Inc()
- d.Metrics.RateLimitHits.WithLabelValues("source").Inc()
- w.Header().Set("Retry-After", strconv.Itoa(int(ttl.Seconds())))
- writeErr(w, http.StatusTooManyRequests, "rate_limited_source", "")
- return
- }
- // 4. Per-company rate limit (cap from config; M2 will pull from DB).
- // For M0 we just use a constant default; replace with config load
- // once that lands.
- if ok, ttl, _ := d.Limiter.Allow(ctx, "company:"+a.CompanyID, 10_000); !ok {
- d.Metrics.AlertsReceived.WithLabelValues("rate_limited").Inc()
- d.Metrics.RateLimitHits.WithLabelValues("company").Inc()
- w.Header().Set("Retry-After", strconv.Itoa(int(ttl.Seconds())))
- writeErr(w, http.StatusTooManyRequests, "rate_limited_company", "")
- return
- }
- // 5b. Dedupe.
- isNew, count, err := d.Deduper.Check(ctx, a.SourceID, a.DedupeKey)
- if err != nil {
- // Fail open on dedupe errors too.
- d.Logger.Warn("dedupe redis error (failing open)", "err", err)
- isNew, count = true, 1
- }
- // Stamp server-side fields.
- a.ID = alert.NewID()
- a.ReceivedAt = now.UTC()
- a.DedupeCount = count
- // Publish to NATS.
- subject := broker.AlertsSubject(a.CompanyID)
- payload, err := json.Marshal(a)
- if err != nil {
- d.Metrics.AlertsReceived.WithLabelValues("invalid").Inc()
- writeErr(w, http.StatusInternalServerError, "marshal_failed", err.Error())
- return
- }
- start := time.Now()
- if err := d.JetStream.PublishAsync(subject, payload); err != nil {
- // Circuit breaker (M9) wraps this. For M0 we just fail loud.
- d.Metrics.AlertsReceived.WithLabelValues("circuit_open").Inc()
- writeErr(w, http.StatusServiceUnavailable, "broker_unavailable", err.Error())
- return
- }
- d.Metrics.PublishLatency.Observe(time.Since(start).Seconds())
- d.Metrics.PayloadBytes.Observe(float64(len(payload)))
- if isNew {
- d.Metrics.AlertsReceived.WithLabelValues("accepted").Inc()
- } else {
- d.Metrics.AlertsReceived.WithLabelValues("deduped").Inc()
- }
- w.Header().Set("Content-Type", "application/json")
- w.WriteHeader(http.StatusAccepted)
- _ = json.NewEncoder(w).Encode(AcceptResponse{
- AlertID: a.ID,
- DedupeCount: count,
- ReceivedAt: a.ReceivedAt.Format(time.RFC3339Nano),
- })
- d.Logger.Info("alert accepted",
- "alert_id", a.ID,
- "company_id", a.CompanyID,
- "source_id", a.SourceID,
- "severity", string(a.Severity),
- "dedupe_count", count,
- )
- }
- // now returns the testable clock.
- func (d *httpDeps) now() time.Time {
- if d.Now != nil {
- return d.Now()
- }
- return time.Now()
- }
- // verifyHMAC parses `X-BA-Signature: t=<unix>,v1=<hex>` and checks
- // HMAC-SHA256(secret, "<unix>.<body>") == hex. Replay window: 5 min.
- func verifyHMAC(header string, secret, body []byte, now time.Time) bool {
- if header == "" || len(secret) == 0 {
- return false
- }
- var tsStr, sigHex string
- for _, part := range strings.Split(header, ",") {
- kv := strings.SplitN(strings.TrimSpace(part), "=", 2)
- if len(kv) != 2 {
- continue
- }
- switch kv[0] {
- case "t":
- tsStr = kv[1]
- case "v1":
- sigHex = kv[1]
- }
- }
- if tsStr == "" || sigHex == "" {
- return false
- }
- tsInt, err := strconv.ParseInt(tsStr, 10, 64)
- if err != nil {
- return false
- }
- ts := time.Unix(tsInt, 0)
- if abs(now.Sub(ts)) > 5*time.Minute {
- return false
- }
- mac := hmac.New(sha256.New, secret)
- mac.Write([]byte(tsStr))
- mac.Write([]byte("."))
- mac.Write(body)
- expected := mac.Sum(nil)
- got, err := hex.DecodeString(sigHex)
- if err != nil {
- return false
- }
- return subtle.ConstantTimeCompare(expected, got) == 1
- }
- func abs(d time.Duration) time.Duration {
- if d < 0 {
- return -d
- }
- return d
- }
- func writeErr(w http.ResponseWriter, code int, reason, detail string) {
- w.Header().Set("Content-Type", "application/json")
- w.WriteHeader(code)
- body := map[string]string{"error": reason}
- if detail != "" {
- body["detail"] = detail
- }
- _ = json.NewEncoder(w).Encode(body)
- }
- // jsPublisher adapts a nats.JetStreamContext to the natsPublisher interface.
- type jsPublisher struct {
- js nats.JetStreamContext
- }
- func (j *jsPublisher) PublishAsync(subj string, data []byte) error {
- _, err := j.js.PublishAsync(subj, data)
- return err
- }
- // newNatsPublisher is the constructor used by main.
- func newNatsPublisher(js nats.JetStreamContext) natsPublisher {
- return &jsPublisher{js: js}
- }
- // loadSourcesFromEnv parses BA_INGESTD_SOURCES as a comma-separated
- // list of company_id:source_id:secret triples. M0 dev-mode only;
- // M2 swaps this for a Postgres lookup.
- func loadSourcesFromEnv(logger *slog.Logger) map[string]SourceConfig {
- out := map[string]SourceConfig{}
- raw := os.Getenv("BA_INGESTD_SOURCES")
- if raw == "" {
- logger.Warn("no BA_INGESTD_SOURCES configured; ingestd will reject all sources")
- return out
- }
- for _, triple := range strings.Split(raw, ",") {
- parts := strings.SplitN(triple, ":", 3)
- if len(parts) != 3 {
- logger.Warn("bad BA_INGESTD_SOURCES entry, expected company:source:secret", "entry", triple)
- continue
- }
- key := parts[0] + ":" + parts[1]
- out[key] = SourceConfig{
- CompanyID: parts[0],
- HMACSecret: []byte(parts[2]),
- RateLimitPerSec: config.GetInt("BA_INGESTD_RATE_LIMIT_PER_SOURCE", 100),
- }
- }
- logger.Info("loaded sources from env", "count", len(out))
- return out
- }
|