@@ -38,25 +38,29 @@ type visitor struct {
3838}
3939
4040func newRateLimiter (rate int , window time.Duration ) * rateLimiter {
41- rl := & rateLimiter {
41+ return & rateLimiter {
4242 visitors : make (map [string ]* visitor ),
4343 rate : rate ,
4444 window : window ,
4545 }
46- go rl .cleanup ()
47- return rl
4846}
4947
50- func (rl * rateLimiter ) cleanup () {
48+ func (rl * rateLimiter ) cleanup (ctx context.Context ) {
49+ ticker := time .NewTicker (time .Minute )
50+ defer ticker .Stop ()
5151 for {
52- time .Sleep (time .Minute )
53- rl .mu .Lock ()
54- for ip , v := range rl .visitors {
55- if time .Since (v .windowStart ) > rl .window {
56- delete (rl .visitors , ip )
52+ select {
53+ case <- ticker .C :
54+ rl .mu .Lock ()
55+ for ip , v := range rl .visitors {
56+ if time .Since (v .windowStart ) > rl .window {
57+ delete (rl .visitors , ip )
58+ }
5759 }
60+ rl .mu .Unlock ()
61+ case <- ctx .Done ():
62+ return
5863 }
59- rl .mu .Unlock ()
6064 }
6165}
6266
@@ -76,12 +80,27 @@ func (rl *rateLimiter) allow(ip string) bool {
7680 return v .count <= rl .rate
7781}
7882
83+ // clientIP extracts the client IP, preferring X-Forwarded-For behind a proxy.
84+ func clientIP (r * http.Request ) string {
85+ if xff := r .Header .Get ("X-Forwarded-For" ); xff != "" {
86+ // First entry is the original client
87+ if ip := strings .TrimSpace (strings .SplitN (xff , "," , 2 )[0 ]); ip != "" {
88+ return ip
89+ }
90+ }
91+ if xri := r .Header .Get ("X-Real-IP" ); xri != "" {
92+ return strings .TrimSpace (xri )
93+ }
94+ ip , _ , err := net .SplitHostPort (r .RemoteAddr )
95+ if err != nil {
96+ return r .RemoteAddr
97+ }
98+ return ip
99+ }
100+
79101func rateLimitMiddleware (limiter * rateLimiter , next http.Handler ) http.Handler {
80102 return http .HandlerFunc (func (w http.ResponseWriter , r * http.Request ) {
81- ip , _ , err := net .SplitHostPort (r .RemoteAddr )
82- if err != nil {
83- ip = r .RemoteAddr
84- }
103+ ip := clientIP (r )
85104 if ! limiter .allow (ip ) {
86105 slog .Warn ("rate limit exceeded" , "ip" , ip )
87106 http .Error (w , "Too Many Requests" , http .StatusTooManyRequests )
@@ -145,9 +164,13 @@ func main() {
145164
146165 queries := db .New (database )
147166
167+ // Cancellable context for background goroutines
168+ bgCtx , bgCancel := context .WithCancel (context .Background ())
169+ defer bgCancel ()
170+
148171 // Session store
149172 store := session .NewStore (queries )
150- go store .CleanupLoop ()
173+ go store .CleanupLoop (bgCtx )
151174
152175 // Auth handler
153176 authHandler := auth .NewHandler (
@@ -156,18 +179,23 @@ func main() {
156179 os .Getenv ("GOOGLE_REDIRECT_URL" ),
157180 queries ,
158181 store ,
182+ isProd ,
159183 )
160184
161- // TMDB client — refresh poster paths in background
185+ // TMDB client -- refresh poster paths in background
162186 tmdbClient := tmdb .NewClient (os .Getenv ("TMDB_API_KEY" ), queries )
163- tmdbClient .StartRefreshLoop (context . Background () )
187+ tmdbClient .StartRefreshLoop (bgCtx )
164188
165189 // Recommendation engine
166190 engine := recommend .NewEngine (queries , os .Getenv ("ANTHROPIC_API_KEY" ), tmdbClient )
167191
168192 // Handlers
169193 h := handlers .New (queries , store , engine , tmdbClient )
170194
195+ // Rate limiter with cancellable cleanup
196+ limiter := newRateLimiter (60 , time .Minute )
197+ go limiter .cleanup (bgCtx )
198+
171199 // Routes
172200 mux := http .NewServeMux ()
173201
@@ -207,7 +235,6 @@ func main() {
207235 })
208236
209237 // Middleware chain
210- limiter := newRateLimiter (60 , time .Minute )
211238 handler := middleware .RequestLogger (
212239 rateLimitMiddleware (limiter ,
213240 middleware .SecurityHeaders (isProd , mux ),
@@ -222,18 +249,27 @@ func main() {
222249 IdleTimeout : 60 * time .Second ,
223250 }
224251
252+ // Use a channel to propagate server errors to main goroutine
253+ serverErr := make (chan error , 1 )
225254 go func () {
226255 slog .Info ("HTTP server listening" , "addr" , srv .Addr )
227256 if err := srv .ListenAndServe (); err != nil && err != http .ErrServerClosed {
228- slog .Error ("server error" , "error" , err )
229- os .Exit (1 )
257+ serverErr <- err
230258 }
231259 }()
232260
233261 quit := make (chan os.Signal , 1 )
234262 signal .Notify (quit , syscall .SIGINT , syscall .SIGTERM )
235- <- quit
236- slog .Info ("shutting down server..." )
263+
264+ select {
265+ case <- quit :
266+ slog .Info ("shutting down server..." )
267+ case err := <- serverErr :
268+ slog .Error ("server error" , "error" , err )
269+ }
270+
271+ // Cancel background goroutines before shutting down the server
272+ bgCancel ()
237273
238274 ctx , cancel := context .WithTimeout (context .Background (), 5 * time .Second )
239275 defer cancel ()
@@ -246,16 +282,21 @@ func main() {
246282// migrateRatingsConstraint recreates the ratings table if the CHECK constraint
247283// does not include 'favourite'. Needed because CREATE TABLE IF NOT EXISTS skips
248284// existing tables, so the schema update alone cannot fix old databases.
249- func migrateRatingsConstraint (db * sql.DB ) error {
285+ func migrateRatingsConstraint (database * sql.DB ) error {
250286 var tableSql string
251- err := db .QueryRow ("SELECT sql FROM sqlite_master WHERE type='table' AND name='ratings'" ).Scan (& tableSql )
287+ err := database .QueryRow ("SELECT sql FROM sqlite_master WHERE type='table' AND name='ratings'" ).Scan (& tableSql )
252288 if err != nil {
253289 return nil // table doesn't exist yet, schema will create it
254290 }
255291 if strings .Contains (tableSql , "'favourite'" ) {
256292 return nil // already migrated
257293 }
258- _ , err = db .Exec (`
294+ tx , err := database .Begin ()
295+ if err != nil {
296+ return err
297+ }
298+ defer tx .Rollback ()
299+ _ , err = tx .Exec (`
259300 CREATE TABLE ratings_new (
260301 user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
261302 show_id TEXT NOT NULL REFERENCES shows(id),
@@ -268,7 +309,10 @@ func migrateRatingsConstraint(db *sql.DB) error {
268309 ALTER TABLE ratings_new RENAME TO ratings;
269310 CREATE INDEX IF NOT EXISTS idx_ratings_user ON ratings(user_id);
270311 ` )
271- return err
312+ if err != nil {
313+ return err
314+ }
315+ return tx .Commit ()
272316}
273317
274318func cacheStaticAssets (next http.Handler ) http.Handler {
0 commit comments