diff --git a/internal/auth/token.go b/internal/auth/token.go index 07f2240..372c67c 100644 --- a/internal/auth/token.go +++ b/internal/auth/token.go @@ -2,12 +2,15 @@ package auth import ( "crypto/hmac" + "crypto/rand" "crypto/sha256" "encoding/base64" + "encoding/hex" "errors" "fmt" "strconv" "strings" + "sync" "time" ) @@ -15,20 +18,43 @@ var ( ErrTokenMalformed = errors.New("token malformed") ErrTokenExpired = errors.New("token expired") ErrTokenInvalid = errors.New("token invalid") + ErrTokenReplayed = errors.New("token already used") ) -// SignToken returns ":". +// Token format: "::" +// jti is a random 16-byte hex string (32 chars) embedded in the HMAC payload +// so it can't be swapped out without breaking the signature. ReplayCache +// records jti -> expiry; a second VerifyTokenAgainst call with the same +// jti returns ErrTokenReplayed. + +// SignToken signs a one-time-use token. Each call generates a fresh jti, so +// distinct calls with the same args still produce distinct tokens. func SignToken(secret, socketID, channel string, expiry time.Time) (string, error) { expMs := expiry.UnixMilli() - payload := fmt.Sprintf("%d|%s|%s", expMs, socketID, channel) + jtiBytes := make([]byte, 16) + if _, err := rand.Read(jtiBytes); err != nil { + return "", err + } + jti := hex.EncodeToString(jtiBytes) + payload := fmt.Sprintf("%d|%s|%s|%s", expMs, socketID, channel, jti) mac := hmac.New(sha256.New, []byte(secret)) mac.Write([]byte(payload)) - return strconv.FormatInt(expMs, 10) + ":" + base64.RawURLEncoding.EncodeToString(mac.Sum(nil)), nil + return strconv.FormatInt(expMs, 10) + ":" + jti + ":" + base64.RawURLEncoding.EncodeToString(mac.Sum(nil)), nil } +// VerifyToken validates a token without replay protection. Provided for +// backward-compatibility; production callers should construct a ReplayCache +// and call VerifyTokenAgainst. func VerifyToken(secret, socketID, channel, tok string) error { - parts := strings.SplitN(tok, ":", 2) - if len(parts) != 2 { + return VerifyTokenAgainst(secret, socketID, channel, tok, nil) +} + +// VerifyTokenAgainst is VerifyToken with optional replay protection. +// When cache is non-nil and the token verifies cleanly, the jti is recorded; +// a subsequent call with the same jti returns ErrTokenReplayed. +func VerifyTokenAgainst(secret, socketID, channel, tok string, cache *ReplayCache) error { + parts := strings.SplitN(tok, ":", 3) + if len(parts) != 3 { return ErrTokenMalformed } expMs, err := strconv.ParseInt(parts[0], 10, 64) @@ -38,14 +64,69 @@ func VerifyToken(secret, socketID, channel, tok string) error { if time.Now().UnixMilli() > expMs { return ErrTokenExpired } - sig, err := base64.RawURLEncoding.DecodeString(parts[1]) + jti := parts[1] + if jti == "" { + return ErrTokenMalformed + } + sig, err := base64.RawURLEncoding.DecodeString(parts[2]) if err != nil { return ErrTokenMalformed } mac := hmac.New(sha256.New, []byte(secret)) - _, _ = fmt.Fprintf(mac, "%d|%s|%s", expMs, socketID, channel) + _, _ = fmt.Fprintf(mac, "%d|%s|%s|%s", expMs, socketID, channel, jti) if !hmac.Equal(sig, mac.Sum(nil)) { return ErrTokenInvalid } + if cache != nil { + if !cache.CheckAndRecord(jti, time.UnixMilli(expMs)) { + return ErrTokenReplayed + } + } return nil } + +// ReplayCache records token jti -> expiry. Safe for concurrent use. Memory +// stays bounded by Sweep, which removes expired entries; callers that don't +// run Sweep periodically will accumulate memory at the rate of issued +// tokens until the next Sweep. +type ReplayCache struct { + mu sync.Mutex + seen map[string]time.Time +} + +func NewReplayCache() *ReplayCache { + return &ReplayCache{seen: map[string]time.Time{}} +} + +// CheckAndRecord returns true the first time it sees jti, false thereafter. +func (c *ReplayCache) CheckAndRecord(jti string, exp time.Time) bool { + c.mu.Lock() + defer c.mu.Unlock() + if _, ok := c.seen[jti]; ok { + return false + } + c.seen[jti] = exp + return true +} + +// Sweep removes entries whose expiry has passed. Returns the count removed. +func (c *ReplayCache) Sweep() int { + now := time.Now() + c.mu.Lock() + defer c.mu.Unlock() + swept := 0 + for jti, exp := range c.seen { + if !now.Before(exp) { + delete(c.seen, jti) + swept++ + } + } + return swept +} + +// Len reports the current number of cached entries (for tests / metrics). +func (c *ReplayCache) Len() int { + c.mu.Lock() + defer c.mu.Unlock() + return len(c.seen) +} diff --git a/internal/auth/token_test.go b/internal/auth/token_test.go index b733022..f7510dc 100644 --- a/internal/auth/token_test.go +++ b/internal/auth/token_test.go @@ -41,3 +41,47 @@ func TestVerifyTokenTampered(t *testing.T) { t.Fatal("expected tamper error") } } + +func TestVerifyTokenAgainstCachePreventsReplay(t *testing.T) { + secret := "s" + tok, _ := SignToken(secret, "sock1", "private-x", time.Now().Add(time.Minute)) + cache := NewReplayCache() + if err := VerifyTokenAgainst(secret, "sock1", "private-x", tok, cache); err != nil { + t.Fatalf("first verify: %v", err) + } + if err := VerifyTokenAgainst(secret, "sock1", "private-x", tok, cache); err == nil { + t.Fatal("expected ErrTokenReplayed on second use, got nil") + } +} + +func TestVerifyTokenAgainstNilCache(t *testing.T) { + secret := "s" + tok, _ := SignToken(secret, "sock1", "private-x", time.Now().Add(time.Minute)) + if err := VerifyTokenAgainst(secret, "sock1", "private-x", tok, nil); err != nil { + t.Fatalf("first verify with nil cache: %v", err) + } + if err := VerifyTokenAgainst(secret, "sock1", "private-x", tok, nil); err != nil { + t.Fatalf("second verify with nil cache: %v", err) + } +} + +func TestReplayCacheSweepRemovesExpired(t *testing.T) { + c := NewReplayCache() + c.CheckAndRecord("expired", time.Now().Add(-time.Minute)) + c.CheckAndRecord("alive", time.Now().Add(time.Minute)) + if got := c.Sweep(); got != 1 { + t.Errorf("Sweep removed %d, want 1", got) + } + if got, want := c.Len(), 1; got != want { + t.Errorf("post-Sweep Len = %d, want %d", got, want) + } +} + +func TestSignTokenJtiUnique(t *testing.T) { + secret := "s" + a, _ := SignToken(secret, "sock1", "x", time.Now().Add(time.Minute)) + b, _ := SignToken(secret, "sock1", "x", time.Now().Add(time.Minute)) + if a == b { + t.Fatal("two SignToken calls with identical args produced identical tokens — jti collision or absent") + } +} diff --git a/internal/conn/conn.go b/internal/conn/conn.go index 5574e46..3d148e2 100644 --- a/internal/conn/conn.go +++ b/internal/conn/conn.go @@ -9,6 +9,7 @@ import ( "sync/atomic" "time" + "github.com/EthanY33/wirefan/internal/auth" "github.com/EthanY33/wirefan/internal/fanout" "github.com/EthanY33/wirefan/internal/hub" "github.com/EthanY33/wirefan/internal/metrics" @@ -37,6 +38,7 @@ type Conn struct { send chan []byte registry registry.Registry signingSecret string + replayCache *auth.ReplayCache fanout fanout.Fanout rateLimit *ratelimit.Limiter // per-API-key bucket; shared across all conns owned by the key connRate *rate.Limiter // per-conn bucket; bounds a single socket's throughput @@ -70,7 +72,11 @@ func (c *Conn) CloseFrame(code websocket.StatusCode, reason string) { } // Run owns the conn for its lifetime. Returns when ctx is canceled or peer disconnects. -func Run(ctx context.Context, ws *websocket.Conn, socketID, apiKeyID string, reg registry.Registry, signingSecret string, fan fanout.Fanout, rl *ratelimit.Limiter, pol Policy, h *hub.Hub) error { +// +// replayCache may be nil; when nil, subscribe-token replay protection is +// disabled (used by tests). Production callers pass a process-wide cache so +// a leaked subscribe token cannot be reused within its 5-minute window. +func Run(ctx context.Context, ws *websocket.Conn, socketID, apiKeyID string, reg registry.Registry, signingSecret string, replayCache *auth.ReplayCache, fan fanout.Fanout, rl *ratelimit.Limiter, pol Policy, h *hub.Hub) error { c := &Conn{ ws: ws, socketID: socketID, @@ -78,6 +84,7 @@ func Run(ctx context.Context, ws *websocket.Conn, socketID, apiKeyID string, reg send: make(chan []byte, sendChanSize), registry: reg, signingSecret: signingSecret, + replayCache: replayCache, fanout: fan, rateLimit: rl, connRate: rate.NewLimiter(rate.Limit(defaultConnPublishRate), defaultConnPublishBurst), diff --git a/internal/conn/conn_test.go b/internal/conn/conn_test.go index 0e7129d..90f2955 100644 --- a/internal/conn/conn_test.go +++ b/internal/conn/conn_test.go @@ -27,7 +27,7 @@ func TestConnectedMessageSent(t *testing.T) { handler := func(c *websocket.Conn) { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() - _ = Run(ctx, c, "01HTEST", "test-key", registry.NewSyncMap(), "test-signing-secret", fanout.NewPerConn(), rl, PolicyDisconnect{}, hub.New()) + _ = Run(ctx, c, "01HTEST", "test-key", registry.NewSyncMap(), "test-signing-secret", nil, fanout.NewPerConn(), rl, PolicyDisconnect{}, hub.New()) } srv := httptest.NewServer(websocketHandler(handler)) diff --git a/internal/conn/handler.go b/internal/conn/handler.go index 84ce226..4e8d51d 100644 --- a/internal/conn/handler.go +++ b/internal/conn/handler.go @@ -81,8 +81,12 @@ func (c *Conn) handleSubscribe(msg incoming) { return } if strings.HasPrefix(msg.Channel, "private-") { - if err := auth.VerifyToken(c.signingSecret, c.socketID, msg.Channel, msg.Token); err != nil { + if err := auth.VerifyTokenAgainst(c.signingSecret, c.socketID, msg.Channel, msg.Token, c.replayCache); err != nil { metrics.AuthFails.Inc() + if errors.Is(err, auth.ErrTokenReplayed) { + c.sendError("AUTH_REPLAYED", "token already used") + return + } c.sendError("AUTH_FAILED", "invalid token") return } diff --git a/internal/conn/handler_test.go b/internal/conn/handler_test.go index 76c716c..c554cc5 100644 --- a/internal/conn/handler_test.go +++ b/internal/conn/handler_test.go @@ -31,7 +31,7 @@ func newTestConn(t *testing.T, signingSecret string) (*websocket.Conn, string) { } ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - _ = Run(ctx, c, socketID, "test-key", registry.NewSyncMap(), signingSecret, fanout.NewPerConn(), rl, PolicyDisconnect{}, hub.New()) + _ = Run(ctx, c, socketID, "test-key", registry.NewSyncMap(), signingSecret, nil, fanout.NewPerConn(), rl, PolicyDisconnect{}, hub.New()) }) srv := httptest.NewServer(handler) t.Cleanup(srv.Close) diff --git a/internal/server/leak_test.go b/internal/server/leak_test.go index 86b2e90..dab1894 100644 --- a/internal/server/leak_test.go +++ b/internal/server/leak_test.go @@ -33,6 +33,7 @@ func TestNoGoroutineLeakAfterChurn(t *testing.T) { []string{"*"}, registry.NewSyncMap(), "test-signing-secret", + nil, fanout.NewPerConn(), rl, conn.PolicyDisconnect{}, diff --git a/internal/server/server.go b/internal/server/server.go index eab0b57..fca002f 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -8,6 +8,7 @@ import ( "net/http/pprof" "time" + "github.com/EthanY33/wirefan/internal/auth" "github.com/EthanY33/wirefan/internal/conn" "github.com/EthanY33/wirefan/internal/fanout" "github.com/EthanY33/wirefan/internal/hub" @@ -33,26 +34,29 @@ type Config struct { } type Server struct { - cfg Config - health *HealthHandler - mux *http.ServeMux - adminMux *http.ServeMux - srv *http.Server - adminSrv *http.Server - store store.Store - hub *hub.Hub + cfg Config + health *HealthHandler + mux *http.ServeMux + adminMux *http.ServeMux + srv *http.Server + adminSrv *http.Server + store store.Store + hub *hub.Hub + replayCache *auth.ReplayCache } // New builds the public and admin muxes. The admin listener is created // only when cfg.AdminAddr is non-empty. func New(cfg Config, st store.Store, adminToken string, reg registry.Registry, signingSecret string, fan fanout.Fanout, rl *ratelimit.Limiter, pol conn.Policy, h *hub.Hub) *Server { + rc := auth.NewReplayCache() s := &Server{ - cfg: cfg, - health: NewHealthHandler(), - mux: http.NewServeMux(), - adminMux: http.NewServeMux(), - store: st, - hub: h, + cfg: cfg, + health: NewHealthHandler(), + mux: http.NewServeMux(), + adminMux: http.NewServeMux(), + store: st, + hub: h, + replayCache: rc, } rest := NewRestHandler(st, adminToken, signingSecret) @@ -60,7 +64,7 @@ func New(cfg Config, st store.Store, adminToken string, reg registry.Registry, s // Public listener: health, /v1/connect (WS), /v1/auth/sign, static client. s.mux.Handle("/v1/health", s.health) rest.RegisterPublic(s.mux) - s.mux.Handle("/v1/connect", NewUpgradeHandler(st, cfg.AllowedOrigins, reg, signingSecret, fan, rl, pol, h)) + s.mux.Handle("/v1/connect", NewUpgradeHandler(st, cfg.AllowedOrigins, reg, signingSecret, rc, fan, rl, pol, h)) s.mux.Handle("/", http.FileServerFS(web.Files)) // Admin listener: metrics, pprof, key management. All gated by @@ -81,6 +85,11 @@ func New(cfg Config, st store.Store, adminToken string, reg registry.Registry, s return s } +// ReplayCache exposes the per-Server token replay cache so the caller can +// run a periodic Sweep goroutine. Public so cmd/wirefan/main.go can drive +// the sweeper without exporting a separate accessor. +func (s *Server) ReplayCache() *auth.ReplayCache { return s.replayCache } + func (s *Server) Run(ctx context.Context) error { errc := make(chan error, 2) go func() { @@ -97,6 +106,7 @@ func (s *Server) Run(ctx context.Context) error { } }() } + go s.sweepReplayCache(ctx) select { case err := <-errc: @@ -113,3 +123,20 @@ func (s *Server) Run(ctx context.Context) error { } return s.srv.Shutdown(shutdownCtx) } + +// sweepReplayCache evicts expired token jti entries every minute. Memory in +// the cache is bounded by the issuance rate * token lifetime (5 minutes by +// default), so a sweep cadence of one minute gives at most ~5 minutes of +// expired entries before reclamation. Loops until ctx is canceled. +func (s *Server) sweepReplayCache(ctx context.Context) { + t := time.NewTicker(time.Minute) + defer t.Stop() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + s.replayCache.Sweep() + } + } +} diff --git a/internal/server/shutdown_test.go b/internal/server/shutdown_test.go index d21117b..22b7093 100644 --- a/internal/server/shutdown_test.go +++ b/internal/server/shutdown_test.go @@ -28,7 +28,7 @@ func TestDrainClosesAllConnections(t *testing.T) { upgrader := NewUpgradeHandler( s, []string{"*"}, registry.NewSyncMap(), "test-secret", - fanout.NewPerConn(), rl, conn.PolicyDisconnect{}, h, + nil, fanout.NewPerConn(), rl, conn.PolicyDisconnect{}, h, ) srv := httptest.NewServer(upgrader) defer srv.Close() diff --git a/internal/server/upgrade.go b/internal/server/upgrade.go index 183ae67..aab3260 100644 --- a/internal/server/upgrade.go +++ b/internal/server/upgrade.go @@ -9,6 +9,7 @@ import ( "strings" "sync" + "github.com/EthanY33/wirefan/internal/auth" "github.com/EthanY33/wirefan/internal/conn" "github.com/EthanY33/wirefan/internal/fanout" "github.com/EthanY33/wirefan/internal/hub" @@ -31,6 +32,7 @@ type UpgradeHandler struct { allowedOrigins []string registry registry.Registry signingSecret string + replayCache *auth.ReplayCache fanout fanout.Fanout rateLimit *ratelimit.Limiter policy conn.Policy @@ -42,12 +44,13 @@ type UpgradeHandler struct { ipCap int } -func NewUpgradeHandler(st store.Store, origins []string, reg registry.Registry, signingSecret string, fan fanout.Fanout, rl *ratelimit.Limiter, pol conn.Policy, h *hub.Hub) *UpgradeHandler { +func NewUpgradeHandler(st store.Store, origins []string, reg registry.Registry, signingSecret string, replayCache *auth.ReplayCache, fan fanout.Fanout, rl *ratelimit.Limiter, pol conn.Policy, h *hub.Hub) *UpgradeHandler { return &UpgradeHandler{ store: st, allowedOrigins: origins, registry: reg, signingSecret: signingSecret, + replayCache: replayCache, fanout: fan, rateLimit: rl, policy: pol, @@ -240,5 +243,5 @@ func (h *UpgradeHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { return } sid := ulid.Make().String() - _ = conn.Run(r.Context(), c, sid, k.ID, h.registry, h.signingSecret, h.fanout, h.rateLimit, h.policy, h.hub) + _ = conn.Run(r.Context(), c, sid, k.ID, h.registry, h.signingSecret, h.replayCache, h.fanout, h.rateLimit, h.policy, h.hub) } diff --git a/internal/server/upgrade_test.go b/internal/server/upgrade_test.go index 4b519ff..ec7924f 100644 --- a/internal/server/upgrade_test.go +++ b/internal/server/upgrade_test.go @@ -38,7 +38,7 @@ func TestUpgradeSucceeds(t *testing.T) { k, _ := s.CreateKey(context.Background(), "t", auth.HashSecret(secret)) rl := ratelimit.New(100, 200, time.Hour) t.Cleanup(rl.Close) - h := NewUpgradeHandler(s, []string{"*"}, registry.NewSyncMap(), "test-signing-secret", fanout.NewPerConn(), rl, conn.PolicyDisconnect{}, hub.New()) + h := NewUpgradeHandler(s, []string{"*"}, registry.NewSyncMap(), "test-signing-secret", nil, fanout.NewPerConn(), rl, conn.PolicyDisconnect{}, hub.New()) srv := httptest.NewServer(h) defer srv.Close() wsURL := strings.Replace(srv.URL, "http", "ws", 1) + "/v1/connect?key=" + k.ID @@ -52,7 +52,7 @@ func TestUpgradeSucceeds(t *testing.T) { func newTestUpgrader(t *testing.T) http.Handler { rl := ratelimit.New(100, 200, time.Hour) t.Cleanup(rl.Close) - return NewUpgradeHandler(store.NewMemory(), []string{"*"}, registry.NewSyncMap(), "test-signing-secret", fanout.NewPerConn(), rl, conn.PolicyDisconnect{}, hub.New()) + return NewUpgradeHandler(store.NewMemory(), []string{"*"}, registry.NewSyncMap(), "test-signing-secret", nil, fanout.NewPerConn(), rl, conn.PolicyDisconnect{}, hub.New()) } func TestParseTrustedProxies(t *testing.T) {