package main import ( "context" "errors" "io" "log" "net" "net/http" "sync" "time" "golang.org/x/time/rate" "github.com/coder/websocket" ) // chatServer enables broadcasting to a set of subscribers. type chatServer struct { // subscriberMessageBuffer controls the max number // of messages that can be queued for a subscriber // before it is kicked. // // Defaults to 16. subscriberMessageBuffer int // publishLimiter controls the rate limit applied to the publish endpoint. // // Defaults to one publish every 100ms with a burst of 8. publishLimiter *rate.Limiter // logf controls where logs are sent. // Defaults to log.Printf. logf func(f string, v ...any) // serveMux routes the various endpoints to the appropriate handler. serveMux http.ServeMux subscribersMu sync.Mutex subscribers map[*subscriber]struct{} } // newChatServer constructs a chatServer with the defaults. func newChatServer() *chatServer { cs := &chatServer{ subscriberMessageBuffer: 16, logf: log.Printf, subscribers: make(map[*subscriber]struct{}), publishLimiter: rate.NewLimiter(rate.Every(time.Millisecond*100), 8), } // Serve static files (index.html, index.js, index.css) from cwd cs.serveMux.Handle("/", http.FileServer(http.Dir("."))) // WebSocket endpoint: clients connect here to receive messages cs.serveMux.HandleFunc("/subscribe", cs.subscribeHandler) // HTTP POST endpoint: clients POST here to send a message cs.serveMux.HandleFunc("/publish", cs.publishHandler) return cs } // subscriber represents a single connected client. // Messages are sent on the msgs channel. If the client // cannot keep up (buffer full), closeSlow is called // to disconnect them. type subscriber struct { msgs chan []byte closeSlow func() } // ServeHTTP makes chatServer implement http.Handler, // delegating to the internal serveMux. func (cs *chatServer) ServeHTTP(w http.ResponseWriter, r *http.Request) { cs.serveMux.ServeHTTP(w, r) } // subscribeHandler accepts the WebSocket connection and then subscribes // it to all future messages. func (cs *chatServer) subscribeHandler(w http.ResponseWriter, r *http.Request) { err := cs.subscribe(w, r) if errors.Is(err, context.Canceled) { return } if websocket.CloseStatus(err) == websocket.StatusNormalClosure || websocket.CloseStatus(err) == websocket.StatusGoingAway { return } if err != nil { cs.logf("%v", err) return } } // publishHandler reads the request body with a limit of 8192 bytes // and then publishes the received message. func (cs *chatServer) publishHandler(w http.ResponseWriter, r *http.Request) { if r.Method != "POST" { http.Error(w, http.StatusText(http.StatusMethodNotAllowed), http.StatusMethodNotAllowed) return } body := http.MaxBytesReader(w, r.Body, 8192) msg, err := io.ReadAll(body) if err != nil { http.Error(w, http.StatusText(http.StatusRequestEntityTooLarge), http.StatusRequestEntityTooLarge) return } cs.publish(msg) w.WriteHeader(http.StatusAccepted) } // subscribe upgrades the HTTP connection to a WebSocket, // registers the subscriber, and then loops forever writing // messages from the subscriber's channel to the WebSocket. // // It uses CloseRead to keep reading control frames (close, ping, pong) // and cancel the context if the connection drops. This means we // never actually read message data from the WebSocket — // messages come in via /publish instead. func (cs *chatServer) subscribe(w http.ResponseWriter, r *http.Request) error { // We need a mutex here because the WebSocket connection might // be closed by closeSlow (from another goroutine) before we've // assigned it to c. The mutex ensures we don't race. var mu sync.Mutex var c *websocket.Conn var closed bool s := &subscriber{ msgs: make(chan []byte, cs.subscriberMessageBuffer), closeSlow: func() { mu.Lock() defer mu.Unlock() closed = true if c != nil { c.Close(websocket.StatusPolicyViolation, "connection too slow to keep up with messages") } }, } // Register the subscriber BEFORE accepting the WebSocket. // This means we won't miss any messages published between // the accept and the registration. cs.addSubscriber(s) defer cs.deleteSubscriber(s) // Upgrade HTTP to WebSocket c2, err := websocket.Accept(w, r, nil) if err != nil { return err } // Check if closeSlow was already called before we got the connection mu.Lock() if closed { mu.Unlock() return net.ErrClosed } c = c2 mu.Unlock() defer c.CloseNow() // CloseRead returns a context that is cancelled when: // - The client sends a close frame // - The connection drops // This frees us from manually handling ping/pong and deadlines. ctx := c.CloseRead(context.Background()) // Main loop: write messages to the WebSocket as they arrive for { select { case msg := <-s.msgs: err := writeTimeout(ctx, time.Second*5, c, msg) if err != nil { return err } case <-ctx.Done(): // Context cancelled = client disconnected return ctx.Err() } } } // publish sends a message to ALL subscribers. // It never blocks — if a subscriber's buffer is full, // they get kicked via closeSlow. func (cs *chatServer) publish(msg []byte) { cs.subscribersMu.Lock() defer cs.subscribersMu.Unlock() // Rate limit: wait for a token before broadcasting cs.publishLimiter.Wait(context.Background()) for s := range cs.subscribers { select { case s.msgs <- msg: // Message delivered to subscriber's buffer default: // Buffer is full — kick the slow subscriber go s.closeSlow() } } } // addSubscriber registers a subscriber. func (cs *chatServer) addSubscriber(s *subscriber) { cs.subscribersMu.Lock() cs.subscribers[s] = struct{}{} cs.subscribersMu.Unlock() } // deleteSubscriber removes a subscriber. func (cs *chatServer) deleteSubscriber(s *subscriber) { cs.subscribersMu.Lock() delete(cs.subscribers, s) cs.subscribersMu.Unlock() } // writeTimeout writes a message to the WebSocket with a deadline. // If the write doesn't complete within the timeout, the context // is cancelled and the operation aborts. func writeTimeout(ctx context.Context, timeout time.Duration, c *websocket.Conn, msg []byte) error { ctx, cancel := context.WithTimeout(ctx, timeout) defer cancel() return c.Write(ctx, websocket.MessageText, msg) }