diff --git a/go.mod b/go.mod index 029c145..6811b87 100644 --- a/go.mod +++ b/go.mod @@ -2,4 +2,7 @@ module l4.pm/jetstream-proxy go 1.25.1 -require github.com/bluesky-social/jetstream v0.0.0-20251009222037-7d7efa58d7f1 // indirect +require ( + github.com/bluesky-social/jetstream v0.0.0-20251009222037-7d7efa58d7f1 // indirect + github.com/gorilla/websocket v1.5.3 // indirect +) diff --git a/go.sum b/go.sum index 9919d25..87f2eae 100644 --- a/go.sum +++ b/go.sum @@ -1,2 +1,4 @@ github.com/bluesky-social/jetstream v0.0.0-20251009222037-7d7efa58d7f1 h1:ovcRKN1iXZnY5WApVg+0Hw2RkwMH0ziA7lSAA8vellU= github.com/bluesky-social/jetstream v0.0.0-20251009222037-7d7efa58d7f1/go.mod h1:5PtGi4r/PjEVBBl+0xWuQn4mBEjr9h6xsfDBADS6cHs= +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= diff --git a/main.go b/main.go index a59f5fa..ae492e2 100644 --- a/main.go +++ b/main.go @@ -8,6 +8,8 @@ import ( "strings" "sync" "time" + + "github.com/gorilla/websocket" ) var DEFAULT_POOL = []string{ @@ -20,6 +22,52 @@ var DEFAULT_POOL = []string{ // want yours here? contact me } +// Broadcaster manages subscribers to Jetstream events +type Broadcaster struct { + listeners []chan []byte + mu sync.Mutex +} + +// Subscribe returns a new channel that will receive Jetstream events +func (b *Broadcaster) Subscribe() chan []byte { + b.mu.Lock() + defer b.mu.Unlock() + + // firehose can be more-than-1k events per second, + // prefer to create a large buffer for the subscribers + ch := make(chan []byte, 10000) + b.listeners = append(b.listeners, ch) + return ch +} + +func (b *Broadcaster) Unsubscribe(ch chan []byte) { + b.mu.Lock() + defer b.mu.Unlock() + + for i, listener := range b.listeners { + if listener == ch { + b.listeners = append(b.listeners[:i], b.listeners[i+1:]...) + close(ch) + break + } + } +} + +func (b *Broadcaster) Broadcast(message []byte) { + b.mu.Lock() + defer b.mu.Unlock() + + for _, ch := range b.listeners { + select { + case ch <- message: + // event sent successfully. we don't want to block + default: + // channel full, skip to avoid blocking + slog.Warn("jetstream broadcast: channel full, dropping event") + } + } +} + type latencyResult struct { url string latency time.Duration @@ -29,8 +77,12 @@ type latencyResult struct { func measureLatency(url string) (time.Duration, error) { httpsURL := strings.Replace(url, "wss://", "https://", 1) + client := &http.Client{ + Timeout: 20 * time.Second, + } + start := time.Now() - resp, err := http.Get(httpsURL) + resp, err := client.Get(httpsURL) if err != nil { return 0, err } @@ -39,6 +91,106 @@ func measureLatency(url string) (time.Duration, error) { return time.Since(start), nil } +var upgrader = websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { + return true // Allow all origins + }, +} + +// handleSubscribe upgrades HTTP connection to websocket and streams events +func handleSubscribe(broadcaster *Broadcaster) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + slog.Error("Failed to upgrade connection", slog.Any("error", err)) + return + } + defer conn.Close() + + // Subscribe to broadcaster + ch := broadcaster.Subscribe() + defer broadcaster.Unsubscribe(ch) + + slog.Info("Client connected", slog.String("remote", r.RemoteAddr)) + + // Stream events to client + for message := range ch { + err := conn.WriteMessage(websocket.TextMessage, message) + if err != nil { + slog.Debug("Client disconnected", slog.String("remote", r.RemoteAddr), slog.Any("error", err)) + break + } + } + + slog.Info("Client disconnected", slog.String("remote", r.RemoteAddr)) + } +} + +// connectToUpstream maintains a connection to the upstream websocket and broadcasts messages +func connectToUpstream(pool []string, broadcaster *Broadcaster) { + backoff := time.Second + maxBackoff := time.Minute + var currentUpstream string + + for { + // Find best upstream (re-evaluate on each connection attempt) + bestUpstream, err := findBestUpstream(pool) + if err != nil { + slog.Error("Failed to find best upstream", slog.Any("error", err)) + time.Sleep(backoff) + backoff *= 2 + if backoff > maxBackoff { + backoff = maxBackoff + } + continue + } + + if bestUpstream != currentUpstream { + slog.Info("Switching to new upstream", slog.String("url", bestUpstream)) + currentUpstream = bestUpstream + } + + slog.Info("Connecting to upstream", slog.String("url", currentUpstream)) + + conn, _, err := websocket.DefaultDialer.Dial(currentUpstream+"/subscribe", nil) + if err != nil { + slog.Error("Failed to connect to upstream", slog.String("url", currentUpstream), slog.Any("error", err)) + time.Sleep(backoff) + backoff *= 2 + if backoff > maxBackoff { + backoff = maxBackoff + } + continue + } + + slog.Info("Connected to upstream", slog.String("url", currentUpstream)) + backoff = time.Second // Reset backoff on successful connection + + // Read messages from upstream and broadcast them + for { + messageType, message, err := conn.ReadMessage() + if err != nil { + slog.Error("Error reading from upstream", slog.Any("error", err)) + conn.Close() + break + } + + // Only broadcast text/binary messages + if messageType == websocket.TextMessage || messageType == websocket.BinaryMessage { + broadcaster.Broadcast(message) + } + } + + // Connection lost, will re-evaluate best upstream on next iteration + slog.Info("Connection lost, finding new upstream", slog.Duration("backoff", backoff)) + time.Sleep(backoff) + backoff *= 2 + if backoff > maxBackoff { + backoff = maxBackoff + } + } +} + func findBestUpstream(pool []string) (string, error) { // Measure latency concurrently @@ -96,10 +248,35 @@ func main() { slog.SetLogLoggerLevel(slog.LevelDebug) } - bestUpstream, err := findBestUpstream(pool) - if err != nil { + envPort := os.Getenv("PORT") + port := envPort + if envPort == "" { + port = "8096" + } + + envHost := os.Getenv("HOST") + host := envHost + if envHost == "" { + // should be running on the same hardware as your service + host = "127.0.0.1" + } + + bindAddr := fmt.Sprintf("%s:%s", host, port) + + // Create broadcaster and start upstream connection + // connectToUpstream will continuously find the best upstream and reconnect on failures + broadcaster := &Broadcaster{} + go connectToUpstream(pool, broadcaster) + + // Setup HTTP server + http.HandleFunc("/subscribe", handleSubscribe(broadcaster)) + + slog.Info("Starting proxy server", slog.String("bind", bindAddr)) + if err := http.ListenAndServe(bindAddr, nil); err != nil { + slog.Error("Server failed", slog.Any("error", err)) panic(err) } - fmt.Println(bestUpstream) + // TODO (future) let zlib compression be env'd + // TODO: the proxy subscribes to all lexicons, but then filters out at client level. add env var for lex filtering too }