package middleware import ( "context" "errors" "net/http" "time" ) // Timeout returns middleware that gives each request limit to finish: // it cancels the request's context once limit has passed, and answers // 504 when the handler then returns without having started its // response. // // It replaces chi's middleware.Timeout, which writes that 504 even // after the handler has sent its own status. A download that outlasts // the limit has already sent its 200 and the whole file, so the late // 504 changes nothing for the client: the access log and the metrics // would record it in place of the 200, and net/http would complain of // a superfluous WriteHeader. func (s *Middleware) Timeout( limit time.Duration, ) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func( w http.ResponseWriter, r *http.Request, ) { ctx, cancel := context.WithTimeout(r.Context(), limit) defer cancel() tw := &timeoutResponseWriter{ResponseWriter: w} next.ServeHTTP(tw, r.WithContext(ctx)) if !tw.started && errors.Is(ctx.Err(), context.DeadlineExceeded) { w.WriteHeader(http.StatusGatewayTimeout) } }) } } // timeoutResponseWriter records whether the handler has started its // response. type timeoutResponseWriter struct { http.ResponseWriter started bool } func (w *timeoutResponseWriter) WriteHeader(code int) { w.started = true w.ResponseWriter.WriteHeader(code) } func (w *timeoutResponseWriter) Write(b []byte) (int, error) { // A Write without a WriteHeader starts the response too: net/http // sends 200 in front of it. w.started = true //nolint:wrapcheck // Pass the writer's own error through unchanged. return w.ResponseWriter.Write(b) } // Unwrap lets http.ResponseController reach the writer underneath, so // a handler can still set a write deadline through this wrapper. func (w *timeoutResponseWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter }