package main import ( "context" "encoding/json" "errors" "io" "log/slog" "net" "net/http" "net/url" "strings" "sync" "time" "golang.org/x/time/rate" "github.com/coder/websocket" ) // chatMessage is the payload sent from the frontend to /publish. type chatMessage struct { To string `json:"to"` Message string `json:"message"` From string `json:"from,omitempty"` // optional, server fills this in } // chatServer enables 1:1 messaging between named subscribers. type chatServer struct { // subscriberMessageBuffer controls the max number // of messages that can be queued for a subscriber // before it is kicked. subscriberMessageBuffer int // publishLimiter controls the rate limit applied to the publish endpoint. publishLimiter *rate.Limiter // logf controls where logs are sent. logf func(f string, v ...any) // serveMux routes the various endpoints. serveMux http.ServeMux // subscribersMu protects the subscribers map. // Key is the username (from query param). subscribersMu sync.RWMutex subscribers map[string]*subscriber } // newChatServer constructs a chatServer with the defaults. func newChatServer() *chatServer { cs := &chatServer{ subscriberMessageBuffer: 16, logf: slog.Info, subscribers: make(map[string]*subscriber), publishLimiter: rate.NewLimiter(rate.Every(time.Millisecond*100), 8), } cs.serveMux.Handle("/", http.FileServer(http.Dir("."))) cs.serveMux.HandleFunc("/subscribe", cs.subscribeHandler) cs.serveMux.HandleFunc("/publish", cs.publishHandler) cs.serveMux.HandleFunc("/users", cs.usersHandler) // list online users return cs } // subscriber represents a single connected client. type subscriber struct { username string // unique identifier from query param msgs chan []byte // buffered channel for outbound messages closeSlow func() // kicks the subscriber if buffer fills cancel context.CancelFunc // cancels the read context on disconnect } // ServeHTTP delegates to the internal serveMux. func (cs *chatServer) ServeHTTP(w http.ResponseWriter, r *http.Request) { cs.serveMux.ServeHTTP(w, r) } // subscribeHandler extracts the username from query params and upgrades to WebSocket. func (cs *chatServer) subscribeHandler(w http.ResponseWriter, r *http.Request) { // Extract username from ?user= parameter username := r.URL.Query().Get("user") if username == "" { http.Error(w, "missing user parameter", http.StatusBadRequest) return } username = url.PathEscape(strings.TrimSpace(username)) err := cs.subscribe(w, r, username) if errors.Is(err, context.Canceled) { return } if websocket.CloseStatus(err) == websocket.StatusNormalClosure || websocket.CloseStatus(err) == websocket.StatusGoingAway { return } if err != nil { cs.logf("subscribe error for user %s: %v", username, err) return } } // subscribe upgrades the HTTP connection to WebSocket and registers the subscriber. func (cs *chatServer) subscribe(w http.ResponseWriter, r *http.Request, username string) error { var mu sync.Mutex var c *websocket.Conn var closed bool s := &subscriber{ username: username, 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 BEFORE accepting WebSocket cs.addSubscriber(s) defer cs.deleteSubscriber(s) // Upgrade to WebSocket c2, err := websocket.Accept(w, r, nil) if err != nil { return err } // Check if already closed mu.Lock() if closed { mu.Unlock() return net.ErrClosed } c = c2 mu.Unlock() defer c.CloseNow() // Create context that cancels on close/read error ctx, cancel := context.WithCancel(context.Background()) s.cancel = cancel // CloseRead handles ping/pong/close frames automatically connCtx := c.CloseRead(ctx) cs.logf("user %s connected", username) for { select { case msg := <-s.msgs: err := writeTimeout(connCtx, time.Second*5, c, msg) if err != nil { return err } case <-connCtx.Done(): cs.logf("user %s disconnected", username) cancel() // trigger cleanup return connCtx.Err() } } } // publishHandler reads the JSON message and routes it to the recipient. func (cs *chatServer) publishHandler(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { http.Error(w, http.StatusText(http.StatusMethodNotAllowed), http.StatusMethodNotAllowed) return } body := http.MaxBytesReader(w, r.Body, 8192) data, err := io.ReadAll(body) if err != nil { http.Error(w, http.StatusText(http.StatusRequestEntityTooLarge), http.StatusRequestEntityTooLarge) return } var msg chatMessage if err := json.Unmarshal(data, &msg); err != nil { http.Error(w, "invalid JSON", http.StatusBadRequest) return } if msg.To == "" || msg.Message == "" { http.Error(w, "missing 'to' or 'message' field", http.StatusBadRequest) return } // Deliver the message to the recipient delivered := cs.deliverMessage(msg.To, msg.Message) if delivered { w.WriteHeader(http.StatusAccepted) } else { http.Error(w, "recipient not found", http.StatusNotFound) } } // deliverMessage sends a message to a specific recipient. func (cs *chatServer) deliverMessage(toUsername, message string) bool { cs.subscribersMu.RLock() recipient, ok := cs.subscribers[toUsername] cs.subscribersMu.RUnlock() if !ok { return false } cs.publishLimiter.Wait(context.Background()) select { case recipient.msgs <- []byte(message): return true default: // Buffer full — kick slow client go recipient.closeSlow() return false } } // usersHandler returns list of currently connected usernames. func (cs *chatServer) usersHandler(w http.ResponseWriter, r *http.Request) { cs.subscribersMu.RLock() users := make([]string, 0, len(cs.subscribers)) for username := range cs.subscribers { users = append(users, username) } cs.subscribersMu.RUnlock() w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(users) } // addSubscriber registers a subscriber. func (cs *chatServer) addSubscriber(s *subscriber) { cs.subscribersMu.Lock() cs.subscribers[s.username] = s cs.subscribersMu.Unlock() } // deleteSubscriber removes a subscriber. func (cs *chatServer) deleteSubscriber(s *subscriber) { cs.subscribersMu.Lock() delete(cs.subscribers, s.username) cs.subscribersMu.Unlock() } 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) }