package web
import (
"embed"
"fmt"
"html/template"
"io"
"io/fs"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"time"
)
//go:embed templates static
var assets embed.FS
// templates compiles each page together with its layout. In dev mode it
// re-reads from disk on every render so edits show up without a restart.
type templates struct {
dev bool
fs fs.FS
funcs template.FuncMap
mu sync.Mutex
cache map[string]*template.Template
}
func newTemplates(dev bool, funcs template.FuncMap) *templates {
t := &templates{dev: dev, funcs: funcs, cache: map[string]*template.Template{}}
t.fs = assets
if dev {
// locate the package directory so `go run` from anywhere still finds the files
_, file, _, _ := runtime.Caller(0)
dir := filepath.Dir(file)
if _, err := os.Stat(filepath.Join(dir, "templates")); err == nil {
t.fs = os.DirFS(dir)
}
}
return t
}
// AssetsFS returns the static files (embedded or on disk in dev mode).
func (t *templates) AssetsFS() fs.FS {
sub, _ := fs.Sub(t.fs, "static")
return sub
}
func layoutFor(name string) string {
if strings.HasPrefix(name, "blog/") {
return "layouts/blog.html"
}
return "layouts/dashboard.html"
}
func (t *templates) get(name string) (*template.Template, error) {
if !t.dev {
t.mu.Lock()
if tpl, ok := t.cache[name]; ok {
t.mu.Unlock()
return tpl, nil
}
t.mu.Unlock()
}
tpl, err := template.New("").Funcs(t.funcs).ParseFS(t.fs,
"templates/"+layoutFor(name), "templates/partials/*.html", "templates/"+name)
if err != nil {
return nil, fmt.Errorf("parse %s: %w", name, err)
}
if !t.dev {
t.mu.Lock()
t.cache[name] = tpl
t.mu.Unlock()
}
return tpl, nil
}
func (t *templates) render(w io.Writer, name string, data any) error {
tpl, err := t.get(name)
if err != nil {
return err
}
return tpl.ExecuteTemplate(w, "layout", data)
}
var funcs = template.FuncMap{
"date": func(t time.Time) string { return t.Format("2 January 2006") },
"rfc": func(t time.Time) string { return t.Format(time.RFC1123Z) },
"html": func(s string) template.HTML { return template.HTML(s) },
"css": func(s string) template.CSS { return template.CSS(s) },
"kb": func(n int) string { return fmt.Sprintf("%.0f KB", float64(n)/1024) },
"add": func(a, b int) int { return a + b },
"sub": func(a, b int) int { return a - b },
"deref": func(p *string) string {
if p == nil {
return ""
}
return *p
},
"lower": strings.ToLower,
// dict builds a map for passing several values to a partial: {{template "x" (dict "a" 1 "b" 2)}}
"dict": func(kv ...any) (map[string]any, error) {
if len(kv)%2 != 0 {
return nil, fmt.Errorf("dict: odd number of arguments")
}
m := make(map[string]any, len(kv)/2)
for i := 0; i < len(kv); i += 2 {
k, ok := kv[i].(string)
if !ok {
return nil, fmt.Errorf("dict: key %v is not a string", kv[i])
}
m[k] = kv[i+1]
}
return m, nil
},
}