Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 41 additions & 0 deletions livestream/auth/access.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,41 @@
package auth

import (
"context"
"net/http"
"time"

"github.com/labstack/echo/v4"
"github.com/spf13/viper"
)

var authorizationClient = &http.Client{
Timeout: 3 * time.Second,
CheckRedirect: func(req *http.Request, via []*http.Request) error {
return http.ErrUseLastResponse
},
}

func CheckAccess(ctx context.Context, header http.Header) error {
url := viper.GetString("jwt.authorization_url")
if url == "" {
return nil
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return echo.NewHTTPError(http.StatusServiceUnavailable, "live stream authorization unavailable")
}
request.Header.Set("Authorization", header.Get("Authorization"))
response, err := authorizationClient.Do(request)
if err != nil {
return echo.NewHTTPError(http.StatusServiceUnavailable, "live stream authorization unavailable")
}
defer func() { _ = response.Body.Close() }()
if response.StatusCode == http.StatusNoContent {
return nil
}
if response.StatusCode == http.StatusUnauthorized || response.StatusCode == http.StatusForbidden {
return echo.NewHTTPError(http.StatusUnauthorized, "live stream access denied")
}
return echo.NewHTTPError(http.StatusServiceUnavailable, "live stream authorization unavailable")
}
65 changes: 65 additions & 0 deletions livestream/auth/access_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
package auth

import (
"context"
"net/http"
"net/http/httptest"
"testing"

"github.com/labstack/echo/v4"
"github.com/spf13/viper"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestCheckAccess(t *testing.T) {
for _, test := range []struct {
status int
want int
}{
{http.StatusNoContent, 0},
{http.StatusOK, http.StatusServiceUnavailable},
{http.StatusFound, http.StatusServiceUnavailable},
{http.StatusUnauthorized, http.StatusUnauthorized},
{http.StatusForbidden, http.StatusUnauthorized},
{http.StatusInternalServerError, http.StatusServiceUnavailable},
} {
t.Run(http.StatusText(test.status), func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, "Bearer test-live-stream-token", r.Header.Get("Authorization"))
if r.URL.Path == "/redirected" {
w.WriteHeader(http.StatusNoContent)
return
}
w.Header().Set("Location", "/redirected")
w.WriteHeader(test.status)
}))
defer server.Close()
viper.Set("jwt.authorization_url", server.URL)
t.Cleanup(func() { viper.Set("jwt.authorization_url", "") })
err := CheckAccess(context.Background(), http.Header{"Authorization": {"Bearer test-live-stream-token"}})
if test.want == 0 {
require.NoError(t, err)
} else {
var httpError *echo.HTTPError
require.ErrorAs(t, err, &httpError)
assert.Equal(t, test.want, httpError.Code)
}
})
}
}

func TestCheckAccessFailsClosedOnConnectionFailure(t *testing.T) {
server := httptest.NewServer(http.NotFoundHandler())
server.Close()
viper.Set("jwt.authorization_url", server.URL)
t.Cleanup(func() { viper.Set("jwt.authorization_url", "") })
var httpError *echo.HTTPError
require.ErrorAs(t, CheckAccess(context.Background(), http.Header{}), &httpError)
assert.Equal(t, http.StatusServiceUnavailable, httpError.Code)
}

func TestCheckAccessBeforeRollout(t *testing.T) {
viper.Set("jwt.authorization_url", "")
require.NoError(t, CheckAccess(context.Background(), http.Header{}))
}
18 changes: 10 additions & 8 deletions livestream/configs/configs.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,8 @@ type MMDBConfig struct {
}

type JWTConfig struct {
Secret string
Secret string
AuthorizationURL string `mapstructure:"authorization_url"`
// Previous secrets still accepted for verification (never used for signing),
// so tokens signed before a key rotation keep working until they expire.
SecretFallbacks []string `mapstructure:"secret_fallbacks"`
Expand All @@ -24,13 +25,13 @@ type SessionRecordingConfig struct {
}

type RedisConfig struct {
Address string `mapstructure:"address"`
Port string `mapstructure:"port"`
TLS bool `mapstructure:"tls"`
FlushIntervalMs int `mapstructure:"flush_interval_ms"`
UsePubSub bool `mapstructure:"use_pub_sub"`
PublishBufferSize int `mapstructure:"publish_buffer_size"`
PublishWorkers int `mapstructure:"publish_workers"`
Address string `mapstructure:"address"`
Port string `mapstructure:"port"`
TLS bool `mapstructure:"tls"`
FlushIntervalMs int `mapstructure:"flush_interval_ms"`
UsePubSub bool `mapstructure:"use_pub_sub"`
PublishBufferSize int `mapstructure:"publish_buffer_size"`
PublishWorkers int `mapstructure:"publish_workers"`
}

// ConsumerConfig holds connection and tuning parameters for a single Kafka consumer.
Expand Down Expand Up @@ -130,6 +131,7 @@ func InitConfigs(filename, configPath string) {
// JWT settings
_ = viper.BindEnv("jwt.secret") // LIVESTREAM_JWT_SECRET
_ = viper.BindEnv("jwt.secret_fallbacks") // LIVESTREAM_JWT_SECRET_FALLBACKS (comma-separated)
_ = viper.BindEnv("jwt.authorization_url")

// Session recording settings
_ = viper.BindEnv("session_recording.max_lru_entries") // LIVESTREAM_SESSION_RECORDING_MAX_LRU_ENTRIES
Expand Down
46 changes: 45 additions & 1 deletion livestream/handlers/handlers.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package handlers

import (
"context"
"encoding/json"
"fmt"
"log"
Expand Down Expand Up @@ -50,6 +51,9 @@ func StatsHandler(stats *events.Stats, sessionStats *events.SessionStats, redisS
if err != nil {
return c.JSON(http.StatusUnauthorized, resp{Error: "wrong token claims"})
}
if err := auth.CheckAccess(c.Request().Context(), c.Request().Header); err != nil {

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Throttle authorization on the 1.5-second stats path

When the Live Events page is open, it polls /stats every 1,500 ms (LiveEventsTable.tsx:32,74), so this new call generates about 40 synchronous Django authorization requests per minute per viewer. Each request performs User and Team lookups and, for access-control organizations, UserPermissions loads every project access-control row across all organizations the user belongs to (posthog/user_permissions.py:113-134), despite checking one team. This multiplies organization-wide database scans by active viewers and can make the three-second fail-closed timeout disconnect users under load; cache or throttle the stats authorization decision to the 15–30-second recheck cadence, or use a team-scoped resolver.

Useful? React with 👍 / 👎.

return err
}

if redisStore != nil {
ctx := c.Request().Context()
Expand Down Expand Up @@ -92,6 +96,29 @@ func StatsHandler(stats *events.Stats, sessionStats *events.SessionStats, redisS

var subID uint64 = 1

func periodicAccessChecks(ctx context.Context, header http.Header, interval time.Duration) <-chan error {
errors := make(chan error, 1)
go func() {
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if err := auth.CheckAccess(ctx, header); err != nil {
select {
case errors <- err:
case <-ctx.Done():
}
return
}
}
}
}()
return errors
}

func StreamEventsHandler(log echo.Logger, subChan chan events.Subscription, unSubChan chan events.Subscription) func(c echo.Context) error {
return func(c echo.Context) error {
log.Debugf("SSE client connected, ip: %v", c.RealIP())
Expand All @@ -107,6 +134,9 @@ func StreamEventsHandler(log echo.Logger, subChan chan events.Subscription, unSu
if err != nil || token == "" || teamID == 0 {
return echo.NewHTTPError(http.StatusUnauthorized, "wrong token")
}
if err := auth.CheckAccess(c.Request().Context(), c.Request().Header); err != nil {
return err
}

eventType := c.QueryParam("eventType")
distinctId := c.QueryParam("distinctId")
Expand Down Expand Up @@ -163,8 +193,14 @@ func StreamEventsHandler(log echo.Logger, subChan chan events.Subscription, unSu
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
timeout := time.After(30 * time.Minute)
accessContext, cancelAccessCheck := context.WithCancel(c.Request().Context())
defer cancelAccessCheck()
accessErrors := periodicAccessChecks(accessContext, c.Request().Header.Clone(), 30*time.Second)
for {
select {
case err := <-accessErrors:
log.Warnf("Live stream authorization check failed: %v", err)
return nil
case <-timeout:
log.Debug("SSE connection to be terminated after timeout")
return nil
Expand Down Expand Up @@ -275,11 +311,15 @@ func NotificationsHandler(redisClient rueidis.Client) func(c echo.Context) error
// Old tokens without organization_id/user_id — no-op until all tokens refresh
return c.NoContent(http.StatusNoContent)
}
if err := auth.CheckAccess(c.Request().Context(), c.Request().Header); err != nil {
return err
}
ctx, cancel := context.WithCancel(c.Request().Context())
defer cancel()

metrics.NotificationSubs.Inc()
defer metrics.NotificationSubs.Dec()

ctx := c.Request().Context()
channel := fmt.Sprintf("notifications:%s", claims.OrganizationID)

// Absorbs publish-rate bursts; drops on overflow to avoid blocking rueidis.
Expand Down Expand Up @@ -307,6 +347,7 @@ func NotificationsHandler(redisClient rueidis.Client) func(c echo.Context) error
heartbeat := time.NewTicker(15 * time.Second)
defer heartbeat.Stop()
timeout := time.After(30 * time.Minute)
accessErrors := periodicAccessChecks(ctx, c.Request().Header.Clone(), 15*time.Second)

for {
select {
Expand All @@ -319,6 +360,9 @@ func NotificationsHandler(redisClient rueidis.Client) func(c echo.Context) error
log.Printf("Redis subscription error: %v", err)
}
return nil
case err := <-accessErrors:
Comment thread
pauldambra marked this conversation as resolved.
log.Printf("Live stream authorization check failed: %v", err)
return nil
case msg := <-msgCh:
cleaned, ok, reason := filterNotificationForUser(msg, claims.UserID)
if !ok {
Expand Down
Loading
Loading