package server import ( "context" _ "embed" "encoding/json" "errors" "fmt" "log" "net" "net/http" "os" "os/signal" "strings" "sync" "syscall" "time" "billboard/internal/state" ) type Config struct { Addr string Token string SnapshotPath string } type Server struct { st *state.State cfg Config hub *hub limit *rateLimiter } func New(st *state.State, cfg Config) *Server { return &Server{ st: st, cfg: cfg, hub: newHub(), limit: newRateLimiter(), } } //go:embed static/index.html var indexHTML []byte func (s *Server) Handler() http.Handler { mux := http.NewServeMux() mux.HandleFunc("GET /{$}", func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/html; charset=utf-8") w.Write(indexHTML) }) mux.HandleFunc("GET /api/state", s.handleState) mux.HandleFunc("GET /api/events", s.handleEvents) mux.HandleFunc("POST /api/pixel", s.handlePixel) mux.HandleFunc("POST /api/pixels", s.handlePixels) mux.HandleFunc("POST /api/image", s.handleImage) mux.HandleFunc("POST /api/message", s.handleMessage) return mux } func (s *Server) Run() error { if s.cfg.SnapshotPath != "" { go s.snapshotLoop() } httpSrv := &http.Server{Addr: s.cfg.Addr, Handler: s.Handler()} go func() { sig := make(chan os.Signal, 1) signal.Notify(sig, os.Interrupt, syscall.SIGTERM) <-sig if s.cfg.SnapshotPath != "" { if err := s.st.Save(s.cfg.SnapshotPath); err != nil { log.Printf("shutdown snapshot save: %v", err) } else { log.Printf("saved snapshot to %s", s.cfg.SnapshotPath) } } httpSrv.Shutdown(context.Background()) }() log.Printf("billboard server listening on %s (canvas %dx%d)", s.cfg.Addr, s.st.Width, s.st.Height) if err := httpSrv.ListenAndServe(); err != http.ErrServerClosed { return err } return nil } func (s *Server) snapshotLoop() { ticker := time.NewTicker(30 * time.Second) defer ticker.Stop() for range ticker.C { if err := s.st.Save(s.cfg.SnapshotPath); err != nil { log.Printf("snapshot save: %v", err) } } } func clientIP(r *http.Request) string { host, _, err := net.SplitHostPort(r.RemoteAddr) if err != nil { return r.RemoteAddr } return host } func (s *Server) authorized(r *http.Request) bool { if s.cfg.Token == "" { return true } return r.Header.Get("Authorization") == "Bearer "+s.cfg.Token } func writeErr(w http.ResponseWriter, code int, msg string) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(code) json.NewEncoder(w).Encode(map[string]string{"error": msg}) } func decode(w http.ResponseWriter, r *http.Request, v any) bool { return decodeLimit(w, r, v, 64<<10) } func decodeLimit(w http.ResponseWriter, r *http.Request, v any, maxBytes int64) bool { r.Body = http.MaxBytesReader(w, r.Body, maxBytes) if err := json.NewDecoder(r.Body).Decode(v); err != nil { writeErr(w, http.StatusBadRequest, "invalid JSON: "+err.Error()) return false } return true } func stateErr(w http.ResponseWriter, err error) { switch { case errors.Is(err, state.ErrOutOfBounds), errors.Is(err, state.ErrInvalidColor), errors.Is(err, state.ErrImageTooBig), errors.Is(err, state.ErrEmptyMessage), errors.Is(err, state.ErrMessageLong): writeErr(w, http.StatusBadRequest, err.Error()) default: writeErr(w, http.StatusInternalServerError, "internal error") } } func (s *Server) guard(w http.ResponseWriter, r *http.Request, cost float64) bool { if !s.authorized(r) { writeErr(w, http.StatusUnauthorized, "missing or invalid token") return false } if !s.limit.allow(clientIP(r), cost) { writeErr(w, http.StatusTooManyRequests, "rate limit exceeded") return false } return true } func (s *Server) handleState(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(s.st.Snapshot()) } type pixelReq struct { X int `json:"x"` Y int `json:"y"` Color uint8 `json:"color"` } func (s *Server) handlePixel(w http.ResponseWriter, r *http.Request) { if !s.guard(w, r, 1) { return } var req pixelReq if !decode(w, r, &req) { return } if err := s.st.SetPixel(req.X, req.Y, req.Color); err != nil { stateErr(w, err) return } s.hub.broadcast("pixel") w.WriteHeader(http.StatusNoContent) } func (s *Server) handlePixels(w http.ResponseWriter, r *http.Request) { var req struct { Pixels []pixelReq `json:"pixels"` } if !decodeLimit(w, r, &req, 256<<10) { return } if len(req.Pixels) == 0 || len(req.Pixels) > 2048 { writeErr(w, http.StatusBadRequest, "batch must contain 1-2048 pixels") return } cost := float64(len(req.Pixels)) / 8 if cost < 1 { cost = 1 } if !s.guard(w, r, cost) { return } for _, p := range req.Pixels { if err := s.st.CheckPixel(p.X, p.Y, p.Color); err != nil { stateErr(w, err) return } } for _, p := range req.Pixels { s.st.SetPixel(p.X, p.Y, p.Color) } s.hub.broadcast("pixel") w.WriteHeader(http.StatusNoContent) } func (s *Server) handleImage(w http.ResponseWriter, r *http.Request) { if !s.guard(w, r, 10) { return } var req struct { X int `json:"x"` Y int `json:"y"` Pixels []state.Row `json:"pixels"` } if !decode(w, r, &req) { return } if err := s.st.StampImage(req.X, req.Y, req.Pixels); err != nil { stateErr(w, err) return } s.hub.broadcast("image") w.WriteHeader(http.StatusNoContent) } func (s *Server) handleMessage(w http.ResponseWriter, r *http.Request) { if !s.guard(w, r, 5) { return } var req struct { Text string `json:"text"` } if !decode(w, r, &req) { return } req.Text = strings.TrimSpace(req.Text) if err := s.st.AddMessage(req.Text); err != nil { stateErr(w, err) return } s.hub.broadcast("message") w.WriteHeader(http.StatusNoContent) } func (s *Server) handleEvents(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.Header().Set("Connection", "keep-alive") flusher, ok := w.(http.Flusher) if !ok { writeErr(w, http.StatusInternalServerError, "streaming unsupported") return } ch := s.hub.subscribe() defer s.hub.unsubscribe(ch) fmt.Fprint(w, "event: hello\ndata: {}\n\n") flusher.Flush() for { select { case <-r.Context().Done(): return case ev := <-ch: fmt.Fprintf(w, "event: %s\ndata: {}\n\n", ev) flusher.Flush() } } } type hub struct { mu sync.Mutex subs map[chan string]struct{} } func newHub() *hub { return &hub{subs: make(map[chan string]struct{})} } func (h *hub) subscribe() chan string { ch := make(chan string, 16) h.mu.Lock() h.subs[ch] = struct{}{} h.mu.Unlock() return ch } func (h *hub) unsubscribe(ch chan string) { h.mu.Lock() delete(h.subs, ch) h.mu.Unlock() } func (h *hub) broadcast(ev string) { h.mu.Lock() defer h.mu.Unlock() for ch := range h.subs { select { case ch <- ev: default: } } } type bucket struct { tokens float64 last time.Time } type rateLimiter struct { mu sync.Mutex buckets map[string]*bucket } func newRateLimiter() *rateLimiter { rl := &rateLimiter{buckets: make(map[string]*bucket)} go rl.gc() return rl } const ( ratePerSec = 4.0 burst = 40.0 ) func (rl *rateLimiter) allow(key string, cost float64) bool { rl.mu.Lock() defer rl.mu.Unlock() b, ok := rl.buckets[key] if !ok { b = &bucket{tokens: burst, last: time.Now()} rl.buckets[key] = b } now := time.Now() b.tokens += now.Sub(b.last).Seconds() * ratePerSec if b.tokens > burst { b.tokens = burst } b.last = now if b.tokens < cost { return false } b.tokens -= cost return true } func (rl *rateLimiter) gc() { ticker := time.NewTicker(10 * time.Minute) defer ticker.Stop() for range ticker.C { rl.mu.Lock() for k, b := range rl.buckets { if time.Since(b.last) > 10*time.Minute { delete(rl.buckets, k) } } rl.mu.Unlock() } }