113 lines
3.4 KiB
Go
113 lines
3.4 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"net/http"
|
|
|
|
"github.com/justinas/nosurf"
|
|
)
|
|
|
|
// Set headers for all
|
|
func commonHeaders(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Security-Policy", "default-src 'self'; style-src 'self' 'unsafe-inline'; script-src 'self' 'unsafe-inline'; img-src 'self' data: https:; media-src 'self' data: https:; frame-src youtube.com https://www.youtube.com;")
|
|
w.Header().Set("Referrer-Policy", "origin-when-cross-origin")
|
|
w.Header().Set("X-Content-Type-Options", "nosniff")
|
|
w.Header().Set("X-Frame-Options", "deny")
|
|
w.Header().Set("X-XSS-Protection", "0")
|
|
w.Header().Set("Server", "Go")
|
|
w.Header().Set("Strict-Transport-Security", "max-age=31536000; includeSubDomains; preload")
|
|
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
|
|
func (app *application) logRequest(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
var (
|
|
ip = r.RemoteAddr
|
|
proto = r.Proto
|
|
method = r.Method
|
|
uri = r.URL.RequestURI()
|
|
)
|
|
|
|
app.logger.Info("received request", "ip", ip, "proto", proto, "method", method, "uri", uri)
|
|
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
|
|
func (app *application) recoverPanic(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
// Create a deferred function (which will always be run in the event
|
|
// of a panic as Go unwinds the stack).
|
|
defer func() {
|
|
// Use the builtin recover function to check if there has been a
|
|
// panic or not. If there has...
|
|
if err := recover(); err != nil {
|
|
// Set a "Connection: close" header on the response.
|
|
w.Header().Set("Connection", "close")
|
|
// Call the app.serverError helper method to return a 500
|
|
// Internal Server response.
|
|
app.serverError(w, r, fmt.Errorf("%v", err))
|
|
}
|
|
}()
|
|
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
|
|
// If the user is not authenticated, redirect them to the login page
|
|
func (app *application) requireAuthentication(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if !app.isAuthenticated(r) {
|
|
app.sessionManager.Put(r.Context(), "redirectPathAfterLogin", r.URL.Path)
|
|
http.Redirect(w, r, "/user/login", http.StatusSeeOther)
|
|
return
|
|
}
|
|
// set the "Cache-Control: no-store" header so that pages
|
|
// require authentication are not stored in the users browser cache (or
|
|
// other intermediary cache).
|
|
w.Header().Add("Cache-Control", "no-store")
|
|
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
|
|
func noSurf(next http.Handler) http.Handler {
|
|
csrfHandler := nosurf.New(next)
|
|
csrfHandler.SetBaseCookie(http.Cookie{
|
|
HttpOnly: true,
|
|
Path: "/",
|
|
Secure: true,
|
|
})
|
|
|
|
return csrfHandler
|
|
}
|
|
|
|
// Retrieves the authenticatedUserID value from the session, otherwise check the DB to see
|
|
// if the ID exists and if true, create a copy of the request and assign it to r.
|
|
func (app *application) authenticate(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
id := app.sessionManager.GetInt(r.Context(), "authenticatedUserID")
|
|
if id == 0 {
|
|
next.ServeHTTP(w, r)
|
|
return
|
|
}
|
|
|
|
exists, err := app.users.Exists(id)
|
|
if err != nil {
|
|
app.serverError(w, r, err)
|
|
return
|
|
}
|
|
|
|
if exists {
|
|
ctx := context.WithValue(r.Context(), isAuthenticatedContextKey, true)
|
|
r = r.WithContext(ctx)
|
|
}
|
|
|
|
// Call the next handler in the chain.
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|