Ronbun/cmd/web/helpers.go
2026-07-11 17:00:58 +02:00

306 lines
7.5 KiB
Go

package main
import (
"bytes"
"errors"
"fmt"
"html/template"
"io"
"io/fs"
"mime/multipart"
"net/http"
"os"
"os/exec"
"path/filepath"
"regexp"
"runtime/debug"
"slices"
"strings"
"time"
"github.com/Thomasorus/Ronbun-CMS/internal/models"
"github.com/go-playground/form/v4"
"github.com/justinas/nosurf"
)
// Writes log entry and sends 500 error to the user
func (app *application) serverError(w http.ResponseWriter, r *http.Request, err error) {
var (
method = r.Method
uri = r.URL.RequestURI()
trace = string(debug.Stack())
)
app.logger.Error(err.Error(), "method", method, "uri", uri, "trace", trace)
if app.debug {
body := fmt.Sprintf("%s\n%s", err, trace)
http.Error(w, body, http.StatusInternalServerError)
return
}
http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError)
}
// Sends specific server error to the user
func (app *application) clientError(w http.ResponseWriter, status int) {
http.Error(w, http.StatusText(status), status)
}
// Sends specific server error to the user + error message
func (app *application) clientErrorMessage(w http.ResponseWriter, error string, status int) {
http.Error(w, error, status)
}
// Renders templates from cache
func (app *application) render(w http.ResponseWriter, r *http.Request, status int, page string, data templateData) {
ts, ok := app.templateCache[page]
if !ok {
err := fmt.Errorf("the template %s does not exist", page)
app.serverError(w, r, err)
return
}
buf := new(bytes.Buffer)
err := ts.ExecuteTemplate(buf, "base", data)
if err != nil {
app.serverError(w, r, err)
return
}
w.WriteHeader(status)
buf.WriteTo(w)
}
// Generic decode form
func (app *application) decodePostForm(r *http.Request, dst any) error {
err := r.ParseForm()
if err != nil {
return err
}
err = app.formDecoder.Decode(dst, r.PostForm)
if err != nil {
// If we try to use an invalid target destination, the Decode() method
// will return an error with the type form.InvalidDecodeError. We use
// errors.AsType() to check for this and panic.
if _, ok := errors.AsType[*form.InvalidDecoderError](err); ok {
panic(err)
}
return err
}
return nil
}
// Helper to return the current year in templateData
func (app *application) newTemplateData(r *http.Request) templateData {
return templateData{
CurrentYear: time.Now().Year(),
Flash: app.sessionManager.PopString(r.Context(), "flash"),
IsAuthenticated: app.isAuthenticated(r),
CSRFToken: nosurf.Token(r),
Files: app.filesList("uploaded"),
Theme: app.getThemeFromRequest(r),
}
}
func (app *application) getThemeFromRequest(r *http.Request) string {
cookie, err := r.Cookie("theme")
if err != nil {
return "color"
}
if cookie.Value != "dark" && cookie.Value != "light" && cookie.Value != "raw" {
return "color"
}
return cookie.Value
}
// Return true if the current request is from an authenticated user, else false
func (app *application) isAuthenticated(r *http.Request) bool {
isAuthenticated, ok := r.Context().Value(isAuthenticatedContextKey).(bool)
if !ok {
return false
}
return isAuthenticated
}
// sanitizeFilename removes dangerous characters and path traversal attempts from filename
func (app *application) sanitizeFilename(filename string) string {
// Get just the base filename (removes any path components)
filename = filepath.Base(filename)
// Remove or replace dangerous characters
re := regexp.MustCompile(`[<>:"/\\|?*\x00-\x1f]`)
filename = re.ReplaceAllString(filename, "_")
// Replace spaces with hyphens
filename = strings.ReplaceAll(filename, " ", "-")
// Limit filename length
if len(filename) > 255 {
ext := filepath.Ext(filename)
name := filename[:255-len(ext)]
filename = name + ext
}
// Prevent empty or dangerous filenames
if filename == "" || filename == "." || filename == ".." {
filename = "upload_file"
}
return filename
}
type File struct {
Filename string
Extension string
Slug string
}
// Return a list of fileNames from the /img folder
func (app *application) filesList(inputPath string) []File {
var filesPaths []File
_ = filepath.WalkDir(inputPath, func(path string, d fs.DirEntry, err error) error {
if err != nil {
return err
}
if !d.IsDir() {
var file File
file.Extension = filepath.Ext(path)
file.Filename = filepath.Base(strings.TrimSuffix(path, file.Extension))
file.Slug = "/" + path
filesPaths = append(filesPaths, file)
}
return nil
})
return filesPaths
}
// Make a number of checks on files
func (app *application) processFile(fileHeader *multipart.FileHeader, safeFilename, filesDir string, processor *ImageProcessor) error {
destPath := filepath.Join(filesDir, safeFilename)
// Validate path is within filesDir
if !strings.HasPrefix(filepath.Clean(destPath), filepath.Clean(filesDir)) {
return fmt.Errorf("invalid filename")
}
// Check if file already exists
if _, err := os.Stat(destPath); err == nil {
return fmt.Errorf("file '%s' already exists", safeFilename)
}
file, err := fileHeader.Open()
if err != nil {
return err
}
defer file.Close()
dst, err := os.Create(destPath)
if err != nil {
return err
}
defer dst.Close()
if _, err := io.Copy(dst, file); err != nil {
os.Remove(destPath)
return err
}
extension := strings.ToLower(filepath.Ext(safeFilename))
if extension == ".heif" || extension == ".heic" {
baseName := strings.TrimSuffix(filepath.Base(destPath), filepath.Ext(destPath))
jpgDest := processor.UploadDir + "/" + baseName + ".jpg"
heifConvertPath := "" // Choose between production and dev
if app.devMode {
heifConvertPath = "/opt/homebrew/bin/heif-convert" // dev
} else {
heifConvertPath = "/usr/bin/heif-convert" // prod
}
cmd := exec.Command(heifConvertPath, destPath, jpgDest)
err := cmd.Run()
if err != nil {
os.Remove(destPath)
return fmt.Errorf("ERROR: %v", err)
}
os.Remove(destPath)
destPath = jpgDest
}
// Process image if it's an image file
if app.isImageFile(fileHeader.Filename) {
if err := processor.ProcessImage(destPath); err != nil {
os.Remove(destPath)
return err
}
}
return nil
}
// Check if file is an image
func (app *application) isImageFile(filename string) bool {
ext := strings.ToLower(filepath.Ext(filename))
imageExts := []string{".jpg", ".jpeg", ".png", ".gif", ".webp", ".avif", ".tiff", ".tif", ".heif", ".heic"}
return slices.Contains(imageExts, ext)
}
// Build a tree based on host ids
func (app *application) BuildTree(pages []*models.Page) template.HTML {
childrenMap := make(map[int][]*models.Page)
for _, p := range pages {
childrenMap[p.Host] = append(childrenMap[p.Host], p)
}
var buf strings.Builder
buf.WriteString("<ul>")
for _, p := range childrenMap[0] {
if p.ID != 0 {
app.buildNode(&buf, p, childrenMap)
}
}
buf.WriteString("</ul>")
return template.HTML(buf.String())
}
func (app *application) buildNode(buf *strings.Builder, p *models.Page, childrenMap map[int][]*models.Page) {
buf.WriteString("<li><a href='")
buf.WriteString(template.HTMLEscapeString(p.Slug))
buf.WriteString("'>")
buf.WriteString(template.HTMLEscapeString(p.Name))
buf.WriteString("</a>")
if children, ok := childrenMap[p.ID]; ok && len(children) > 0 {
buf.WriteString("<ul>")
for _, child := range children {
app.buildNode(buf, child, childrenMap)
}
buf.WriteString("</ul>")
}
buf.WriteString("</li>")
}
// Checks if redirection path is safe
func (app *application) isSafeLocalPath(p string) bool {
if p == "" || p[0] != '/' {
return false
}
if len(p) > 1 && p[1] == '/' {
return false
}
return true
}