From a51ac9752ba04d18128a180313c20402ae8b35e8 Mon Sep 17 00:00:00 2001 From: Bloomy Date: Thu, 30 Jul 2026 17:04:12 +0200 Subject: [PATCH 1/3] refactor: route with chi and per-resource sub-routers Handlers were constructor closures over *sql.DB, and main.go listed every route by hand. Turn them into methods on a Handlers struct holding the shared dependencies, and let each resource file own its own sub-router. Router() is now the only place that knows URL prefixes, so a resource file never repeats its own mount path. Two routing details worth calling out: - Bind /static to GET and HEAD explicitly. http.FileServer ignores the request method, so a single catch-all Handle answered POST, DELETE and TRACE with the file body. - Canonicalise paths with CleanPath and StripSlashes, so "/clients//" and "/clients/1/edit/" reach their handler instead of the catch-all. Route parity with the previous net/http mux was verified route by route: same paths, same status codes, same Location headers. Co-Authored-By: Claude Opus 5 --- cmd/time-tracker/main.go | 50 +--- go.mod | 3 +- go.sum | 2 + internal/handlers/account.go | 121 ++++---- internal/handlers/auth.go | 77 +++-- internal/handlers/clients.go | 396 ++++++++++++------------- internal/handlers/handlers.go | 59 ++++ internal/handlers/home.go | 2 +- internal/handlers/locale.go | 5 +- internal/handlers/periods.go | 147 +++++----- internal/handlers/redirect_test.go | 14 +- internal/handlers/router.go | 52 ++++ internal/handlers/router_test.go | 81 +++++ internal/handlers/task_types.go | 124 ++++---- internal/handlers/tasks.go | 422 +++++++++++++-------------- internal/handlers/time_management.go | 49 ++-- 16 files changed, 843 insertions(+), 761 deletions(-) create mode 100644 internal/handlers/handlers.go create mode 100644 internal/handlers/router.go create mode 100644 internal/handlers/router_test.go diff --git a/cmd/time-tracker/main.go b/cmd/time-tracker/main.go index 407975f..85c33de 100644 --- a/cmd/time-tracker/main.go +++ b/cmd/time-tracker/main.go @@ -14,7 +14,6 @@ import ( "github.com/bloomyindev/time-tracker/internal/db" "github.com/bloomyindev/time-tracker/internal/handlers" "github.com/bloomyindev/time-tracker/internal/i18n" - "github.com/bloomyindev/time-tracker/internal/redirect" "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/urfave/cli/v3" ) @@ -67,58 +66,15 @@ func serveCommand() *cli.Command { return fmt.Errorf("open db: %w", err) } - authSvc := auth.NewService(conn, cfg.JWTSecret) - - mux := http.NewServeMux() - - mux.HandleFunc("GET /", handlers.Home) - mux.HandleFunc("GET /lang/{code}", handlers.SetLocale) - mux.HandleFunc("GET /login", handlers.Login) - mux.HandleFunc("POST /login", handlers.LoginSubmit(authSvc)) - mux.HandleFunc("GET /logout", handlers.Logout(authSvc)) - - mux.Handle("GET /clients", authSvc.RequireAuth(handlers.ListClients(conn))) - mux.Handle("POST /clients", authSvc.RequireAuth(handlers.CreateClient(conn))) - mux.Handle("GET /clients/{id}", authSvc.RequireAuth(handlers.ClientDetail(conn))) - mux.Handle("GET /clients/{id}/report", authSvc.RequireAuth(handlers.ClientReport(conn))) - mux.Handle("GET /clients/{id}/edit", authSvc.RequireAuth(handlers.EditClientForm(conn))) - mux.Handle("POST /clients/{id}/edit", authSvc.RequireAuth(handlers.UpdateClient(conn))) - mux.Handle("POST /clients/{id}/delete", authSvc.RequireAuth(handlers.DeleteClient(conn))) - - mux.Handle("GET /task-types", authSvc.RequireAuth(handlers.ListTaskTypes(conn))) - mux.Handle("POST /task-types", authSvc.RequireAuth(handlers.CreateTaskType(conn))) - mux.Handle("GET /task-types/{id}/edit", authSvc.RequireAuth(handlers.EditTaskTypeForm(conn))) - mux.Handle("POST /task-types/{id}/rename", authSvc.RequireAuth(handlers.RenameTaskType(conn))) - mux.Handle("POST /task-types/{id}/delete", authSvc.RequireAuth(handlers.DeleteTaskType(conn))) - - mux.Handle("GET /periods", authSvc.RequireAuth(handlers.ListPeriods(conn))) - mux.Handle("POST /periods", authSvc.RequireAuth(handlers.CreatePeriod(conn))) - mux.Handle("POST /periods/{id}/default", authSvc.RequireAuth(handlers.SetDefaultPeriod(conn))) - mux.Handle("GET /periods/{id}/edit", authSvc.RequireAuth(handlers.EditPeriodForm(conn))) - mux.Handle("POST /periods/{id}/rename", authSvc.RequireAuth(handlers.RenamePeriod(conn))) - mux.Handle("POST /periods/{id}/delete", authSvc.RequireAuth(handlers.DeletePeriod(conn))) - - mux.Handle("GET /tasks", authSvc.RequireAuth(handlers.ListTasks(conn))) - mux.Handle("POST /tasks", authSvc.RequireAuth(handlers.CreateTask(conn))) - mux.Handle("GET /tasks/{id}/edit", authSvc.RequireAuth(handlers.EditTaskForm(conn))) - mux.Handle("POST /tasks/{id}/update", authSvc.RequireAuth(handlers.UpdateTask(conn))) - mux.Handle("POST /tasks/{id}/delete", authSvc.RequireAuth(handlers.DeleteTask(conn))) - - mux.Handle("GET /time", authSvc.RequireAuth(handlers.ListTimeEntries(conn))) - mux.Handle("GET /time/report", authSvc.RequireAuth(handlers.TimeReport(conn))) - - mux.Handle("GET /account", authSvc.RequireAuth(handlers.Account(conn))) - mux.Handle("POST /account/hours", authSvc.RequireAuth(handlers.UpdateDailyHours(conn))) - mux.Handle("POST /account/password", authSvc.RequireAuth(handlers.ChangePassword(conn, authSvc))) - staticFS, err := fs.Sub(assets.Static, "static") if err != nil { return fmt.Errorf("mount static assets: %w", err) } - mux.Handle("GET /static/", http.StripPrefix("/static/", http.FileServerFS(staticFS))) + + h := handlers.New(conn, auth.NewService(conn, cfg.JWTSecret)) log.Printf("listening on port %d", cfg.Port) - return http.ListenAndServe(fmt.Sprintf(":%d", cfg.Port), i18n.Middleware(redirect.Middleware(mux))) + return http.ListenAndServe(fmt.Sprintf(":%d", cfg.Port), h.Router(staticFS)) }, } } diff --git a/go.mod b/go.mod index 83cecdd..473e505 100644 --- a/go.mod +++ b/go.mod @@ -9,6 +9,8 @@ tool ( require ( github.com/a-h/templ v0.3.1020 + github.com/bryanvaz/go-templ-lucide-icons v0.480.0 + github.com/go-chi/chi/v5 v5.3.1 github.com/golang-jwt/jwt/v5 v5.3.1 github.com/invopop/ctxi18n v0.9.0 github.com/urfave/cli/v3 v3.10.1 @@ -23,7 +25,6 @@ require ( github.com/andybalholm/brotli v1.2.0 // indirect github.com/bep/godartsass/v2 v2.5.0 // indirect github.com/bep/golibsass v1.2.0 // indirect - github.com/bryanvaz/go-templ-lucide-icons v0.480.0 // indirect github.com/cenkalti/backoff/v4 v4.3.0 // indirect github.com/cli/browser v1.3.0 // indirect github.com/dustin/go-humanize v1.0.1 // indirect diff --git a/go.sum b/go.sum index 9489ca4..922f0d6 100644 --- a/go.sum +++ b/go.sum @@ -75,6 +75,8 @@ github.com/getkin/kin-openapi v0.133.0 h1:pJdmNohVIJ97r4AUFtEXRXwESr8b0bD721u/Tz github.com/getkin/kin-openapi v0.133.0/go.mod h1:boAciF6cXk5FhPqe/NQeBTeenbjqU4LhWBf09ILVvWE= github.com/ghodss/yaml v1.0.0 h1:wQHKEahhL6wmXdzwWG11gIVCkOv05bNOh+Rxn0yngAk= github.com/ghodss/yaml v1.0.0/go.mod h1:4dBDuWmgqj2HViK6kFavaiC9ZROes6MMH2rRYeMEF04= +github.com/go-chi/chi/v5 v5.3.1 h1:3j4HZLGZQ3JpMCrPJF/Jl3mYJfWLKBfNJ6quurUGCf8= +github.com/go-chi/chi/v5 v5.3.1/go.mod h1:R+tYY2hNuVUUjxoPtqUdgBqevM9s9njzkTLutVsOCto= github.com/go-openapi/jsonpointer v0.21.0 h1:YgdVicSA9vH5RiHs9TZW5oyafXZFc6+2Vc1rr/O9oNQ= github.com/go-openapi/jsonpointer v0.21.0/go.mod h1:IUyH9l/+uyhIYQ/PXVA41Rexl+kOkAPDdXEYns6fzUY= github.com/go-openapi/swag v0.23.0 h1:vsEVJDUo2hPJ2tu0/Xc+4noaxyEffXNIs3cOULZ+GrE= diff --git a/internal/handlers/account.go b/internal/handlers/account.go index fe82d53..1e07daa 100644 --- a/internal/handlers/account.go +++ b/internal/handlers/account.go @@ -1,7 +1,6 @@ package handlers import ( - "database/sql" "errors" "net/http" "strconv" @@ -9,80 +8,78 @@ import ( "github.com/bloomyindev/time-tracker/internal/db" "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/bloomyindev/time-tracker/internal/templates" + "github.com/go-chi/chi/v5" "github.com/invopop/ctxi18n/i18n" ) -func Account(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - user, err := db.GetUser(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - templates.Account(user, "", "").Render(r.Context(), w) - } +func (h *Handlers) AccountRouter() chi.Router { + r := chi.NewRouter() + r.Get("/", h.account) + r.Post("/hours", h.updateDailyHours) + r.Post("/password", h.changePassword) + return r } -func UpdateDailyHours(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - user, err := db.GetUser(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } - - var hours [7]float64 - for i := 0; i < 7; i++ { - raw := r.FormValue("hours_" + strconv.Itoa(i)) - if raw == "" { - continue - } - h, err := strconv.ParseFloat(raw, 64) - if err != nil { - http.Error(w, "invalid hours", http.StatusBadRequest) - return - } - hours[i] = h - } - - if err := db.UpdateDailyHours(conn, userID, hours); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - user.DailyHours = hours - templates.Account(user, "", i18n.T(r.Context(), "account.hours_saved")).Render(r.Context(), w) +func (h *Handlers) account(w http.ResponseWriter, r *http.Request) { + user, err := db.GetUser(h.DB, userID(r)) + if err != nil { + fail(w, err) + return } + templates.Account(user, "", "").Render(r.Context(), w) } -func ChangePassword(conn *sql.DB, svc *auth.Service) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - user, err := db.GetUser(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } +func (h *Handlers) updateDailyHours(w http.ResponseWriter, r *http.Request) { + user, err := db.GetUser(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + if !parseForm(w, r) { + return + } - err = svc.ChangePassword(userID, r.FormValue("current_password"), r.FormValue("new_password")) - if errors.Is(err, auth.ErrInvalidCredentials) { - templates.Account(user, i18n.T(r.Context(), "account.wrong_current_password"), "").Render(r.Context(), w) - return + var hours [7]float64 + for i := 0; i < 7; i++ { + raw := r.FormValue("hours_" + strconv.Itoa(i)) + if raw == "" { + continue } + hrs, err := strconv.ParseFloat(raw, 64) if err != nil { - templates.Account(user, i18n.T(r.Context(), "login.something_went_wrong"), "").Render(r.Context(), w) + http.Error(w, "invalid hours", http.StatusBadRequest) return } + hours[i] = hrs + } - templates.Account(user, "", i18n.T(r.Context(), "account.password_changed")).Render(r.Context(), w) + if err := db.UpdateDailyHours(h.DB, userID(r), hours); err != nil { + fail(w, err) + return } + user.DailyHours = hours + templates.Account(user, "", i18n.T(r.Context(), "account.hours_saved")).Render(r.Context(), w) +} + +func (h *Handlers) changePassword(w http.ResponseWriter, r *http.Request) { + user, err := db.GetUser(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + if !parseForm(w, r) { + return + } + + err = h.Auth.ChangePassword(userID(r), r.FormValue("current_password"), r.FormValue("new_password")) + if errors.Is(err, auth.ErrInvalidCredentials) { + templates.Account(user, i18n.T(r.Context(), "account.wrong_current_password"), "").Render(r.Context(), w) + return + } + if err != nil { + templates.Account(user, i18n.T(r.Context(), "login.something_went_wrong"), "").Render(r.Context(), w) + return + } + + templates.Account(user, "", i18n.T(r.Context(), "account.password_changed")).Render(r.Context(), w) } diff --git a/internal/handlers/auth.go b/internal/handlers/auth.go index 4fa9122..198fb62 100644 --- a/internal/handlers/auth.go +++ b/internal/handlers/auth.go @@ -11,55 +11,50 @@ import ( "github.com/invopop/ctxi18n/i18n" ) -func Login(w http.ResponseWriter, r *http.Request) { +func (h *Handlers) loginForm(w http.ResponseWriter, r *http.Request) { dest := redirect.Sanitize(r.URL.Query().Get(redirect.Param), "") templates.Login("", dest).Render(r.Context(), w) } -func LoginSubmit(svc *auth.Service) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } - - dest := redirect.Sanitize(r.FormValue(redirect.Param), "/") +func (h *Handlers) loginSubmit(w http.ResponseWriter, r *http.Request) { + if !parseForm(w, r) { + return + } - token, err := svc.Login(r.FormValue("email"), r.FormValue("password")) - if errors.Is(err, auth.ErrInvalidCredentials) { - templates.Login(i18n.T(r.Context(), "login.invalid_credentials"), dest).Render(r.Context(), w) - return - } - if err != nil { - templates.Login(i18n.T(r.Context(), "login.something_went_wrong"), dest).Render(r.Context(), w) - return - } + dest := redirect.Sanitize(r.FormValue(redirect.Param), "/") - http.SetCookie(w, &http.Cookie{ - Name: auth.CookieName, - Value: token, - Path: "/", - HttpOnly: true, - SameSite: http.SameSiteLaxMode, - MaxAge: int((24 * time.Hour).Seconds()), - }) - http.Redirect(w, r, dest, http.StatusSeeOther) + token, err := h.Auth.Login(r.FormValue("email"), r.FormValue("password")) + if errors.Is(err, auth.ErrInvalidCredentials) { + templates.Login(i18n.T(r.Context(), "login.invalid_credentials"), dest).Render(r.Context(), w) + return + } + if err != nil { + templates.Login(i18n.T(r.Context(), "login.something_went_wrong"), dest).Render(r.Context(), w) + return } -} -func Logout(svc *auth.Service) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - if cookie, err := r.Cookie(auth.CookieName); err == nil { - svc.Logout(cookie.Value) - } + http.SetCookie(w, &http.Cookie{ + Name: auth.CookieName, + Value: token, + Path: "/", + HttpOnly: true, + SameSite: http.SameSiteLaxMode, + MaxAge: int((24 * time.Hour).Seconds()), + }) + http.Redirect(w, r, dest, http.StatusSeeOther) +} - http.SetCookie(w, &http.Cookie{ - Name: auth.CookieName, - Value: "", - Path: "/", - HttpOnly: true, - MaxAge: -1, - }) - http.Redirect(w, r, "/login", http.StatusSeeOther) +func (h *Handlers) logout(w http.ResponseWriter, r *http.Request) { + if cookie, err := r.Cookie(auth.CookieName); err == nil { + h.Auth.Logout(cookie.Value) } + + http.SetCookie(w, &http.Cookie{ + Name: auth.CookieName, + Value: "", + Path: "/", + HttpOnly: true, + MaxAge: -1, + }) + http.Redirect(w, r, "/login", http.StatusSeeOther) } diff --git a/internal/handlers/clients.go b/internal/handlers/clients.go index 35f9149..cd7a449 100644 --- a/internal/handlers/clients.go +++ b/internal/handlers/clients.go @@ -7,36 +7,41 @@ import ( "github.com/bloomyindev/time-tracker/internal/db" "github.com/bloomyindev/time-tracker/internal/models" - "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/bloomyindev/time-tracker/internal/templates" + "github.com/go-chi/chi/v5" ) -func ListClients(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - clients, err := db.ListClientsOrderedByName(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - templates.Clients(clients).Render(r.Context(), w) +func (h *Handlers) ClientsRouter() chi.Router { + r := chi.NewRouter() + r.Get("/", h.listClients) + r.Post("/", h.createClient) + r.Get("/{id}", h.clientDetail) + r.Get("/{id}/report", h.clientReport) + r.Get("/{id}/edit", h.editClientForm) + r.Post("/{id}/edit", h.updateClient) + r.Post("/{id}/delete", h.deleteClient) + return r +} + +func (h *Handlers) listClients(w http.ResponseWriter, r *http.Request) { + clients, err := db.ListClientsOrderedByName(h.DB, userID(r)) + if err != nil { + fail(w, err) + return } + templates.Clients(clients).Render(r.Context(), w) } -func CreateClient(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } +func (h *Handlers) createClient(w http.ResponseWriter, r *http.Request) { + if !parseForm(w, r) { + return + } - if _, err := db.CreateClient(conn, userID, r.FormValue("name")); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/clients", http.StatusSeeOther) + if _, err := db.CreateClient(h.DB, userID(r), r.FormValue("name")); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/clients", http.StatusSeeOther) } // taskTypeChoices lists the user's task types, each flagged with whether @@ -62,59 +67,50 @@ func taskTypeChoices(conn *sql.DB, userID, clientID int64) ([]templates.TaskType return choices, nil } -// EditClientForm renders the single page that edits everything about a +// editClientForm renders the single page that edits everything about a // client: its name, its archived state and its allowed task types. -func EditClientForm(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - client, err := db.GetClient(conn, userID, id) - if err != nil { - http.Error(w, "client not found", http.StatusNotFound) - return - } - choices, err := taskTypeChoices(conn, userID, id) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - templates.EditClient(client, choices).Render(r.Context(), w) +func (h *Handlers) editClientForm(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } + client, err := db.GetClient(h.DB, userID(r), id) + if err != nil { + http.Error(w, "client not found", http.StatusNotFound) + return + } + choices, err := taskTypeChoices(h.DB, userID(r), id) + if err != nil { + fail(w, err) + return } + templates.EditClient(client, choices).Render(r.Context(), w) } -// UpdateClient saves the whole edit form: name, archived flag and the +// updateClient saves the whole edit form: name, archived flag and the // client's allowed task types. -func UpdateClient(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - if _, err := db.GetClient(conn, userID, id); err != nil { - http.Error(w, "client not found", http.StatusNotFound) - return - } - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } +func (h *Handlers) updateClient(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } + if _, err := db.GetClient(h.DB, userID(r), id); err != nil { + http.Error(w, "client not found", http.StatusNotFound) + return + } + if !parseForm(w, r) { + return + } - if err := db.UpdateClient(conn, userID, id, r.FormValue("name"), r.FormValue("is_archived") == "1"); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - if err := syncClientTaskTypes(conn, userID, id, r.Form["task_type_id"]); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/clients", http.StatusSeeOther) + if err := db.UpdateClient(h.DB, userID(r), id, r.FormValue("name"), r.FormValue("is_archived") == "1"); err != nil { + fail(w, err) + return + } + if err := syncClientTaskTypes(h.DB, userID(r), id, r.Form["task_type_id"]); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/clients", http.StatusSeeOther) } // syncClientTaskTypes assigns exactly the checked task types to the @@ -146,81 +142,67 @@ func syncClientTaskTypes(conn *sql.DB, userID, clientID int64, checkedIDs []stri return nil } -func DeleteClient(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } +func (h *Handlers) deleteClient(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } - if err := db.DeleteClient(conn, userID, id); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/clients", http.StatusSeeOther) + if err := db.DeleteClient(h.DB, userID(r), id); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/clients", http.StatusSeeOther) } -func ClientDetail(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - - client, err := db.GetClient(conn, userID, id) - if err != nil { - http.Error(w, "client not found", http.StatusNotFound) - return - } +func (h *Handlers) clientDetail(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } - allTypes, err := db.ListTaskTypes(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } + client, err := db.GetClient(h.DB, userID(r), id) + if err != nil { + http.Error(w, "client not found", http.StatusNotFound) + return + } - periods, err := db.ListPeriods(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } + allTypes, err := db.ListTaskTypes(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } - var selectedPeriodID int64 - if raw := r.URL.Query().Get("period_id"); raw != "" { - selectedPeriodID, err = strconv.ParseInt(raw, 10, 64) - if err != nil { - http.Error(w, "invalid period_id", http.StatusBadRequest) - return - } - } - var selectedTaskTypeID int64 - if raw := r.URL.Query().Get("task_type_id"); raw != "" { - selectedTaskTypeID, err = strconv.ParseInt(raw, 10, 64) - if err != nil { - http.Error(w, "invalid task_type_id", http.StatusBadRequest) - return - } - } + periods, err := db.ListPeriods(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } - tasks, err := db.ListTasksByClientFiltered(conn, userID, id, selectedPeriodID, selectedTaskTypeID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - var totalHours float64 - hoursByType := make(map[int64]float64) - for _, t := range tasks { - totalHours += t.HoursSpent - hoursByType[t.TaskTypeID] += t.HoursSpent - } + selectedPeriodID, err := clientFilterID(r, "period_id") + if err != nil { + http.Error(w, "invalid period_id", http.StatusBadRequest) + return + } + selectedTaskTypeID, err := clientFilterID(r, "task_type_id") + if err != nil { + http.Error(w, "invalid task_type_id", http.StatusBadRequest) + return + } - templates.ClientDetail(client, tasks, totalHours, hoursByType, allTypes, periods, selectedPeriodID, selectedTaskTypeID).Render(r.Context(), w) + tasks, err := db.ListTasksByClientFiltered(h.DB, userID(r), id, selectedPeriodID, selectedTaskTypeID) + if err != nil { + fail(w, err) + return } + var totalHours float64 + hoursByType := make(map[int64]float64) + for _, t := range tasks { + totalHours += t.HoursSpent + hoursByType[t.TaskTypeID] += t.HoursSpent + } + + templates.ClientDetail(client, tasks, totalHours, hoursByType, allTypes, periods, selectedPeriodID, selectedTaskTypeID).Render(r.Context(), w) } // clientFilterID reads an optional int64 query param; a blank value means @@ -233,92 +215,88 @@ func clientFilterID(r *http.Request, key string) (int64, error) { return strconv.ParseInt(raw, 10, 64) } -// ClientReport renders a print-friendly page for a client: the total hours on +// clientReport renders a print-friendly page for a client: the total hours on // top, then one table per task type ("project"), honoring the active // period/task-type filters. -func ClientReport(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - client, err := db.GetClient(conn, userID, id) - if err != nil { - http.Error(w, "client not found", http.StatusNotFound) - return - } +func (h *Handlers) clientReport(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } + client, err := db.GetClient(h.DB, userID(r), id) + if err != nil { + http.Error(w, "client not found", http.StatusNotFound) + return + } - periodID, err := clientFilterID(r, "period_id") - if err != nil { - http.Error(w, "invalid period_id", http.StatusBadRequest) - return - } - taskTypeID, err := clientFilterID(r, "task_type_id") - if err != nil { - http.Error(w, "invalid task_type_id", http.StatusBadRequest) - return - } + periodID, err := clientFilterID(r, "period_id") + if err != nil { + http.Error(w, "invalid period_id", http.StatusBadRequest) + return + } + taskTypeID, err := clientFilterID(r, "task_type_id") + if err != nil { + http.Error(w, "invalid task_type_id", http.StatusBadRequest) + return + } - tasks, err := db.ListTasksByClientFiltered(conn, userID, id, periodID, taskTypeID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - allTypes, err := db.ListTaskTypes(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - periods, err := db.ListPeriods(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } + tasks, err := db.ListTasksByClientFiltered(h.DB, userID(r), id, periodID, taskTypeID) + if err != nil { + fail(w, err) + return + } + allTypes, err := db.ListTaskTypes(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + periods, err := db.ListPeriods(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } - // One table per task type, in the app's task-type order, keeping - // only types that actually have tasks in the filtered set. - var total float64 - tasksByType := make(map[int64][]models.Task) - for _, t := range tasks { - total += t.HoursSpent - tasksByType[t.TaskTypeID] = append(tasksByType[t.TaskTypeID], t) + // One table per task type, in the app's task-type order, keeping + // only types that actually have tasks in the filtered set. + var total float64 + tasksByType := make(map[int64][]models.Task) + for _, t := range tasks { + total += t.HoursSpent + tasksByType[t.TaskTypeID] = append(tasksByType[t.TaskTypeID], t) + } + var groups []templates.ClientTypeGroup + for _, tt := range allTypes { + ts, ok := tasksByType[tt.ID] + if !ok { + continue } - var groups []templates.ClientTypeGroup - for _, tt := range allTypes { - ts, ok := tasksByType[tt.ID] - if !ok { - continue - } - var h float64 - for _, t := range ts { - h += t.HoursSpent - } - groups = append(groups, templates.ClientTypeGroup{Name: tt.Name, Tasks: ts, Hours: h}) + var hrs float64 + for _, t := range ts { + hrs += t.HoursSpent } + groups = append(groups, templates.ClientTypeGroup{Name: tt.Name, Tasks: ts, Hours: hrs}) + } - var periodLabel string - for _, p := range periods { - if p.ID == periodID { - periodLabel = p.Name - } + var periodLabel string + for _, p := range periods { + if p.ID == periodID { + periodLabel = p.Name } - var taskTypeLabel string - for _, tt := range allTypes { - if tt.ID == taskTypeID { - taskTypeLabel = tt.Name - } + } + var taskTypeLabel string + for _, tt := range allTypes { + if tt.ID == taskTypeID { + taskTypeLabel = tt.Name } + } - view := templates.ClientReportView{ - ClientName: client.Name, - PeriodName: periodLabel, - TaskTypeName: taskTypeLabel, - Total: total, - Groups: groups, - Periods: periods, - } - templates.ClientReport(view).Render(r.Context(), w) + view := templates.ClientReportView{ + ClientName: client.Name, + PeriodName: periodLabel, + TaskTypeName: taskTypeLabel, + Total: total, + Groups: groups, + Periods: periods, } + templates.ClientReport(view).Render(r.Context(), w) } diff --git a/internal/handlers/handlers.go b/internal/handlers/handlers.go new file mode 100644 index 0000000..d188f7e --- /dev/null +++ b/internal/handlers/handlers.go @@ -0,0 +1,59 @@ +package handlers + +import ( + "database/sql" + "net/http" + "strconv" + + "github.com/bloomyindev/time-tracker/internal/service/auth" + "github.com/go-chi/chi/v5" +) + +// Handlers holds the dependencies every handler shares. Handlers are +// methods on it, which is what lets each resource file expose its own +// sub-router. +type Handlers struct { + DB *sql.DB + Auth *auth.Service +} + +func New(conn *sql.DB, authSvc *auth.Service) *Handlers { + return &Handlers{DB: conn, Auth: authSvc} +} + +// userID returns the authenticated user for the request. It panics if the +// route was mounted outside the auth group: that is a wiring mistake, and +// failing loudly beats silently querying user 0. +func userID(r *http.Request) int64 { + id, ok := auth.UserIDFromContext(r.Context()) + if !ok { + panic("userID called on a route without RequireAuth") + } + return id +} + +// pathID reads the {id} path param. It writes its own 400 and reports +// false when the value isn't an integer. +func pathID(w http.ResponseWriter, r *http.Request) (int64, bool) { + id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64) + if err != nil { + http.Error(w, "invalid id", http.StatusBadRequest) + return 0, false + } + return id, true +} + +// parseForm parses the request body, writing its own 400 on failure. +func parseForm(w http.ResponseWriter, r *http.Request) bool { + if err := r.ParseForm(); err != nil { + http.Error(w, "bad request", http.StatusBadRequest) + return false + } + return true +} + +// fail writes an internal error; handlers use it for anything the user +// can't act on. +func fail(w http.ResponseWriter, err error) { + http.Error(w, err.Error(), http.StatusInternalServerError) +} diff --git a/internal/handlers/home.go b/internal/handlers/home.go index 37d4a88..cc9f7fa 100644 --- a/internal/handlers/home.go +++ b/internal/handlers/home.go @@ -2,6 +2,6 @@ package handlers import "net/http" -func Home(w http.ResponseWriter, r *http.Request) { +func (h *Handlers) home(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "/tasks", http.StatusSeeOther) } diff --git a/internal/handlers/locale.go b/internal/handlers/locale.go index 1bdf036..0df44da 100644 --- a/internal/handlers/locale.go +++ b/internal/handlers/locale.go @@ -6,10 +6,11 @@ import ( "github.com/bloomyindev/time-tracker/internal/i18n" "github.com/bloomyindev/time-tracker/internal/redirect" + "github.com/go-chi/chi/v5" ) -func SetLocale(w http.ResponseWriter, r *http.Request) { - code := r.PathValue("code") +func (h *Handlers) setLocale(w http.ResponseWriter, r *http.Request) { + code := chi.URLParam(r, "code") valid := false for _, l := range i18n.SupportedLocales { diff --git a/internal/handlers/periods.go b/internal/handlers/periods.go index 3c2610c..2d845ea 100644 --- a/internal/handlers/periods.go +++ b/internal/handlers/periods.go @@ -1,110 +1,95 @@ package handlers import ( - "database/sql" "net/http" - "strconv" "github.com/bloomyindev/time-tracker/internal/db" - "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/bloomyindev/time-tracker/internal/templates" + "github.com/go-chi/chi/v5" ) -func ListPeriods(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - periods, err := db.ListPeriods(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - templates.Periods(periods).Render(r.Context(), w) +func (h *Handlers) PeriodsRouter() chi.Router { + r := chi.NewRouter() + r.Get("/", h.listPeriods) + r.Post("/", h.createPeriod) + r.Post("/{id}/default", h.setDefaultPeriod) + r.Get("/{id}/edit", h.editPeriodForm) + r.Post("/{id}/rename", h.renamePeriod) + r.Post("/{id}/delete", h.deletePeriod) + return r +} + +func (h *Handlers) listPeriods(w http.ResponseWriter, r *http.Request) { + periods, err := db.ListPeriods(h.DB, userID(r)) + if err != nil { + fail(w, err) + return } + templates.Periods(periods).Render(r.Context(), w) } -func CreatePeriod(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } +func (h *Handlers) createPeriod(w http.ResponseWriter, r *http.Request) { + if !parseForm(w, r) { + return + } - if _, err := db.CreatePeriod(conn, userID, r.FormValue("name")); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/periods", http.StatusSeeOther) + if _, err := db.CreatePeriod(h.DB, userID(r), r.FormValue("name")); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/periods", http.StatusSeeOther) } -func SetDefaultPeriod(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } +func (h *Handlers) setDefaultPeriod(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } - if err := db.SetDefaultPeriod(conn, userID, id); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/periods", http.StatusSeeOther) + if err := db.SetDefaultPeriod(h.DB, userID(r), id); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/periods", http.StatusSeeOther) } -func EditPeriodForm(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - period, err := db.GetPeriod(conn, userID, id) - if err != nil { - http.Error(w, "period not found", http.StatusNotFound) - return - } - templates.EditPeriod(period).Render(r.Context(), w) +func (h *Handlers) editPeriodForm(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return } + period, err := db.GetPeriod(h.DB, userID(r), id) + if err != nil { + http.Error(w, "period not found", http.StatusNotFound) + return + } + templates.EditPeriod(period).Render(r.Context(), w) } -func RenamePeriod(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } - if err := db.UpdatePeriod(conn, userID, id, r.FormValue("name")); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/periods", http.StatusSeeOther) +func (h *Handlers) renamePeriod(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } + if !parseForm(w, r) { + return } + if err := db.UpdatePeriod(h.DB, userID(r), id, r.FormValue("name")); err != nil { + fail(w, err) + return + } + http.Redirect(w, r, "/periods", http.StatusSeeOther) } -func DeletePeriod(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } +func (h *Handlers) deletePeriod(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } - if err := db.DeletePeriod(conn, userID, id); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/periods", http.StatusSeeOther) + if err := db.DeletePeriod(h.DB, userID(r), id); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/periods", http.StatusSeeOther) } diff --git a/internal/handlers/redirect_test.go b/internal/handlers/redirect_test.go index 2d5b7e5..59a94c1 100644 --- a/internal/handlers/redirect_test.go +++ b/internal/handlers/redirect_test.go @@ -5,12 +5,17 @@ import ( "net/http/httptest" "strings" "testing" + "testing/fstest" appi18n "github.com/bloomyindev/time-tracker/internal/i18n" "github.com/bloomyindev/time-tracker/internal/redirect" "github.com/invopop/ctxi18n" ) +// testHandlers builds a Handlers with no dependencies, for the routes +// that touch neither the database nor a session. +func testHandlers() *Handlers { return &Handlers{} } + func TestMain(m *testing.M) { if err := appi18n.Load(); err != nil { panic(err) @@ -37,7 +42,7 @@ func localeReq(target string) *http.Request { func TestLoginPageCarriesRedirect(t *testing.T) { w := httptest.NewRecorder() - Login(w, localeReq("/login?redirect=%2Ftasks%3Fclient%3D3")) + testHandlers().loginForm(w, localeReq("/login?redirect=%2Ftasks%3Fclient%3D3")) if body := w.Body.String(); !strings.Contains(body, `name="redirect" value="/tasks?client=3"`) { t.Errorf("hidden redirect field missing, body:\n%s", body) @@ -46,7 +51,7 @@ func TestLoginPageCarriesRedirect(t *testing.T) { func TestLoginPageDropsUnsafeRedirect(t *testing.T) { w := httptest.NewRecorder() - Login(w, localeReq("/login?redirect=https%3A%2F%2Fevil.example")) + testHandlers().loginForm(w, localeReq("/login?redirect=https%3A%2F%2Fevil.example")) if strings.Contains(w.Body.String(), `name="redirect"`) { t.Error("unsafe redirect kept as the login form destination") @@ -55,7 +60,7 @@ func TestLoginPageDropsUnsafeRedirect(t *testing.T) { func TestLocaleLinksReturnToCurrentPage(t *testing.T) { w := httptest.NewRecorder() - Login(w, localeReq("/login?redirect=%2Ftasks")) + testHandlers().loginForm(w, localeReq("/login?redirect=%2Ftasks")) body := w.Body.String() for _, want := range []string{ @@ -82,9 +87,8 @@ func TestSetLocaleRedirect(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { r := httptest.NewRequest(http.MethodGet, "/lang/fr"+tt.query, nil) - r.SetPathValue("code", "fr") w := httptest.NewRecorder() - SetLocale(w, r) + testHandlers().Router(fstest.MapFS{}).ServeHTTP(w, r) if w.Code != http.StatusSeeOther { t.Errorf("status = %d, want %d", w.Code, http.StatusSeeOther) diff --git a/internal/handlers/router.go b/internal/handlers/router.go new file mode 100644 index 0000000..071386c --- /dev/null +++ b/internal/handlers/router.go @@ -0,0 +1,52 @@ +package handlers + +import ( + "io/fs" + "net/http" + + "github.com/bloomyindev/time-tracker/internal/i18n" + "github.com/bloomyindev/time-tracker/internal/redirect" + "github.com/go-chi/chi/v5" + "github.com/go-chi/chi/v5/middleware" +) + +// Router assembles the whole site. It is the only place that knows the +// URL prefixes; a resource file never repeats its own mount path. +func (h *Handlers) Router(static fs.FS) http.Handler { + r := chi.NewRouter() + r.Use(middleware.Recoverer) + // Canonicalise the path before routing, so "/clients/" and + // "/clients//" reach the same handler as "/clients" instead of + // falling through to the catch-all. + r.Use(middleware.CleanPath, middleware.StripSlashes) + r.Use(i18n.Middleware, redirect.Middleware) + + // Assets are read-only: bind GET and HEAD explicitly, because + // http.FileServer ignores the method and would otherwise serve a + // file body for POST, DELETE or TRACE too. + assets := http.StripPrefix("/static/", http.FileServerFS(static)) + r.Method(http.MethodGet, "/static/*", assets) + r.Method(http.MethodHead, "/static/*", assets) + + // Public pages: the landing redirect, the login flow and the locale + // switch. Unknown paths land on the app rather than a bare 404. + r.Get("/", h.home) + r.NotFound(h.home) + r.Get("/login", h.loginForm) + r.Post("/login", h.loginSubmit) + r.Get("/logout", h.logout) + r.Get("/lang/{code}", h.setLocale) + + // Everything below needs a session. + r.Group(func(r chi.Router) { + r.Use(h.Auth.RequireAuth) + r.Mount("/clients", h.ClientsRouter()) + r.Mount("/task-types", h.TaskTypesRouter()) + r.Mount("/periods", h.PeriodsRouter()) + r.Mount("/tasks", h.TasksRouter()) + r.Mount("/time", h.TimeRouter()) + r.Mount("/account", h.AccountRouter()) + }) + + return r +} diff --git a/internal/handlers/router_test.go b/internal/handlers/router_test.go new file mode 100644 index 0000000..2583e9c --- /dev/null +++ b/internal/handlers/router_test.go @@ -0,0 +1,81 @@ +package handlers + +import ( + "net/http" + "net/http/httptest" + "testing" + "testing/fstest" +) + +// staticFS is a stand-in for the embedded asset tree. +var staticFS = fstest.MapFS{"css/style.css": &fstest.MapFile{Data: []byte("body{}")}} + +// TestStaticServesReadMethodsOnly guards the asset route: http.FileServer +// ignores the request method, so the route has to bind GET and HEAD itself +// or a POST/DELETE/TRACE would be answered with the file body. +func TestStaticServesReadMethodsOnly(t *testing.T) { + tests := []struct { + method string + wantCode int + wantBody string + }{ + {http.MethodGet, http.StatusOK, "body{}"}, + {http.MethodHead, http.StatusOK, ""}, + {http.MethodPost, http.StatusMethodNotAllowed, ""}, + {http.MethodDelete, http.StatusMethodNotAllowed, ""}, + {http.MethodPut, http.StatusMethodNotAllowed, ""}, + {http.MethodTrace, http.StatusMethodNotAllowed, ""}, + } + + for _, tt := range tests { + t.Run(tt.method, func(t *testing.T) { + w := httptest.NewRecorder() + r := httptest.NewRequest(tt.method, "/static/css/style.css", nil) + testHandlers().Router(staticFS).ServeHTTP(w, r) + + if w.Code != tt.wantCode { + t.Errorf("status = %d, want %d", w.Code, tt.wantCode) + } + if got := w.Body.String(); got != tt.wantBody { + t.Errorf("body = %q, want %q", got, tt.wantBody) + } + }) + } +} + +// TestTrailingSlashesReachTheSameRoute checks that a stray or doubled +// slash still lands inside the mounted resource router rather than falling +// through to the catch-all. Unauthenticated, the tell is the login +// redirect: the catch-all would send the user to /tasks instead. +func TestTrailingSlashesReachTheSameRoute(t *testing.T) { + paths := []string{ + "/clients", "/clients/", "/clients//", + "/clients/1/edit", "/clients/1/edit/", + "/tasks/", "/time/report/", "/account/", + } + + for _, path := range paths { + t.Run(path, func(t *testing.T) { + w := httptest.NewRecorder() + testHandlers().Router(staticFS).ServeHTTP(w, httptest.NewRequest(http.MethodGet, path, nil)) + + if w.Code != http.StatusSeeOther { + t.Fatalf("status = %d, want %d", w.Code, http.StatusSeeOther) + } + if loc := w.Header().Get("Location"); loc == "/tasks" { + t.Errorf("%s fell through to the catch-all instead of the mounted route", path) + } + }) + } +} + +// TestRootIsNotStripped makes sure canonicalising slashes leaves "/" +// itself alone, so the landing redirect still works. +func TestRootIsNotStripped(t *testing.T) { + w := httptest.NewRecorder() + testHandlers().Router(staticFS).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/", nil)) + + if w.Code != http.StatusSeeOther || w.Header().Get("Location") != "/tasks" { + t.Errorf("got %d -> %q, want 303 -> /tasks", w.Code, w.Header().Get("Location")) + } +} diff --git a/internal/handlers/task_types.go b/internal/handlers/task_types.go index 9f0bbd0..53d8b20 100644 --- a/internal/handlers/task_types.go +++ b/internal/handlers/task_types.go @@ -1,93 +1,81 @@ package handlers import ( - "database/sql" "net/http" - "strconv" "github.com/bloomyindev/time-tracker/internal/db" - "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/bloomyindev/time-tracker/internal/templates" + "github.com/go-chi/chi/v5" ) -func ListTaskTypes(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - types, err := db.ListTaskTypes(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - templates.TaskTypes(types).Render(r.Context(), w) +func (h *Handlers) TaskTypesRouter() chi.Router { + r := chi.NewRouter() + r.Get("/", h.listTaskTypes) + r.Post("/", h.createTaskType) + r.Get("/{id}/edit", h.editTaskTypeForm) + r.Post("/{id}/rename", h.renameTaskType) + r.Post("/{id}/delete", h.deleteTaskType) + return r +} + +func (h *Handlers) listTaskTypes(w http.ResponseWriter, r *http.Request) { + types, err := db.ListTaskTypes(h.DB, userID(r)) + if err != nil { + fail(w, err) + return } + templates.TaskTypes(types).Render(r.Context(), w) } -func CreateTaskType(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } +func (h *Handlers) createTaskType(w http.ResponseWriter, r *http.Request) { + if !parseForm(w, r) { + return + } - if _, err := db.CreateTaskType(conn, userID, r.FormValue("name")); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/task-types", http.StatusSeeOther) + if _, err := db.CreateTaskType(h.DB, userID(r), r.FormValue("name")); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/task-types", http.StatusSeeOther) } -func EditTaskTypeForm(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - taskType, err := db.GetTaskType(conn, userID, id) - if err != nil { - http.Error(w, "task type not found", http.StatusNotFound) - return - } - templates.EditTaskType(taskType).Render(r.Context(), w) +func (h *Handlers) editTaskTypeForm(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } + taskType, err := db.GetTaskType(h.DB, userID(r), id) + if err != nil { + http.Error(w, "task type not found", http.StatusNotFound) + return } + templates.EditTaskType(taskType).Render(r.Context(), w) } -func RenameTaskType(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } - if err := db.UpdateTaskType(conn, userID, id, r.FormValue("name")); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/task-types", http.StatusSeeOther) +func (h *Handlers) renameTaskType(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return } + if !parseForm(w, r) { + return + } + if err := db.UpdateTaskType(h.DB, userID(r), id, r.FormValue("name")); err != nil { + fail(w, err) + return + } + http.Redirect(w, r, "/task-types", http.StatusSeeOther) } -func DeleteTaskType(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } +func (h *Handlers) deleteTaskType(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } - if err := db.DeleteTaskType(conn, userID, id); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/task-types", http.StatusSeeOther) + if err := db.DeleteTaskType(h.DB, userID(r), id); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/task-types", http.StatusSeeOther) } diff --git a/internal/handlers/tasks.go b/internal/handlers/tasks.go index 6e5a945..541239a 100644 --- a/internal/handlers/tasks.go +++ b/internal/handlers/tasks.go @@ -8,47 +8,53 @@ import ( "github.com/bloomyindev/time-tracker/internal/db" "github.com/bloomyindev/time-tracker/internal/models" - "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/bloomyindev/time-tracker/internal/templates" + "github.com/go-chi/chi/v5" ) -func ListTasks(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - - tasks, err := db.ListTasks(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - clients, err := db.ListClientsOrderedByName(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - types, err := db.ListTaskTypes(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - periods, err := db.ListPeriods(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - byClient, err := db.ListTaskTypesByClient(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } +func (h *Handlers) TasksRouter() chi.Router { + r := chi.NewRouter() + r.Get("/", h.listTasks) + r.Post("/", h.createTask) + r.Get("/{id}/edit", h.editTaskForm) + r.Post("/{id}/update", h.updateTask) + r.Post("/{id}/delete", h.deleteTask) + return r +} - var defaultPeriodID int64 - if p, err := db.GetDefaultPeriod(conn, userID); err == nil { - defaultPeriodID = p.ID - } +func (h *Handlers) listTasks(w http.ResponseWriter, r *http.Request) { + tasks, err := db.ListTasks(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + clients, err := db.ListClientsOrderedByName(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + types, err := db.ListTaskTypes(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + periods, err := db.ListPeriods(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + byClient, err := db.ListTaskTypesByClient(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } - templates.Tasks(clients, types, periods, byClient, groupByDay(tasks), time.Now().Format("2006-01-02"), defaultPeriodID).Render(r.Context(), w) + var defaultPeriodID int64 + if p, err := db.GetDefaultPeriod(h.DB, userID(r)); err == nil { + defaultPeriodID = p.ID } + + templates.Tasks(clients, types, periods, byClient, groupByDay(tasks), time.Now().Format("2006-01-02"), defaultPeriodID).Render(r.Context(), w) } // parsePeriodID reads an optional period_id form value; a blank or @@ -113,214 +119,190 @@ func taskTypeAllowedForClient(conn *sql.DB, clientID, taskTypeID int64) (bool, e return false, nil } -func CreateTask(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } - - clientID, err := strconv.ParseInt(r.FormValue("client_id"), 10, 64) - if err != nil { - http.Error(w, "invalid client_id", http.StatusBadRequest) - return - } - taskTypeID, err := strconv.ParseInt(r.FormValue("task_type_id"), 10, 64) - if err != nil { - http.Error(w, "invalid task_type_id", http.StatusBadRequest) - return - } - hoursSpent, err := strconv.ParseFloat(r.FormValue("hours_spent"), 64) - if err != nil { - http.Error(w, "invalid hours_spent", http.StatusBadRequest) - return - } - date, err := time.Parse("2006-01-02", r.FormValue("date")) - if err != nil { - http.Error(w, "invalid date", http.StatusBadRequest) - return - } - periodID, err := parsePeriodID(r) - if err != nil { - http.Error(w, "invalid period_id", http.StatusBadRequest) - return - } - - active, err := clientAcceptsTasks(conn, userID, clientID) - if err != nil { - http.Error(w, "client not found", http.StatusBadRequest) - return - } - if !active { - http.Error(w, "client is archived", http.StatusBadRequest) - return - } +// taskForm holds the fields shared by the create and update forms. +type taskForm struct { + clientID int64 + taskTypeID int64 + periodID int64 + title string + hoursSpent float64 + date time.Time +} - ok, err := taskTypeAllowedForClient(conn, clientID, taskTypeID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - if !ok { - http.Error(w, "task type not allowed for this client", http.StatusBadRequest) - return - } +// parseTaskForm reads and validates the task form, writing its own 400 +// and reporting false on the first bad field. +func parseTaskForm(w http.ResponseWriter, r *http.Request) (taskForm, bool) { + if !parseForm(w, r) { + return taskForm{}, false + } - _, err = db.CreateTask(conn, models.Task{ - UserID: userID, - ClientID: clientID, - TaskTypeID: taskTypeID, - PeriodID: periodID, - Title: r.FormValue("title"), - HoursSpent: hoursSpent, - Date: date, - }) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/tasks", http.StatusSeeOther) + var f taskForm + var err error + if f.clientID, err = strconv.ParseInt(r.FormValue("client_id"), 10, 64); err != nil { + http.Error(w, "invalid client_id", http.StatusBadRequest) + return taskForm{}, false + } + if f.taskTypeID, err = strconv.ParseInt(r.FormValue("task_type_id"), 10, 64); err != nil { + http.Error(w, "invalid task_type_id", http.StatusBadRequest) + return taskForm{}, false } + if f.hoursSpent, err = strconv.ParseFloat(r.FormValue("hours_spent"), 64); err != nil { + http.Error(w, "invalid hours_spent", http.StatusBadRequest) + return taskForm{}, false + } + if f.date, err = time.Parse("2006-01-02", r.FormValue("date")); err != nil { + http.Error(w, "invalid date", http.StatusBadRequest) + return taskForm{}, false + } + if f.periodID, err = parsePeriodID(r); err != nil { + http.Error(w, "invalid period_id", http.StatusBadRequest) + return taskForm{}, false + } + f.title = r.FormValue("title") + return f, true } -func EditTaskForm(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } +func (h *Handlers) createTask(w http.ResponseWriter, r *http.Request) { + f, ok := parseTaskForm(w, r) + if !ok { + return + } - task, err := db.GetTask(conn, userID, id) - if err != nil { - http.Error(w, "task not found", http.StatusNotFound) - return - } - clients, err := db.ListClients(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - types, err := db.ListTaskTypes(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - periods, err := db.ListPeriods(conn, userID) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } + active, err := clientAcceptsTasks(h.DB, userID(r), f.clientID) + if err != nil { + http.Error(w, "client not found", http.StatusBadRequest) + return + } + if !active { + http.Error(w, "client is archived", http.StatusBadRequest) + return + } - templates.EditTask(task, clients, types, periods).Render(r.Context(), w) + allowed, err := taskTypeAllowedForClient(h.DB, f.clientID, f.taskTypeID) + if err != nil { + fail(w, err) + return + } + if !allowed { + http.Error(w, "task type not allowed for this client", http.StatusBadRequest) + return } + + _, err = db.CreateTask(h.DB, models.Task{ + UserID: userID(r), + ClientID: f.clientID, + TaskTypeID: f.taskTypeID, + PeriodID: f.periodID, + Title: f.title, + HoursSpent: f.hoursSpent, + Date: f.date, + }) + if err != nil { + fail(w, err) + return + } + http.Redirect(w, r, "/tasks", http.StatusSeeOther) } -func UpdateTask(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } - existing, err := db.GetTask(conn, userID, id) - if err != nil { - http.Error(w, "task not found", http.StatusNotFound) - return - } +func (h *Handlers) editTaskForm(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } - if err := r.ParseForm(); err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } + task, err := db.GetTask(h.DB, userID(r), id) + if err != nil { + http.Error(w, "task not found", http.StatusNotFound) + return + } + clients, err := db.ListClients(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + types, err := db.ListTaskTypes(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } + periods, err := db.ListPeriods(h.DB, userID(r)) + if err != nil { + fail(w, err) + return + } - clientID, err := strconv.ParseInt(r.FormValue("client_id"), 10, 64) - if err != nil { - http.Error(w, "invalid client_id", http.StatusBadRequest) - return - } - taskTypeID, err := strconv.ParseInt(r.FormValue("task_type_id"), 10, 64) - if err != nil { - http.Error(w, "invalid task_type_id", http.StatusBadRequest) - return - } - hoursSpent, err := strconv.ParseFloat(r.FormValue("hours_spent"), 64) - if err != nil { - http.Error(w, "invalid hours_spent", http.StatusBadRequest) - return - } - date, err := time.Parse("2006-01-02", r.FormValue("date")) - if err != nil { - http.Error(w, "invalid date", http.StatusBadRequest) - return - } - periodID, err := parsePeriodID(r) - if err != nil { - http.Error(w, "invalid period_id", http.StatusBadRequest) - return - } + templates.EditTask(task, clients, types, periods).Render(r.Context(), w) +} - // Moving a task onto an archived client is a new assignment, so - // it's refused; a task already on an archived client stays - // editable. - if clientID != existing.ClientID { - active, err := clientAcceptsTasks(conn, userID, clientID) - if err != nil { - http.Error(w, "client not found", http.StatusBadRequest) - return - } - if !active { - http.Error(w, "client is archived", http.StatusBadRequest) - return - } - } +func (h *Handlers) updateTask(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } + existing, err := db.GetTask(h.DB, userID(r), id) + if err != nil { + http.Error(w, "task not found", http.StatusNotFound) + return + } + + f, ok := parseTaskForm(w, r) + if !ok { + return + } - ok, err := taskTypeAllowedForClient(conn, clientID, taskTypeID) + // Moving a task onto an archived client is a new assignment, so + // it's refused; a task already on an archived client stays + // editable. + if f.clientID != existing.ClientID { + active, err := clientAcceptsTasks(h.DB, userID(r), f.clientID) if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) + http.Error(w, "client not found", http.StatusBadRequest) return } - if !ok { - http.Error(w, "task type not allowed for this client", http.StatusBadRequest) + if !active { + http.Error(w, "client is archived", http.StatusBadRequest) return } + } - err = db.UpdateTask(conn, models.Task{ - ID: id, - UserID: userID, - ClientID: clientID, - TaskTypeID: taskTypeID, - PeriodID: periodID, - Title: r.FormValue("title"), - HoursSpent: hoursSpent, - Date: date, - }) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - // Anchor the reload on the edited row so the browser restores the - // scroll position instead of jumping to the top of the list. - http.Redirect(w, r, "/tasks#task-"+strconv.FormatInt(id, 10), http.StatusSeeOther) + allowed, err := taskTypeAllowedForClient(h.DB, f.clientID, f.taskTypeID) + if err != nil { + fail(w, err) + return + } + if !allowed { + http.Error(w, "task type not allowed for this client", http.StatusBadRequest) + return } + + err = db.UpdateTask(h.DB, models.Task{ + ID: id, + UserID: userID(r), + ClientID: f.clientID, + TaskTypeID: f.taskTypeID, + PeriodID: f.periodID, + Title: f.title, + HoursSpent: f.hoursSpent, + Date: f.date, + }) + if err != nil { + fail(w, err) + return + } + // Anchor the reload on the edited row so the browser restores the + // scroll position instead of jumping to the top of the list. + http.Redirect(w, r, "/tasks#task-"+strconv.FormatInt(id, 10), http.StatusSeeOther) } -func DeleteTask(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) - if err != nil { - http.Error(w, "invalid id", http.StatusBadRequest) - return - } +func (h *Handlers) deleteTask(w http.ResponseWriter, r *http.Request) { + id, ok := pathID(w, r) + if !ok { + return + } - if err := db.DeleteTask(conn, userID, id); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - http.Redirect(w, r, "/tasks", http.StatusSeeOther) + if err := db.DeleteTask(h.DB, userID(r), id); err != nil { + fail(w, err) + return } + http.Redirect(w, r, "/tasks", http.StatusSeeOther) } diff --git a/internal/handlers/time_management.go b/internal/handlers/time_management.go index 219fcb0..7a95fb9 100644 --- a/internal/handlers/time_management.go +++ b/internal/handlers/time_management.go @@ -6,10 +6,17 @@ import ( "time" "github.com/bloomyindev/time-tracker/internal/db" - "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/bloomyindev/time-tracker/internal/templates" + "github.com/go-chi/chi/v5" ) +func (h *Handlers) TimeRouter() chi.Router { + r := chi.NewRouter() + r.Get("/", h.listTimeEntries) + r.Get("/report", h.timeReport) + return r +} + // weekdayIndex maps a date to models.User.DailyHours order (0 = Monday // .. 6 = Sunday). Go's time.Weekday has Sunday = 0, so shift by 6. func weekdayIndex(t time.Time) int { @@ -105,33 +112,27 @@ func buildTimeView(conn *sql.DB, userID int64, from, to string) (templates.TimeV return view, nil } -// ListTimeEntries shows one row per day with the total hours logged that +// listTimeEntries shows one row per day with the total hours logged that // day, grouped by month, plus an optional date-to-date range filter. -func ListTimeEntries(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - view, err := buildTimeView(conn, userID, r.URL.Query().Get("from"), r.URL.Query().Get("to")) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - templates.TimeEntries(view).Render(r.Context(), w) +func (h *Handlers) listTimeEntries(w http.ResponseWriter, r *http.Request) { + view, err := buildTimeView(h.DB, userID(r), r.URL.Query().Get("from"), r.URL.Query().Get("to")) + if err != nil { + fail(w, err) + return } + templates.TimeEntries(view).Render(r.Context(), w) } -// TimeReport renders a print-friendly, standalone page of the same breakdown +// timeReport renders a print-friendly, standalone page of the same breakdown // (no navbar) so the browser's print dialog can save it as a PDF. -func TimeReport(conn *sql.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, _ := auth.UserIDFromContext(r.Context()) - view, err := buildTimeView(conn, userID, r.URL.Query().Get("from"), r.URL.Query().Get("to")) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - // By default the report lists only days off target; all=1 keeps - // the days that hit their target exactly. - includeOnTarget := r.URL.Query().Get("all") == "1" - templates.TimeReport(view, includeOnTarget).Render(r.Context(), w) +func (h *Handlers) timeReport(w http.ResponseWriter, r *http.Request) { + view, err := buildTimeView(h.DB, userID(r), r.URL.Query().Get("from"), r.URL.Query().Get("to")) + if err != nil { + fail(w, err) + return } + // By default the report lists only days off target; all=1 keeps + // the days that hit their target exactly. + includeOnTarget := r.URL.Query().Get("all") == "1" + templates.TimeReport(view, includeOnTarget).Render(r.Context(), w) } From e63c8037ceb60c2b3866fc8d9741c0730273f283 Mon Sep 17 00:00:00 2001 From: Bloomy Date: Thu, 30 Jul 2026 17:55:11 +0200 Subject: [PATCH 2/3] feat: enforce foreign keys and secure cookies Two data-integrity problems and one transport one, plus the migration machinery the first of them needed. Foreign keys were never actually on. The pragma was issued as a statement after connecting, which configures only the connection that happened to serve it, and database/sql pools connections freely; every other one ran with enforcement off. Deleting a client therefore left its tasks behind, pointing at a row that no longer existed, where they rendered with a blank name and still counted toward every hours total. The pragmas move into the DSN, so no connection can miss them. Turning enforcement on changes what a delete means. Tasks are the records this app exists to keep, so their foreign keys stay RESTRICT and the refusal surfaces as db.ErrInUse, which the handlers report as a 409 rather than a 500 full of SQL. The client/task-type join table is bookkeeping instead, so it cascades; that needed a table rebuild, since SQLite can't ALTER a constraint. Periods sit between the two: they only group hours, so deleting one detaches its tasks in a transaction rather than refusing. Expressing that rebuild meant replacing the old startup schema block and its list of ALTERs, which could only ever add columns. Migrations are now embedded .sql files applied in order and recorded in schema_migrations. Migration 0001 is the schema as it already stands, written idempotently so existing databases are recognised as migrated instead of rebuilt. Session cookies were never marked Secure, so the token travelled in cleartext over any plain-HTTP hop. It defaults on now, since the app is normally reached through a TLS-terminating proxy, and a deployment that never considered it should land on the safe setting; plain-HTTP setups opt out with TRACKER_SECURE_COOKIES=false, which compose.yml does. Login and logout share one helper so their attributes can't drift apart. Environment variables gain a TRACKER_ prefix. This is a breaking change for existing deployments: DB_PATH becomes TRACKER_DB_PATH and JWT_SECRET becomes TRACKER_JWT_SECRET. The bundled Dockerfile and compose.yml are updated; anything overriding them needs the same rename. WAL stays opt-in behind TRACKER_SQLITE_WAL, since it leaves -wal and -shm files next to the database. Capping the pool at one connection, which SQLite's single-writer rule wants anyway, removes the contention that would otherwise be the reason to reach for it. Verified against a copy of a live database: 927 tasks and their 1063.5 hours preserved, all 52 assignments carried across the rebuild, both cascade clauses in place, no foreign key violations. Co-Authored-By: Claude Opus 5 --- Dockerfile | 2 +- README.md | 29 +- cmd/time-tracker/main.go | 11 +- compose.yml | 11 +- internal/config/config.go | 43 ++- internal/db/clients.go | 6 + internal/db/db.go | 163 ++++++----- internal/db/db_test.go | 274 +++++++++++++++++++ internal/db/periods.go | 21 +- internal/db/task_types.go | 5 + internal/handlers/auth.go | 23 +- internal/handlers/auth_test.go | 117 ++++++++ internal/handlers/clients.go | 8 +- internal/handlers/handlers.go | 10 +- internal/handlers/locale.go | 1 + internal/handlers/task_types.go | 8 +- migrations/0001_init.sql | 56 ++++ migrations/0002_task_type_client_cascade.sql | 25 ++ migrations/migrations.go | 9 + 19 files changed, 715 insertions(+), 107 deletions(-) create mode 100644 internal/db/db_test.go create mode 100644 internal/handlers/auth_test.go create mode 100644 migrations/0001_init.sql create mode 100644 migrations/0002_task_type_client_cascade.sql create mode 100644 migrations/migrations.go diff --git a/Dockerfile b/Dockerfile index 8260544..02c203b 100644 --- a/Dockerfile +++ b/Dockerfile @@ -23,7 +23,7 @@ WORKDIR /app COPY --from=build /out/time-tracker /app/ EXPOSE 8080 -ENV DB_PATH=/data/time-tracker.db +ENV TRACKER_DB_PATH=/data/time-tracker.db VOLUME ["/data"] ENTRYPOINT ["/app/time-tracker"] diff --git a/README.md b/README.md index e05a34a..2487787 100644 --- a/README.md +++ b/README.md @@ -46,13 +46,14 @@ version tags, for `linux/amd64`, `linux/arm64`, `linux/arm/v7`, `linux/arm/v6`, ```sh docker run -d -p 8080:8080 \ - -e JWT_SECRET=change-me \ + -e TRACKER_JWT_SECRET=change-me \ + -e TRACKER_SECURE_COOKIES=false \ -v time-tracker-data:/data \ ghcr.io/bloomyindev/time-tracker:latest ``` The container runs `time-tracker serve` by default. The database is stored in -the `/data` volume (`DB_PATH=/data/time-tracker.db`). Run admin commands inside +the `/data` volume (`TRACKER_DB_PATH=/data/time-tracker.db`). Run admin commands inside the container with `docker exec /app/time-tracker ` (see [CLI](#cli)). @@ -82,7 +83,7 @@ archive contains the single `time-tracker` binary. Prebuilt targets: ```sh tar -xzf time-tracker_*_linux_amd64.tar.gz cd time-tracker_*_linux_amd64 -JWT_SECRET=change-me ./time-tracker serve +TRACKER_JWT_SECRET=change-me TRACKER_SECURE_COOKIES=false ./time-tracker serve ``` ### From source @@ -93,7 +94,7 @@ Building from source requires Go 1.26 or later. See ## First run 1. **Start the server.** It listens on . In production, - always set `JWT_SECRET` (see [Configuration](#configuration)). + always set `TRACKER_JWT_SECRET` (see [Configuration](#configuration)). 2. **Create a user.** There is no public sign-up: ```sh ./time-tracker register --email you@example.com --password secret @@ -106,13 +107,21 @@ Building from source requires Go 1.26 or later. See The app is configured through environment variables: -| Variable | Default | Description | -|--------------|------------------------|---------------------------------------------------------------------------------------------| -| `DB_PATH` | `time-tracker.db` | Path to the SQLite database file. | -| `JWT_SECRET` | `dev-secret-change-me` | Secret for signing JWTs used by the (currently unused) bearer-token API flow. Set this in production. | +| Variable | Default | Description | +|---------------------------|------------------------|------------------------------------------------------------------------------------------------------| +| `TRACKER_DB_PATH` | `time-tracker.db` | Path to the SQLite database file. | +| `TRACKER_JWT_SECRET` | `dev-secret-change-me` | Secret for signing JWTs used by the (currently unused) bearer-token API flow. Set this in production. | +| `TRACKER_SECURE_COOKIES` | `true` | Mark cookies `Secure`, so browsers only send them over HTTPS. Set to `false` to serve plain HTTP. | +| `TRACKER_SQLITE_WAL` | `false` | Enable SQLite write-ahead logging. Adds `-wal` and `-shm` files next to the database. | + +**Serving over plain HTTP?** Set `TRACKER_SECURE_COOKIES=false`. The default +assumes a TLS-terminating reverse proxy in front. Left on over plain HTTP, the +browser accepts the session cookie and then never sends it back, so logging in +appears to do nothing. The server always listens on port `8080`. The database path can also be passed -with `--db-path`. The schema is created and migrated automatically on startup. +with `--db-path`. Migrations are embedded in the binary and applied on startup; +each one is recorded in a `schema_migrations` table so it runs exactly once. ## CLI @@ -124,7 +133,7 @@ time-tracker register --email --password # create a user time-tracker export-users # dump users as JSON (no password hashes) ``` -Every command accepts `--db-path` (or the `DB_PATH` environment variable). +Every command accepts `--db-path` (or the `TRACKER_DB_PATH` environment variable). ## Contributing diff --git a/cmd/time-tracker/main.go b/cmd/time-tracker/main.go index 85c33de..3e8ebe4 100644 --- a/cmd/time-tracker/main.go +++ b/cmd/time-tracker/main.go @@ -34,19 +34,19 @@ func main() { } } -// dbPathFlag overrides the database path (env: DB_PATH). A fresh flag is -// returned per command so each owns its own value. +// dbPathFlag overrides the database path (env: TRACKER_DB_PATH). A fresh flag +// is returned per command so each owns its own value. func dbPathFlag() *cli.StringFlag { return &cli.StringFlag{ Name: "db-path", Usage: "path to the sqlite database file", Value: config.Load().DBPath, - Sources: cli.EnvVars("DB_PATH"), + Sources: cli.EnvVars("TRACKER_DB_PATH"), } } func openDB(cmd *cli.Command) (*sql.DB, error) { - return db.Open(cmd.String("db-path")) + return db.Open(cmd.String("db-path"), db.Options{WAL: config.Load().SQLiteWAL}) } func serveCommand() *cli.Command { @@ -65,13 +65,14 @@ func serveCommand() *cli.Command { if err != nil { return fmt.Errorf("open db: %w", err) } + defer conn.Close() staticFS, err := fs.Sub(assets.Static, "static") if err != nil { return fmt.Errorf("mount static assets: %w", err) } - h := handlers.New(conn, auth.NewService(conn, cfg.JWTSecret)) + h := handlers.New(conn, auth.NewService(conn, cfg.JWTSecret), cfg) log.Printf("listening on port %d", cfg.Port) return http.ListenAndServe(fmt.Sprintf(":%d", cfg.Port), h.Router(staticFS)) diff --git a/compose.yml b/compose.yml index 2f0638a..3eea82b 100644 --- a/compose.yml +++ b/compose.yml @@ -7,7 +7,16 @@ services: - "8080:8080" environment: # Set a strong secret in production. - JWT_SECRET: change-me + TRACKER_JWT_SECRET: change-me + # This file publishes port 8080 directly, so the app is reached over + # plain HTTP and the session cookie must not be marked Secure — the + # browser would accept it and then never send it back, and login would + # appear to do nothing. Drop this line the moment a TLS-terminating + # reverse proxy sits in front. + TRACKER_SECURE_COOKIES: "false" + # Write-ahead logging: readers and writers stop blocking each other, at + # the cost of two extra files next to the database. + # TRACKER_SQLITE_WAL: "true" volumes: - data:/data restart: unless-stopped diff --git a/internal/config/config.go b/internal/config/config.go index cb484ee..6aae4a7 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -1,18 +1,37 @@ +// Package config reads runtime settings from the environment. Every variable +// is namespaced with TRACKER_, so the app can't pick up a generic name another +// process on the same host happens to export. package config -import "os" +import ( + "log" + "os" + "strconv" +) type Config struct { Port int DBPath string JWTSecret string + // SQLiteWAL turns on write-ahead logging. It trades two extra files + // next to the database for readers and writers that no longer block + // each other. + SQLiteWAL bool + // SecureCookies marks cookies Secure, so a browser only ever sends + // them back over HTTPS. It defaults to on: the app is normally reached + // through a TLS-terminating reverse proxy, and a deployment that never + // thought about this should land on the safe setting. Turn it off to + // serve plain HTTP, or the session cookie gets set and never returned. + SecureCookies bool } func Load() Config { return Config{ - Port: 8080, - DBPath: getEnv("DB_PATH", "time-tracker.db"), - JWTSecret: getEnv("JWT_SECRET", "dev-secret-change-me"), + Port: 8080, + DBPath: getEnv("TRACKER_DB_PATH", "time-tracker.db"), + JWTSecret: getEnv("TRACKER_JWT_SECRET", "dev-secret-change-me"), + SQLiteWAL: getBool("TRACKER_SQLITE_WAL", false), + SecureCookies: getBool("TRACKER_SECURE_COOKIES", true), } } @@ -22,3 +41,19 @@ func getEnv(key, fallback string) string { } return fallback } + +// getBool reads a boolean in any form strconv accepts ("1", "true", "off"). +// An unparseable value is a typo in the deployment, not a reason to run with a +// setting nobody chose, so it is reported and the default stands. +func getBool(key string, fallback bool) bool { + raw := os.Getenv(key) + if raw == "" { + return fallback + } + parsed, err := strconv.ParseBool(raw) + if err != nil { + log.Printf("%s: %q isn't a boolean, using %t", key, raw, fallback) + return fallback + } + return parsed +} diff --git a/internal/db/clients.go b/internal/db/clients.go index 4519886..5a51ca5 100644 --- a/internal/db/clients.go +++ b/internal/db/clients.go @@ -64,7 +64,13 @@ func UpdateClient(conn *sql.DB, userID, id int64, name string, archived bool) er return err } +// DeleteClient removes a client. It returns ErrInUse when tasks are still +// logged against it: those hours are the point of the app, so the client has +// to be archived rather than deleted. Its task type assignments cascade away. func DeleteClient(conn *sql.DB, userID, id int64) error { _, err := conn.Exec(`DELETE FROM clients WHERE id = ? AND user_id = ?`, id, userID) + if isForeignKeyErr(err) { + return ErrInUse + } return err } diff --git a/internal/db/db.go b/internal/db/db.go index 114bb8d..92067fa 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -2,88 +2,115 @@ package db import ( "database/sql" - "strings" + "errors" + "fmt" + "io/fs" + "sort" - _ "modernc.org/sqlite" -) - -const schema = ` -CREATE TABLE IF NOT EXISTS users ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - email TEXT NOT NULL UNIQUE, - password_hash TEXT NOT NULL -); + "github.com/bloomyindev/time-tracker/migrations" -CREATE TABLE IF NOT EXISTS clients ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name TEXT NOT NULL, - user_id INTEGER NOT NULL REFERENCES users(id) -); - -CREATE TABLE IF NOT EXISTS task_types ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - user_id INTEGER NOT NULL REFERENCES users(id), - name TEXT NOT NULL -); + sqlite "modernc.org/sqlite" +) -CREATE TABLE IF NOT EXISTS task_types_for_client ( - client_id INTEGER NOT NULL REFERENCES clients(id), - task_type_id INTEGER NOT NULL REFERENCES task_types(id), - PRIMARY KEY (client_id, task_type_id) -); +// ErrInUse is returned when a row can't be deleted because another row still +// references it (foreign key constraint). +var ErrInUse = errors.New("resource is still referenced by other records") -CREATE TABLE IF NOT EXISTS periods ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - user_id INTEGER NOT NULL REFERENCES users(id), - name TEXT NOT NULL, - is_default INTEGER NOT NULL DEFAULT 0 -); +// sqliteConstraintForeignKey is SQLite's extended result code +// SQLITE_CONSTRAINT_FOREIGNKEY, returned when a delete/update violates a +// foreign key constraint. +const sqliteConstraintForeignKey = 787 -CREATE TABLE IF NOT EXISTS tasks ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - user_id INTEGER NOT NULL REFERENCES users(id), - client_id INTEGER NOT NULL REFERENCES clients(id), - task_type_id INTEGER NOT NULL REFERENCES task_types(id), - title TEXT NOT NULL, - hours_spent DOUBLE NOT NULL, - date DATE NOT NULL -); -` +// isForeignKeyErr reports whether err is a foreign-key constraint violation. +func isForeignKeyErr(err error) bool { + var se *sqlite.Error + return errors.As(err, &se) && se.Code() == sqliteConstraintForeignKey +} -// migrations holds schema changes that ALTER existing tables, which -// CREATE TABLE IF NOT EXISTS can't express. Each is safe to run -// repeatedly: "duplicate column" errors from an already-applied -// migration are ignored. -var migrations = []string{ - `ALTER TABLE tasks ADD COLUMN period_id INTEGER REFERENCES periods(id)`, - `ALTER TABLE users ADD COLUMN hours_mon DOUBLE NOT NULL DEFAULT 0`, - `ALTER TABLE users ADD COLUMN hours_tue DOUBLE NOT NULL DEFAULT 0`, - `ALTER TABLE users ADD COLUMN hours_wed DOUBLE NOT NULL DEFAULT 0`, - `ALTER TABLE users ADD COLUMN hours_thu DOUBLE NOT NULL DEFAULT 0`, - `ALTER TABLE users ADD COLUMN hours_fri DOUBLE NOT NULL DEFAULT 0`, - `ALTER TABLE users ADD COLUMN hours_sat DOUBLE NOT NULL DEFAULT 0`, - `ALTER TABLE users ADD COLUMN hours_sun DOUBLE NOT NULL DEFAULT 0`, - `ALTER TABLE clients ADD COLUMN is_archived INTEGER NOT NULL DEFAULT 0`, +// Options tunes the SQLite connection. +type Options struct { + // WAL switches on write-ahead logging, which stops readers and writers + // from blocking each other. It costs two extra files next to the + // database (-wal and -shm), so it stays opt-in: capping the pool at a + // single connection already removes contention within this process. + WAL bool } -func Open(path string) (*sql.DB, error) { - conn, err := sql.Open("sqlite", path) +// Open connects to the SQLite database at path with foreign keys enforced, +// limits writes to a single connection (SQLite allows one writer), and applies +// all embedded migrations in alphabetical order. +// +// The pragmas live in the DSN rather than in a PRAGMA statement issued after +// connecting, because database/sql pools connections: a statement configures +// only whichever connection happened to serve it, leaving every other +// connection in the pool on the defaults. +func Open(path string, opts Options) (*sql.DB, error) { + dsn := path + "?_pragma=foreign_keys(1)&_pragma=busy_timeout(5000)" + if opts.WAL { + dsn += "&_pragma=journal_mode(WAL)" + } + + conn, err := sql.Open("sqlite", dsn) if err != nil { - return nil, err + return nil, fmt.Errorf("connect sqlite: %w", err) } - if _, err := conn.Exec("PRAGMA foreign_keys = ON;"); err != nil { + // SQLite permits only one writer at a time; a single connection avoids + // "database is locked" errors under concurrent writes. Every + // transaction must therefore run its statements on the tx, never on + // the pool, or it would wait on a connection it is itself holding. + conn.SetMaxOpenConns(1) + + if err := migrate(conn); err != nil { conn.Close() return nil, err } - if _, err := conn.Exec(schema); err != nil { - conn.Close() - return nil, err + return conn, nil +} + +// migrate applies embedded migration files in alphabetical order, skipping any +// already recorded in schema_migrations. The first migration is written to be +// idempotent, so it runs safely against databases that predate this ledger. +func migrate(conn *sql.DB) error { + if _, err := conn.Exec( + `CREATE TABLE IF NOT EXISTS schema_migrations (name TEXT PRIMARY KEY)`, + ); err != nil { + return fmt.Errorf("create schema_migrations: %w", err) + } + + entries, err := fs.ReadDir(migrations.FS, ".") + if err != nil { + return fmt.Errorf("read migrations: %w", err) } - for _, m := range migrations { - if _, err := conn.Exec(m); err != nil && !strings.Contains(err.Error(), "duplicate column") { - conn.Close() - return nil, err + names := make([]string, 0, len(entries)) + for _, e := range entries { + if !e.IsDir() { + names = append(names, e.Name()) } } - return conn, nil + sort.Strings(names) + + for _, name := range names { + var applied int + if err := conn.QueryRow( + `SELECT COUNT(*) FROM schema_migrations WHERE name = ?`, name, + ).Scan(&applied); err != nil { + return fmt.Errorf("check migration %s: %w", name, err) + } + if applied > 0 { + continue + } + stmts, err := fs.ReadFile(migrations.FS, name) + if err != nil { + return fmt.Errorf("read migration %s: %w", name, err) + } + if _, err := conn.Exec(string(stmts)); err != nil { + return fmt.Errorf("apply migration %s: %w", name, err) + } + if _, err := conn.Exec( + `INSERT INTO schema_migrations (name) VALUES (?)`, name, + ); err != nil { + return fmt.Errorf("record migration %s: %w", name, err) + } + } + return nil } diff --git a/internal/db/db_test.go b/internal/db/db_test.go new file mode 100644 index 0000000..9424fa0 --- /dev/null +++ b/internal/db/db_test.go @@ -0,0 +1,274 @@ +package db + +import ( + "database/sql" + "errors" + "path/filepath" + "testing" + + _ "modernc.org/sqlite" +) + +// open builds a migrated database in a temp dir. +func open(t *testing.T) *sql.DB { + t.Helper() + conn, err := Open(filepath.Join(t.TempDir(), "test.db"), Options{}) + if err != nil { + t.Fatalf("Open: %v", err) + } + t.Cleanup(func() { conn.Close() }) + return conn +} + +// seed creates one user with a client, a task type assigned to it, a period +// and a task tying them together. +func seed(t *testing.T, conn *sql.DB) { + t.Helper() + stmts := []string{ + `INSERT INTO users (id, email, password_hash) VALUES (1, 'a@b.c', 'x')`, + `INSERT INTO clients (id, name, user_id) VALUES (1, 'Acme', 1)`, + `INSERT INTO task_types (id, user_id, name) VALUES (1, 1, 'Dev')`, + `INSERT INTO periods (id, user_id, name, is_default) VALUES (1, 1, 'Q1', 0)`, + `INSERT INTO task_types_for_client (client_id, task_type_id) VALUES (1, 1)`, + `INSERT INTO tasks (id, user_id, client_id, task_type_id, period_id, title, hours_spent, date) + VALUES (1, 1, 1, 1, 1, 'work', 2, '2026-01-01')`, + } + for _, s := range stmts { + if _, err := conn.Exec(s); err != nil { + t.Fatalf("seed %q: %v", s, err) + } + } +} + +func count(t *testing.T, conn *sql.DB, query string, args ...any) int { + t.Helper() + var n int + if err := conn.QueryRow(query, args...).Scan(&n); err != nil { + t.Fatalf("%s: %v", query, err) + } + return n +} + +// TestForeignKeysOnEveryConnection is the point of putting the pragma in the +// DSN. Setting it with a PRAGMA statement configures only the connection that +// served the statement, and database/sql hands out others freely. +func TestForeignKeysOnEveryConnection(t *testing.T) { + conn := open(t) + conn.SetMaxIdleConns(0) // force a fresh connection each round + + for i := range 5 { + var on int + if err := conn.QueryRow(`PRAGMA foreign_keys`).Scan(&on); err != nil { + t.Fatal(err) + } + if on != 1 { + t.Errorf("connection %d: foreign_keys = %d, want 1", i, on) + } + } +} + +// TestWALIsOptIn guards the default: WAL leaves -wal and -shm files next to +// the database, so it only happens when asked for. +func TestWALIsOptIn(t *testing.T) { + tests := []struct { + name string + opts Options + want string + }{ + {"default", Options{}, "delete"}, + {"opt in", Options{WAL: true}, "wal"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + conn, err := Open(filepath.Join(t.TempDir(), "test.db"), tt.opts) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + + var mode string + if err := conn.QueryRow(`PRAGMA journal_mode`).Scan(&mode); err != nil { + t.Fatal(err) + } + if mode != tt.want { + t.Errorf("journal_mode = %q, want %q", mode, tt.want) + } + }) + } +} + +// TestBadReferenceRejected covers what foreign keys buy beyond tidy deletes: +// a task can no longer point at a task type that doesn't exist, or at one +// belonging to somebody else. +func TestBadReferenceRejected(t *testing.T) { + conn := open(t) + seed(t, conn) + + _, err := conn.Exec( + `INSERT INTO tasks (user_id, client_id, task_type_id, title, hours_spent, date) + VALUES (1, 1, 999, 'work', 1, '2026-01-01')`) + if !isForeignKeyErr(err) { + t.Errorf("insert with an unknown task_type_id = %v, want a foreign key error", err) + } +} + +// TestDeleteRefusedWhileReferenced covers the ErrInUse paths. Tasks are the +// records the app exists to keep, so nothing that still has hours logged +// against it can be deleted out from under them. +func TestDeleteRefusedWhileReferenced(t *testing.T) { + tests := []struct { + name string + remove func(*sql.DB) error + }{ + {"client", func(c *sql.DB) error { return DeleteClient(c, 1, 1) }}, + {"task type", func(c *sql.DB) error { return DeleteTaskType(c, 1, 1) }}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + conn := open(t) + seed(t, conn) + + if err := tt.remove(conn); !errors.Is(err, ErrInUse) { + t.Fatalf("delete = %v, want ErrInUse", err) + } + if n := count(t, conn, `SELECT count(*) FROM tasks`); n != 1 { + t.Errorf("tasks = %d, want the task left untouched", n) + } + }) + } +} + +// TestDeleteCascadesAssignments checks migration 0002: the join table is +// bookkeeping, so a delete that is otherwise allowed shouldn't be blocked by +// a checkbox someone ticked on an edit page. +func TestDeleteCascadesAssignments(t *testing.T) { + conn := open(t) + seed(t, conn) + + if _, err := conn.Exec(`DELETE FROM tasks`); err != nil { + t.Fatal(err) + } + if err := DeleteClient(conn, 1, 1); err != nil { + t.Fatalf("DeleteClient after removing its tasks: %v", err) + } + if n := count(t, conn, `SELECT count(*) FROM task_types_for_client`); n != 0 { + t.Errorf("task_types_for_client = %d rows, want the assignment cascaded away", n) + } +} + +// TestDeletePeriodDetachesTasks covers the one delete that is allowed to +// proceed while tasks reference it: a period only groups hours, so removing it +// must not remove them. +func TestDeletePeriodDetachesTasks(t *testing.T) { + conn := open(t) + seed(t, conn) + + if err := DeletePeriod(conn, 1, 1); err != nil { + t.Fatalf("DeletePeriod: %v", err) + } + if n := count(t, conn, `SELECT count(*) FROM periods`); n != 0 { + t.Errorf("periods = %d, want 0", n) + } + if n := count(t, conn, `SELECT count(*) FROM tasks WHERE period_id IS NULL`); n != 1 { + t.Errorf("detached tasks = %d, want the task kept with no period", n) + } +} + +// TestMigrationsAreIdempotent reopens the same file: every migration is +// already recorded, so the second open must be a no-op rather than an error. +func TestMigrationsAreIdempotent(t *testing.T) { + path := filepath.Join(t.TempDir(), "test.db") + + conn, err := Open(path, Options{}) + if err != nil { + t.Fatal(err) + } + seed(t, conn) + conn.Close() + + conn, err = Open(path, Options{}) + if err != nil { + t.Fatalf("reopen: %v", err) + } + defer conn.Close() + + if n := count(t, conn, `SELECT count(*) FROM tasks`); n != 1 { + t.Errorf("tasks = %d, want the seeded row intact", n) + } + if n := count(t, conn, `SELECT count(*) FROM schema_migrations`); n != 2 { + t.Errorf("schema_migrations = %d rows, want one per migration file", n) + } +} + +// legacySchema is the schema exactly as the pre-ledger startup code left it: +// a CREATE TABLE block plus a list of ALTER TABLE ADD COLUMN statements, run +// on every open. Databases in the wild are in this shape, and migration 0001 +// has to recognise them as already migrated. +const legacySchema = ` +CREATE TABLE users (id INTEGER PRIMARY KEY AUTOINCREMENT, email TEXT NOT NULL UNIQUE, password_hash TEXT NOT NULL); +CREATE TABLE clients (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, user_id INTEGER NOT NULL REFERENCES users(id)); +CREATE TABLE task_types (id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL REFERENCES users(id), name TEXT NOT NULL); +CREATE TABLE task_types_for_client (client_id INTEGER NOT NULL REFERENCES clients(id), task_type_id INTEGER NOT NULL REFERENCES task_types(id), PRIMARY KEY (client_id, task_type_id)); +CREATE TABLE periods (id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL REFERENCES users(id), name TEXT NOT NULL, is_default INTEGER NOT NULL DEFAULT 0); +CREATE TABLE tasks (id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL REFERENCES users(id), client_id INTEGER NOT NULL REFERENCES clients(id), task_type_id INTEGER NOT NULL REFERENCES task_types(id), title TEXT NOT NULL, hours_spent DOUBLE NOT NULL, date DATE NOT NULL); +ALTER TABLE tasks ADD COLUMN period_id INTEGER REFERENCES periods(id); +ALTER TABLE users ADD COLUMN hours_mon DOUBLE NOT NULL DEFAULT 0; +ALTER TABLE users ADD COLUMN hours_tue DOUBLE NOT NULL DEFAULT 0; +ALTER TABLE users ADD COLUMN hours_wed DOUBLE NOT NULL DEFAULT 0; +ALTER TABLE users ADD COLUMN hours_thu DOUBLE NOT NULL DEFAULT 0; +ALTER TABLE users ADD COLUMN hours_fri DOUBLE NOT NULL DEFAULT 0; +ALTER TABLE users ADD COLUMN hours_sat DOUBLE NOT NULL DEFAULT 0; +ALTER TABLE users ADD COLUMN hours_sun DOUBLE NOT NULL DEFAULT 0; +ALTER TABLE clients ADD COLUMN is_archived INTEGER NOT NULL DEFAULT 0; +` + +// TestMigratesLegacyDatabase is the upgrade path for existing installs: an +// untracked database keeps its rows, gains the ledger, and comes out with the +// cascade from 0002 applied. +func TestMigratesLegacyDatabase(t *testing.T) { + path := filepath.Join(t.TempDir(), "legacy.db") + + legacy, err := sql.Open("sqlite", path) + if err != nil { + t.Fatal(err) + } + if _, err := legacy.Exec(legacySchema); err != nil { + t.Fatal(err) + } + seed(t, legacy) + legacy.Close() + + conn, err := Open(path, Options{}) + if err != nil { + t.Fatalf("migrating a legacy database: %v", err) + } + defer conn.Close() + + if n := count(t, conn, `SELECT count(*) FROM tasks`); n != 1 { + t.Errorf("tasks = %d, want the existing row preserved", n) + } + if n := count(t, conn, `SELECT count(*) FROM task_types_for_client`); n != 1 { + t.Errorf("assignments = %d, want the rebuild to have copied it", n) + } + // The user's daily-hours columns came from the old ALTERs; 0001 must + // not have replaced the table and dropped them. + if n := count(t, conn, `SELECT count(*) FROM pragma_table_info('users') WHERE name LIKE 'hours_%'`); n != 7 { + t.Errorf("hours_* columns = %d, want 7", n) + } + if n := count(t, conn, `SELECT count(*) FROM schema_migrations`); n != 2 { + t.Errorf("schema_migrations = %d rows, want one per migration file", n) + } + + // The cascade only exists if 0002 actually rebuilt the table. + if _, err := conn.Exec(`DELETE FROM tasks`); err != nil { + t.Fatal(err) + } + if err := DeleteClient(conn, 1, 1); err != nil { + t.Fatalf("DeleteClient: %v", err) + } + if n := count(t, conn, `SELECT count(*) FROM task_types_for_client`); n != 0 { + t.Errorf("task_types_for_client = %d rows, want the assignment cascaded away", n) + } +} diff --git a/internal/db/periods.go b/internal/db/periods.go index d3b811e..6408d29 100644 --- a/internal/db/periods.go +++ b/internal/db/periods.go @@ -71,7 +71,24 @@ func UpdatePeriod(conn *sql.DB, userID, id int64, name string) error { return err } +// DeletePeriod removes a period and detaches the tasks filed under it. A +// period is only a label for grouping, so losing it must not lose the hours: +// the tasks stay, with no period. Both statements run on the transaction, not +// on the pool, which is capped at a single connection. func DeletePeriod(conn *sql.DB, userID, id int64) error { - _, err := conn.Exec(`DELETE FROM periods WHERE id = ? AND user_id = ?`, id, userID) - return err + tx, err := conn.Begin() + if err != nil { + return err + } + defer tx.Rollback() + + if _, err := tx.Exec( + `UPDATE tasks SET period_id = NULL WHERE period_id = ? AND user_id = ?`, id, userID, + ); err != nil { + return err + } + if _, err := tx.Exec(`DELETE FROM periods WHERE id = ? AND user_id = ?`, id, userID); err != nil { + return err + } + return tx.Commit() } diff --git a/internal/db/task_types.go b/internal/db/task_types.go index 001de6e..ee16a8a 100644 --- a/internal/db/task_types.go +++ b/internal/db/task_types.go @@ -48,7 +48,12 @@ func UpdateTaskType(conn *sql.DB, userID, id int64, name string) error { return err } +// DeleteTaskType removes a task type. It returns ErrInUse when tasks still +// carry it; the client assignments referencing it cascade away. func DeleteTaskType(conn *sql.DB, userID, id int64) error { _, err := conn.Exec(`DELETE FROM task_types WHERE id = ? AND user_id = ?`, id, userID) + if isForeignKeyErr(err) { + return ErrInUse + } return err } diff --git a/internal/handlers/auth.go b/internal/handlers/auth.go index 198fb62..6da1f30 100644 --- a/internal/handlers/auth.go +++ b/internal/handlers/auth.go @@ -33,15 +33,24 @@ func (h *Handlers) loginSubmit(w http.ResponseWriter, r *http.Request) { return } + h.setSessionCookie(w, token, int((24 * time.Hour).Seconds())) + http.Redirect(w, r, dest, http.StatusSeeOther) +} + +// setSessionCookie writes the session cookie. Login and logout both go through +// it so their attributes can't drift: a browser replaces a cookie only when the +// name, path and domain all match, so a clear that disagrees with the set +// leaves the original in place. +func (h *Handlers) setSessionCookie(w http.ResponseWriter, value string, maxAge int) { http.SetCookie(w, &http.Cookie{ Name: auth.CookieName, - Value: token, + Value: value, Path: "/", HttpOnly: true, SameSite: http.SameSiteLaxMode, - MaxAge: int((24 * time.Hour).Seconds()), + Secure: h.Config.SecureCookies, + MaxAge: maxAge, }) - http.Redirect(w, r, dest, http.StatusSeeOther) } func (h *Handlers) logout(w http.ResponseWriter, r *http.Request) { @@ -49,12 +58,6 @@ func (h *Handlers) logout(w http.ResponseWriter, r *http.Request) { h.Auth.Logout(cookie.Value) } - http.SetCookie(w, &http.Cookie{ - Name: auth.CookieName, - Value: "", - Path: "/", - HttpOnly: true, - MaxAge: -1, - }) + h.setSessionCookie(w, "", -1) http.Redirect(w, r, "/login", http.StatusSeeOther) } diff --git a/internal/handlers/auth_test.go b/internal/handlers/auth_test.go new file mode 100644 index 0000000..58aab10 --- /dev/null +++ b/internal/handlers/auth_test.go @@ -0,0 +1,117 @@ +package handlers + +import ( + "net/http" + "net/http/httptest" + "net/url" + "path/filepath" + "strings" + "testing" + "testing/fstest" + + "github.com/bloomyindev/time-tracker/internal/config" + "github.com/bloomyindev/time-tracker/internal/db" + "github.com/bloomyindev/time-tracker/internal/service/auth" +) + +// loggedIn builds a router backed by a real database holding one user, and +// returns it alongside that user's session token. +func loggedIn(t *testing.T, cfg config.Config) (http.Handler, string) { + t.Helper() + + conn, err := db.Open(filepath.Join(t.TempDir(), "test.db"), db.Options{}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { conn.Close() }) + + svc := auth.NewService(conn, "test-secret") + if err := svc.Register("a@b.c", "hunter2"); err != nil { + t.Fatal(err) + } + token, err := svc.Login("a@b.c", "hunter2") + if err != nil { + t.Fatal(err) + } + return New(conn, svc, cfg).Router(fstest.MapFS{}), token +} + +// findCookie returns the named cookie from the response, failing if absent. +func findCookie(t *testing.T, w *httptest.ResponseRecorder, name string) *http.Cookie { + t.Helper() + for _, c := range w.Result().Cookies() { + if c.Name == name { + return c + } + } + t.Fatalf("no %q cookie in the response", name) + return nil +} + +// TestSessionCookieSecureFollowsConfig covers both directions. The default is +// on, for a deployment behind a TLS-terminating proxy; plain-HTTP setups have +// to turn it off, or the browser accepts the cookie and never sends it back. +func TestSessionCookieSecureFollowsConfig(t *testing.T) { + for _, secure := range []bool{true, false} { + t.Run(map[bool]string{true: "secure", false: "insecure"}[secure], func(t *testing.T) { + router, _ := loggedIn(t, config.Config{SecureCookies: secure}) + + body := url.Values{"email": {"a@b.c"}, "password": {"hunter2"}} + r := httptest.NewRequest(http.MethodPost, "/login", strings.NewReader(body.Encode())) + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + w := httptest.NewRecorder() + router.ServeHTTP(w, r) + + if w.Code != http.StatusSeeOther { + t.Fatalf("login status = %d, want %d", w.Code, http.StatusSeeOther) + } + c := findCookie(t, w, auth.CookieName) + if c.Secure != secure { + t.Errorf("Secure = %t, want %t", c.Secure, secure) + } + if !c.HttpOnly { + t.Error("HttpOnly = false, want true") + } + if c.SameSite != http.SameSiteLaxMode { + t.Errorf("SameSite = %v, want Lax", c.SameSite) + } + }) + } +} + +// TestLogoutCookieMatchesLogin is why both go through one helper: a browser +// only replaces a cookie whose name, path and domain match, so a clear that +// disagrees with the set would leave the session cookie in place. +func TestLogoutCookieMatchesLogin(t *testing.T) { + cfg := config.Config{SecureCookies: true} + router, token := loggedIn(t, cfg) + + r := httptest.NewRequest(http.MethodGet, "/logout", nil) + r.AddCookie(&http.Cookie{Name: auth.CookieName, Value: token}) + w := httptest.NewRecorder() + router.ServeHTTP(w, r) + + c := findCookie(t, w, auth.CookieName) + if c.Value != "" { + t.Errorf("value = %q, want empty", c.Value) + } + if c.MaxAge >= 0 { + t.Errorf("MaxAge = %d, want negative so the browser drops it", c.MaxAge) + } + if !c.Secure || !c.HttpOnly || c.Path != "/" { + t.Errorf("attributes drifted from the login cookie: %+v", c) + } +} + +// TestLocaleCookieSecureFollowsConfig keeps the language cookie consistent +// with the session one; there is no reason for them to differ. +func TestLocaleCookieSecureFollowsConfig(t *testing.T) { + router, _ := loggedIn(t, config.Config{SecureCookies: true}) + + w := httptest.NewRecorder() + router.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/lang/fr", nil)) + + if c := findCookie(t, w, "lang"); !c.Secure { + t.Error("Secure = false, want true") + } +} diff --git a/internal/handlers/clients.go b/internal/handlers/clients.go index cd7a449..fb151aa 100644 --- a/internal/handlers/clients.go +++ b/internal/handlers/clients.go @@ -2,6 +2,7 @@ package handlers import ( "database/sql" + "errors" "net/http" "strconv" @@ -148,7 +149,12 @@ func (h *Handlers) deleteClient(w http.ResponseWriter, r *http.Request) { return } - if err := db.DeleteClient(h.DB, userID(r), id); err != nil { + err := db.DeleteClient(h.DB, userID(r), id) + if errors.Is(err, db.ErrInUse) { + http.Error(w, "client still has tasks; archive it instead", http.StatusConflict) + return + } + if err != nil { fail(w, err) return } diff --git a/internal/handlers/handlers.go b/internal/handlers/handlers.go index d188f7e..b19ce93 100644 --- a/internal/handlers/handlers.go +++ b/internal/handlers/handlers.go @@ -5,6 +5,7 @@ import ( "net/http" "strconv" + "github.com/bloomyindev/time-tracker/internal/config" "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/go-chi/chi/v5" ) @@ -13,12 +14,13 @@ import ( // methods on it, which is what lets each resource file expose its own // sub-router. type Handlers struct { - DB *sql.DB - Auth *auth.Service + DB *sql.DB + Auth *auth.Service + Config config.Config } -func New(conn *sql.DB, authSvc *auth.Service) *Handlers { - return &Handlers{DB: conn, Auth: authSvc} +func New(conn *sql.DB, authSvc *auth.Service, cfg config.Config) *Handlers { + return &Handlers{DB: conn, Auth: authSvc, Config: cfg} } // userID returns the authenticated user for the request. It panics if the diff --git a/internal/handlers/locale.go b/internal/handlers/locale.go index 0df44da..8499250 100644 --- a/internal/handlers/locale.go +++ b/internal/handlers/locale.go @@ -30,6 +30,7 @@ func (h *Handlers) setLocale(w http.ResponseWriter, r *http.Request) { Path: "/", MaxAge: int((365 * 24 * time.Hour).Seconds()), SameSite: http.SameSiteLaxMode, + Secure: h.Config.SecureCookies, }) dest := redirect.Sanitize(r.URL.Query().Get(redirect.Param), "/") diff --git a/internal/handlers/task_types.go b/internal/handlers/task_types.go index 53d8b20..da7c9ab 100644 --- a/internal/handlers/task_types.go +++ b/internal/handlers/task_types.go @@ -1,6 +1,7 @@ package handlers import ( + "errors" "net/http" "github.com/bloomyindev/time-tracker/internal/db" @@ -73,7 +74,12 @@ func (h *Handlers) deleteTaskType(w http.ResponseWriter, r *http.Request) { return } - if err := db.DeleteTaskType(h.DB, userID(r), id); err != nil { + err := db.DeleteTaskType(h.DB, userID(r), id) + if errors.Is(err, db.ErrInUse) { + http.Error(w, "task type is still used by tasks", http.StatusConflict) + return + } + if err != nil { fail(w, err) return } diff --git a/migrations/0001_init.sql b/migrations/0001_init.sql new file mode 100644 index 0000000..cefae13 --- /dev/null +++ b/migrations/0001_init.sql @@ -0,0 +1,56 @@ +-- 0001_init: the schema as it stood before migrations were tracked in a +-- ledger. Every statement is idempotent, because databases created by the +-- older startup code already have all of this: they ran a CREATE TABLE IF NOT +-- EXISTS block plus a list of ALTER TABLE ADD COLUMN statements on every open. +-- The columns those ALTERs added are inlined here, so a fresh database ends up +-- with exactly the same shape an existing one already has. + +CREATE TABLE IF NOT EXISTS users ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + email TEXT NOT NULL UNIQUE, + password_hash TEXT NOT NULL, + hours_mon DOUBLE NOT NULL DEFAULT 0, + hours_tue DOUBLE NOT NULL DEFAULT 0, + hours_wed DOUBLE NOT NULL DEFAULT 0, + hours_thu DOUBLE NOT NULL DEFAULT 0, + hours_fri DOUBLE NOT NULL DEFAULT 0, + hours_sat DOUBLE NOT NULL DEFAULT 0, + hours_sun DOUBLE NOT NULL DEFAULT 0 +); + +CREATE TABLE IF NOT EXISTS clients ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + user_id INTEGER NOT NULL REFERENCES users(id), + is_archived INTEGER NOT NULL DEFAULT 0 +); + +CREATE TABLE IF NOT EXISTS task_types ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + name TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS task_types_for_client ( + client_id INTEGER NOT NULL REFERENCES clients(id), + task_type_id INTEGER NOT NULL REFERENCES task_types(id), + PRIMARY KEY (client_id, task_type_id) +); + +CREATE TABLE IF NOT EXISTS periods ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + name TEXT NOT NULL, + is_default INTEGER NOT NULL DEFAULT 0 +); + +CREATE TABLE IF NOT EXISTS tasks ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + client_id INTEGER NOT NULL REFERENCES clients(id), + task_type_id INTEGER NOT NULL REFERENCES task_types(id), + title TEXT NOT NULL, + hours_spent DOUBLE NOT NULL, + date DATE NOT NULL, + period_id INTEGER REFERENCES periods(id) +); diff --git a/migrations/0002_task_type_client_cascade.sql b/migrations/0002_task_type_client_cascade.sql new file mode 100644 index 0000000..f4f904c --- /dev/null +++ b/migrations/0002_task_type_client_cascade.sql @@ -0,0 +1,25 @@ +-- 0002_task_type_client_cascade: foreign keys are enforced from this release +-- on, and the join table between clients and task types is the one place where +-- a delete should take its rows with it rather than be refused. Deleting a +-- client or a task type must not be blocked by a checkbox someone ticked on an +-- edit page; the rows carry no information of their own. +-- +-- Tasks are deliberately left out of this: their foreign keys stay RESTRICT so +-- deleting a client can't silently destroy the time logged against it. That +-- refusal surfaces as db.ErrInUse and, in the UI, as a 409. +-- +-- SQLite can't ALTER a constraint, so the table is rebuilt. Nothing references +-- it, which is what makes the drop-and-rename safe to do in place. + +CREATE TABLE task_types_for_client_new ( + client_id INTEGER NOT NULL REFERENCES clients(id) ON DELETE CASCADE, + task_type_id INTEGER NOT NULL REFERENCES task_types(id) ON DELETE CASCADE, + PRIMARY KEY (client_id, task_type_id) +); + +INSERT INTO task_types_for_client_new (client_id, task_type_id) +SELECT client_id, task_type_id FROM task_types_for_client; + +DROP TABLE task_types_for_client; + +ALTER TABLE task_types_for_client_new RENAME TO task_types_for_client; diff --git a/migrations/migrations.go b/migrations/migrations.go new file mode 100644 index 0000000..c0b08bb --- /dev/null +++ b/migrations/migrations.go @@ -0,0 +1,9 @@ +// Package migrations embeds the versioned SQL migration files so they ship +// inside the binary. It lives in migrations/ because //go:embed only reaches +// files at or below the embedding source file's directory. +package migrations + +import "embed" + +//go:embed *.sql +var FS embed.FS From 48d6ec67d7a0e0b0b87db94da7e781aa5057ddb9 Mon Sep 17 00:00:00 2001 From: Bloomy Date: Thu, 30 Jul 2026 18:01:28 +0200 Subject: [PATCH 3/3] feat: make the listen port configurable The port was the one setting hardcoded in Config, so the only way to move the server was to remap it outside the process. It reads TRACKER_PORT now, alongside the other settings. An out-of-range or unparseable value falls back to 8080 with a warning rather than failing to start, matching how the boolean settings already behave. Zero is rejected with the rest: it would ask the kernel for an arbitrary free port, which is never what a server someone has to reach was meant to do. Config is read once into a package-level variable instead of per command. Load logs as it falls back, and every command builds its flags at startup, so a single malformed value used to print its warning five times. Co-Authored-By: Claude Opus 5 --- README.md | 5 ++- cmd/time-tracker/admin.go | 3 +- cmd/time-tracker/main.go | 11 +++-- internal/config/config.go | 22 ++++++++- internal/config/config_test.go | 81 ++++++++++++++++++++++++++++++++++ 5 files changed, 113 insertions(+), 9 deletions(-) create mode 100644 internal/config/config_test.go diff --git a/README.md b/README.md index 2487787..a0402d1 100644 --- a/README.md +++ b/README.md @@ -109,6 +109,7 @@ The app is configured through environment variables: | Variable | Default | Description | |---------------------------|------------------------|------------------------------------------------------------------------------------------------------| +| `TRACKER_PORT` | `8080` | TCP port the server listens on. | | `TRACKER_DB_PATH` | `time-tracker.db` | Path to the SQLite database file. | | `TRACKER_JWT_SECRET` | `dev-secret-change-me` | Secret for signing JWTs used by the (currently unused) bearer-token API flow. Set this in production. | | `TRACKER_SECURE_COOKIES` | `true` | Mark cookies `Secure`, so browsers only send them over HTTPS. Set to `false` to serve plain HTTP. | @@ -119,8 +120,8 @@ assumes a TLS-terminating reverse proxy in front. Left on over plain HTTP, the browser accepts the session cookie and then never sends it back, so logging in appears to do nothing. -The server always listens on port `8080`. The database path can also be passed -with `--db-path`. Migrations are embedded in the binary and applied on startup; +The database path can also be passed with `--db-path`, which takes precedence +over the environment. Migrations are embedded in the binary and applied on startup; each one is recorded in a `schema_migrations` table so it runs exactly once. ## CLI diff --git a/cmd/time-tracker/admin.go b/cmd/time-tracker/admin.go index 8530533..771f81d 100644 --- a/cmd/time-tracker/admin.go +++ b/cmd/time-tracker/admin.go @@ -6,7 +6,6 @@ import ( "fmt" "os" - "github.com/bloomyindev/time-tracker/internal/config" "github.com/bloomyindev/time-tracker/internal/db" "github.com/bloomyindev/time-tracker/internal/service/auth" "github.com/urfave/cli/v3" @@ -28,7 +27,7 @@ func registerCommand() *cli.Command { } defer conn.Close() - svc := auth.NewService(conn, config.Load().JWTSecret) + svc := auth.NewService(conn, cfg.JWTSecret) email := cmd.String("email") if err := svc.Register(email, cmd.String("password")); err != nil { return fmt.Errorf("register: %w", err) diff --git a/cmd/time-tracker/main.go b/cmd/time-tracker/main.go index 3e8ebe4..b4f86f2 100644 --- a/cmd/time-tracker/main.go +++ b/cmd/time-tracker/main.go @@ -18,6 +18,11 @@ import ( "github.com/urfave/cli/v3" ) +// cfg is read once for the whole process. Load reports malformed values as it +// falls back to a default, and every command builds its flags at startup, so +// reading it per command would print each of those warnings several times over. +var cfg = config.Load() + func main() { cmd := &cli.Command{ Name: "time-tracker", @@ -40,13 +45,13 @@ func dbPathFlag() *cli.StringFlag { return &cli.StringFlag{ Name: "db-path", Usage: "path to the sqlite database file", - Value: config.Load().DBPath, + Value: cfg.DBPath, Sources: cli.EnvVars("TRACKER_DB_PATH"), } } func openDB(cmd *cli.Command) (*sql.DB, error) { - return db.Open(cmd.String("db-path"), db.Options{WAL: config.Load().SQLiteWAL}) + return db.Open(cmd.String("db-path"), db.Options{WAL: cfg.SQLiteWAL}) } func serveCommand() *cli.Command { @@ -55,8 +60,6 @@ func serveCommand() *cli.Command { Usage: "run the web server", Flags: []cli.Flag{dbPathFlag()}, Action: func(ctx context.Context, cmd *cli.Command) error { - cfg := config.Load() - if err := i18n.Load(); err != nil { return fmt.Errorf("load locales: %w", err) } diff --git a/internal/config/config.go b/internal/config/config.go index 6aae4a7..dd649f4 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -27,7 +27,7 @@ type Config struct { func Load() Config { return Config{ - Port: 8080, + Port: getPort("TRACKER_PORT", 8080), DBPath: getEnv("TRACKER_DB_PATH", "time-tracker.db"), JWTSecret: getEnv("TRACKER_JWT_SECRET", "dev-secret-change-me"), SQLiteWAL: getBool("TRACKER_SQLITE_WAL", false), @@ -42,6 +42,26 @@ func getEnv(key, fallback string) string { return fallback } +// getPort reads a TCP port. Zero is rejected along with the out-of-range +// values: asking the kernel to pick a free port is never what a server someone +// has to reach was meant to do. +func getPort(key string, fallback int) int { + raw := os.Getenv(key) + if raw == "" { + return fallback + } + parsed, err := strconv.Atoi(raw) + if err != nil { + log.Printf("%s: %q isn't a number, using %d", key, raw, fallback) + return fallback + } + if parsed < 1 || parsed > 65535 { + log.Printf("%s: %d isn't a valid port, using %d", key, parsed, fallback) + return fallback + } + return parsed +} + // getBool reads a boolean in any form strconv accepts ("1", "true", "off"). // An unparseable value is a typo in the deployment, not a reason to run with a // setting nobody chose, so it is reported and the default stands. diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..9f9cf95 --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,81 @@ +package config + +import "testing" + +// TestLoadDefaults pins the defaults that matter for a deployment nobody +// configured: cookies secure, WAL off, the usual port. The variables are +// blanked first, since every getter reads an empty value as unset and the +// developer running this may well have some of them exported. +func TestLoadDefaults(t *testing.T) { + for _, key := range []string{ + "TRACKER_PORT", "TRACKER_DB_PATH", "TRACKER_JWT_SECRET", + "TRACKER_SQLITE_WAL", "TRACKER_SECURE_COOKIES", + } { + t.Setenv(key, "") + } + + cfg := Load() + + if cfg.Port != 8080 { + t.Errorf("Port = %d, want 8080", cfg.Port) + } + if !cfg.SecureCookies { + t.Error("SecureCookies = false, want true so an unconfigured deploy is the safe one") + } + if cfg.SQLiteWAL { + t.Error("SQLiteWAL = true, want false so no -wal/-shm files appear unasked") + } +} + +func TestGetPort(t *testing.T) { + tests := []struct { + name string + set string + want int + }{ + {"unset falls back", "", 8080}, + {"valid port", "3000", 3000}, + {"lowest valid", "1", 1}, + {"highest valid", "65535", 65535}, + {"zero rejected", "0", 8080}, + {"negative rejected", "-1", 8080}, + {"above range rejected", "65536", 8080}, + {"not a number", "http", 8080}, + {"empty is unset", "", 8080}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Setenv("TRACKER_TEST_PORT", tt.set) + if got := getPort("TRACKER_TEST_PORT", 8080); got != tt.want { + t.Errorf("getPort(%q) = %d, want %d", tt.set, got, tt.want) + } + }) + } +} + +func TestGetBool(t *testing.T) { + tests := []struct { + name string + set string + fallback bool + want bool + }{ + {"unset keeps default", "", true, true}, + {"true", "true", false, true}, + {"one", "1", false, true}, + {"false", "false", true, false}, + {"zero", "0", true, false}, + {"garbage keeps default", "yes please", true, true}, + {"garbage keeps a false default", "nope", false, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Setenv("TRACKER_TEST_BOOL", tt.set) + if got := getBool("TRACKER_TEST_BOOL", tt.fallback); got != tt.want { + t.Errorf("getBool(%q, %t) = %t, want %t", tt.set, tt.fallback, got, tt.want) + } + }) + } +}