| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249 |
- // WebSocket ingest endpoint for ingestd (M5). Implements the
- // seven protection layers from SPEC §22 in the same order as
- // HTTP/MQTT, with layer 2 (per-IP concurrency cap) being new
- // in M5.
- //
- // Endpoint:
- //
- // GET /v1/ingest/ws HTTP Upgrade → WebSocket
- //
- // Protocol:
- //
- // client → server: {"api_key": "company:source:secret"} (auth frame)
- // server → client: {"ready": true} (auth ack)
- // client → server: {"alert": {...}, "auth": "t=...,v1=..."} (alert frame)
- // server → client: {"alert_id": "...", "dedupe_count": N} (per-alert ack)
- // server → client: {"error": "reason", "detail": "..."} (rejection)
- //
- // One WS connection = one (company_id, source_id) pair. After
- // the auth frame, the server only accepts alert frames for the
- // authenticated source. A second auth frame on the same
- // connection is a protocol error → close 1008.
- //
- // Why text frames: easier to debug with wscat / websocat. The
- // JSON shape is identical to the HTTP POST body. Binary frames
- // are rejected as a protocol error.
- package main
- import (
- "encoding/json"
- "net/http"
- "strings"
- "time"
- "git3.techno-world.net/lrosales/broad-announce/internal/concurrency"
- "github.com/gorilla/websocket"
- )
- // upgrader is shared by /v1/ingest/ws and /v1/tail/ws. We allow
- // any Origin in dev; the production cutover uses a same-origin
- // check.
- var upgrader = websocket.Upgrader{
- ReadBufferSize: 4 << 10, // 4 KB
- WriteBufferSize: 4 << 10,
- // Origins: enforce same-origin or empty Origin in prod.
- // M5 dev: allow any.
- CheckOrigin: func(r *http.Request) bool { return true },
- }
- // wsIngestDeps is the WS-handler-scoped wrapper. The body of
- // the handler is still a single ProcessAlert call; this struct
- // only adds the WS-specific bits (PerIP gate, max frame size,
- // per-conn read deadline).
- type wsIngestDeps struct {
- processDeps
- // PerIP is the shared per-IP concurrency cap. Incremented
- // on upgrade success; decremented on close.
- PerIP *concurrency.PerIP
- // MaxFrameBytes is the per-frame payload cap (SPEC §22
- // layer 1). Default 256 KB; matches HTTP's MaxBytes.
- MaxFrameBytes int64
- // ReadDeadline is the per-read deadline for the auth frame
- // and every subsequent alert frame. A slow client gets a
- // 1008 close.
- ReadDeadline time.Duration
- // WriteDeadline is the per-write deadline for the auth
- // ack and per-alert acks.
- WriteDeadline time.Duration
- }
- // authFrame is the first frame sent by the client. It is the
- // same shape as the HTTP `X-BA-Key` header but as a JSON object
- // so the server can also accept a per-conn token in the future.
- type wsAuthFrame struct {
- APIKey string `json:"api_key"`
- }
- // wsAckFrame is the per-alert ack.
- type wsAckFrame struct {
- AlertID string `json:"alert_id,omitempty"`
- DedupeCount uint32 `json:"dedupe_count,omitempty"`
- Error string `json:"error,omitempty"`
- Detail string `json:"detail,omitempty"`
- // Transport is the per-alert "result" label, useful for
- // tests that want to assert which metric advanced.
- Result string `json:"result,omitempty"`
- }
- // RegisterWSRoutes wires the WS endpoints onto the given mux.
- func RegisterWSRoutes(mux *http.ServeMux, d *wsIngestDeps, tail *wsTailDeps) {
- mux.HandleFunc("GET /v1/ingest/ws", d.handleIngest)
- if tail != nil {
- mux.HandleFunc("GET /v1/tail/ws", tail.handleTail)
- }
- }
- // handleIngest is the WS upgrade handler for /v1/ingest/ws.
- func (d *wsIngestDeps) handleIngest(w http.ResponseWriter, r *http.Request) {
- ip := clientIP(r)
- // Layer 2 — per-IP concurrency cap. Acquire BEFORE the
- // upgrade so a DoS'd client doesn't even reach the WS
- // handshake. We release on close (any reason).
- if !d.PerIP.Acquire(ip) {
- d.Metrics.ConnectionRejected.WithLabelValues("ws").Inc()
- d.Metrics.WSConnections.WithLabelValues("closed_per_ip_cap").Inc()
- http.Error(w, "per-IP connection cap exceeded", http.StatusTooManyRequests)
- return
- }
- conn, err := upgrader.Upgrade(w, r, nil)
- if err != nil {
- d.PerIP.Release(ip)
- d.Logger.Warn("ws upgrade", "err", err, "ip", ip)
- return
- }
- // From here on, every exit path must call
- // defer d.PerIP.Release(ip).
- defer d.PerIP.Release(ip)
- d.Metrics.WSConnections.WithLabelValues("open").Inc()
- // Auth frame. 5-second deadline: a client that dials and
- // doesn't send the auth frame fast enough is treated as
- // unauthenticated and gets a 1008 close.
- _ = conn.SetReadDeadline(time.Now().Add(5 * time.Second))
- _, msg, err := conn.ReadMessage()
- if err != nil {
- d.Metrics.WSConnections.WithLabelValues("closed_unauth").Inc()
- _ = conn.Close()
- return
- }
- var auth wsAuthFrame
- if err := json.Unmarshal(msg, &auth); err != nil || auth.APIKey == "" {
- d.Metrics.WSConnections.WithLabelValues("closed_unauth").Inc()
- _ = conn.WriteJSON(wsAckFrame{Error: "unauthorized", Detail: "bad auth frame"})
- _ = conn.Close()
- return
- }
- parts := strings.SplitN(auth.APIKey, ":", 3)
- if len(parts) != 3 {
- d.Metrics.WSConnections.WithLabelValues("closed_unauth").Inc()
- _ = conn.WriteJSON(wsAckFrame{Error: "unauthorized", Detail: "api_key must be company:source:secret"})
- _ = conn.Close()
- return
- }
- companyID, sourceID, secret := parts[0], parts[1], parts[2]
- src, ok := d.Sources[companyID+":"+sourceID]
- if !ok || string(src.HMACSecret) != secret {
- d.Metrics.WSConnections.WithLabelValues("closed_unauth").Inc()
- _ = conn.WriteJSON(wsAckFrame{Error: "unauthorized", Detail: "unknown source or wrong secret"})
- _ = conn.Close()
- return
- }
- // Replace the in-process processDeps.Sources with a single-
- // source map so a misrouted alert on this conn (e.g. one
- // with the wrong company_id) is rejected with unknown_source
- // — same as HTTP and MQTT.
- scoped := d.processDeps
- scoped.Sources = map[string]SourceConfig{companyID + ":" + sourceID: {
- CompanyID: companyID,
- HMACSecret: src.HMACSecret,
- RateLimitPerSec: src.RateLimitPerSec,
- AllowedTargets: src.AllowedTargets,
- }}
- // Auth ack
- _ = conn.SetWriteDeadline(time.Now().Add(d.WriteDeadline))
- if err := conn.WriteJSON(wsAckFrame{Result: "ready"}); err != nil {
- d.Metrics.WSConnections.WithLabelValues("closed_protocol_error").Inc()
- _ = conn.Close()
- return
- }
- // Loop. One frame per alert. We reset the read deadline on
- // every successful read; the conn-level deadline is
- // ReadDeadline from the deps.
- conn.SetReadLimit(d.MaxFrameBytes)
- closeState := "closed_clean"
- for {
- _ = conn.SetReadDeadline(time.Now().Add(d.ReadDeadline))
- mt, body, err := conn.ReadMessage()
- if err != nil {
- if websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) {
- closeState = "closed_clean"
- } else {
- closeState = "closed_protocol_error"
- }
- break
- }
- if mt != websocket.TextMessage {
- // Binary frames are a protocol error — M5 is text-only.
- _ = conn.WriteJSON(wsAckFrame{Error: "protocol_error", Detail: "binary frames not supported"})
- closeState = "closed_protocol_error"
- break
- }
- // Layer 1: MaxFrameBytes is enforced by SetReadLimit
- // above; gorilla returns an error if the frame exceeds
- // it. The error path closes the conn.
- // Sniff envelope (same as MQTT). Accept either
- // {alert, auth} envelope or bare alert body.
- alertBody, sigHeader := body, ""
- if len(body) > 0 && body[0] == '{' {
- var env struct {
- Alert json.RawMessage `json:"alert"`
- Auth string `json:"auth"`
- }
- if err := json.Unmarshal(body, &env); err == nil && len(env.Alert) > 0 {
- alertBody = env.Alert
- sigHeader = env.Auth
- }
- }
- d.Metrics.WSMessages.WithLabelValues("received").Inc()
- res := scoped.ProcessAlert(r.Context(), alertBody, sigHeader)
- _ = conn.SetWriteDeadline(time.Now().Add(d.WriteDeadline))
- if !res.Accepted {
- d.Metrics.WSMessages.WithLabelValues(res.RejectReason).Inc()
- _ = conn.WriteJSON(wsAckFrame{Error: res.RejectReason, Detail: res.Detail, Result: res.RejectReason})
- continue
- }
- d.Metrics.WSMessages.WithLabelValues("accepted").Inc()
- if !res.IsNew {
- d.Metrics.WSMessages.WithLabelValues("deduped").Inc()
- }
- _ = conn.WriteJSON(wsAckFrame{AlertID: res.AlertID, DedupeCount: res.DedupeCount, Result: "accepted"})
- _ = conn.SetWriteDeadline(time.Time{}) // reset
- }
- d.Metrics.WSConnections.WithLabelValues(closeState).Inc()
- _ = conn.Close()
- }
- // clientIP returns the best-effort source IP for r. We trust
- // the X-Forwarded-For header only when behind a known proxy —
- // in M5 dev, we use r.RemoteAddr. The production cutover reads
- // the proxy chain from config.
- func clientIP(r *http.Request) string {
- // gorilla's RemoteAddr is "host:port"; strip the port.
- addr := r.RemoteAddr
- if i := strings.LastIndex(addr, ":"); i > 0 {
- // Handle IPv6 "::1:80" — r.RemoteAddr can be "[::1]:80".
- if addr[0] == '[' {
- if j := strings.Index(addr, "]"); j > 0 {
- return addr[1:j]
- }
- }
- return addr[:i]
- }
- return addr
- }
|