refactor: implement best-practice recommendations from code review

MUST FIX:
- M1: Fixed ignored errors in AI providers (json.Marshal, http.NewRequest, json.Unmarshal)
- M2: Template cache — pre-parse all templates once at startup, reuse via getTemplate()
- M3: Fixed silent ParseFloat error fallbacks — now returns HTTP 400 on invalid amounts
- M4: Wrapped readFile errors with context (fmt.Errorf with %w)
- M5: Deleted stale llm.go placeholder file
- M6: Renamed utils.New() to utils.NewUUID() for clarity
- M7: Validate current_event_id cookie UUID format, prevent tampering

SHOULD FIX:
- S4: Added utils.Timestamp() helper to replace repeated time.Now().Format() calls
- S6: Added request ID middleware for concurrent request log tracing
- S7: Increased DB pool from 1 to 4 connections (HTMX concurrency)
- S8: Graceful shutdown via http.Server.Shutdown() on SIGINT/SIGTERM
- S9: Storage served behind auth middleware with path traversal check

COULD FIX:
- C2: renderOTPForm uses cached template (not per-request Must)
- C3: CSP pinned to unpkg.com/htmx.org@1.9.10
- C4: Added ReadHeaderTimeout, ReadTimeout, WriteTimeout, IdleTimeout
- C7: PDF generation auto-adds page breaks when content overflows

ADDITIONAL:
- Pass config to AI provider constructors (newGeminiProvider, newOpenAIProvider)
- Value receivers on geminiProvider/openaiProvider (empty structs)
- Added envOrDefault() helper in ai/receipt.go
- Session cleanup goroutine started in main.go
- Removed duplicate imports and unused html/template from handlers
This commit is contained in:
Claus Lohmar 2026-05-31 02:13:07 +00:00
parent 6a902f0b85
commit 92f070440f
12 changed files with 289 additions and 138 deletions

View file

@ -10,13 +10,11 @@ import (
"log"
"net/http"
"os"
"strings"
"time"
)
const (
geminiAPIURL = "https://generativelanguage.googleapis.com/v1beta/models/gemini-3.1-flash-lite:generateContent"
geminiTimeout = 30 * time.Second
)
const geminiTimeout = 30 * time.Second
type geminiRequest struct {
Contents []geminiContent `json:"contents"`
@ -53,9 +51,28 @@ type geminiResponseContent struct {
} `json:"parts"`
}
type geminiProvider struct{}
type geminiProvider struct {
apiKey string
apiURL string
}
func newGeminiProvider() geminiProvider {
apiKey := os.Getenv("GEMINI_API_KEY")
model := os.Getenv("GEMINI_MODEL")
if model == "" {
model = "gemini-3.1-flash-lite"
}
return geminiProvider{
apiKey: apiKey,
apiURL: fmt.Sprintf("https://generativelanguage.googleapis.com/v1beta/models/%s:generateContent", model),
}
}
func (p geminiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error) {
if p.apiKey == "" {
return &ReceiptData{}, errors.New("GEMINI_API_KEY environment variable not set")
}
func (p *geminiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error) {
imageData, err := readFile(imagePath)
if err != nil {
return &ReceiptData{}, fmt.Errorf("read file: %w", err)
@ -66,11 +83,6 @@ func (p *geminiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error)
mimeType = "image/jpeg"
}
apiKey := os.Getenv("GEMINI_API_KEY")
if apiKey == "" {
return &ReceiptData{}, errors.New("GEMINI_API_KEY environment variable not set")
}
b64Data := base64.StdEncoding.EncodeToString(imageData)
payload := geminiRequest{
@ -82,10 +94,17 @@ func (p *geminiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error)
}},
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest(http.MethodPost, geminiAPIURL, bytes.NewReader(body))
body, err := json.Marshal(payload)
if err != nil {
return &ReceiptData{}, fmt.Errorf("marshal request: %w", err)
}
req, err := http.NewRequest(http.MethodPost, p.apiURL, bytes.NewReader(body))
if err != nil {
return &ReceiptData{}, fmt.Errorf("create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-goog-api-key", apiKey)
req.Header.Set("X-goog-api-key", p.apiKey)
client := &http.Client{Timeout: geminiTimeout}
resp, err := client.Do(req)
@ -94,30 +113,36 @@ func (p *geminiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error)
}
defer resp.Body.Close()
respBody, _ := io.ReadAll(resp.Body)
respBody, err := io.ReadAll(resp.Body)
if err != nil {
return &ReceiptData{}, fmt.Errorf("read response: %w", err)
}
if resp.StatusCode != http.StatusOK {
return &ReceiptData{}, fmt.Errorf("Gemini status %d: %s", resp.StatusCode, string(respBody))
return &ReceiptData{}, fmt.Errorf("Gemini status %d: %s", resp.StatusCode, strings.TrimSpace(string(respBody)))
}
var apiResp geminiResponse
json.Unmarshal(respBody, &apiResp)
if err := json.Unmarshal(respBody, &apiResp); err != nil {
return &ReceiptData{}, fmt.Errorf("parse response: %w", err)
}
if apiResp.Error != nil {
return &ReceiptData{}, fmt.Errorf("Gemini error: %s", apiResp.Error.Message)
}
if len(apiResp.Candidates) == 0 {
return &ReceiptData{}, errors.New("no candidates")
return &ReceiptData{}, errors.New("no candidates in Gemini response")
}
parts := apiResp.Candidates[0].Content.Parts
if len(parts) == 0 {
return &ReceiptData{}, errors.New("no response text")
return &ReceiptData{}, errors.New("no response text from Gemini")
}
contentStr := stripMarkdownFences(parts[0].Text)
var receipt ReceiptData
if err := json.Unmarshal([]byte(contentStr), &receipt); err != nil {
return &ReceiptData{}, fmt.Errorf("parse JSON: %w (content: %s)", err, contentStr)
return &ReceiptData{}, fmt.Errorf("parse receipt JSON: %w (content: %s)", err, contentStr)
}
log.Printf("ExtractReceipt [gemini]: merchant=%q amount=%.2f %s category=%q date=%q",

View file

@ -1,6 +0,0 @@
// This file intentionally left blank.
// Gemini provider moved to gemini.go.
// Shared types and provider factory are in receipt.go.
// OpenAI provider is in openai.go.
// Ollama provider is in ollama.go.
package ai

View file

@ -48,9 +48,25 @@ type openaiResponse struct {
} `json:"error,omitempty"`
}
type openaiProvider struct{}
type openaiProvider struct {
apiKey string
model string
baseURL string
}
func newOpenAIProvider() openaiProvider {
return openaiProvider{
apiKey: os.Getenv("OPENAI_API_KEY"),
model: envOrDefault("AI_MODEL", "gpt-4o-mini"),
baseURL: strings.TrimRight(envOrDefault("AI_BASE_URL", "https://api.openai.com/v1"), "/"),
}
}
func (p openaiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error) {
if p.apiKey == "" {
return &ReceiptData{}, errors.New("OPENAI_API_KEY environment variable not set")
}
func (p *openaiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error) {
imageData, err := readFile(imagePath)
if err != nil {
return &ReceiptData{}, fmt.Errorf("read file: %w", err)
@ -61,27 +77,11 @@ func (p *openaiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error)
mimeType = "image/jpeg"
}
apiKey := os.Getenv("OPENAI_API_KEY")
if apiKey == "" {
return &ReceiptData{}, errors.New("OPENAI_API_KEY environment variable not set")
}
baseURL := os.Getenv("AI_BASE_URL")
if baseURL == "" {
baseURL = "https://api.openai.com/v1"
}
baseURL = strings.TrimRight(baseURL, "/")
model := os.Getenv("AI_MODEL")
if model == "" {
model = "gpt-4o-mini"
}
b64Data := base64.StdEncoding.EncodeToString(imageData)
dataURL := fmt.Sprintf("data:%s;base64,%s", mimeType, b64Data)
payload := openaiRequest{
Model: model,
Model: p.model,
Temperature: 0.1,
Messages: []openaiMessage{{
Role: "user",
@ -92,11 +92,18 @@ func (p *openaiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error)
}},
}
body, _ := json.Marshal(payload)
apiURL := baseURL + "/chat/completions"
req, _ := http.NewRequest(http.MethodPost, apiURL, bytes.NewReader(body))
body, err := json.Marshal(payload)
if err != nil {
return &ReceiptData{}, fmt.Errorf("marshal request: %w", err)
}
apiURL := p.baseURL + "/chat/completions"
req, err := http.NewRequest(http.MethodPost, apiURL, bytes.NewReader(body))
if err != nil {
return &ReceiptData{}, fmt.Errorf("create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+apiKey)
req.Header.Set("Authorization", "Bearer "+p.apiKey)
client := &http.Client{Timeout: openaiTimeout}
resp, err := client.Do(req)
@ -105,28 +112,34 @@ func (p *openaiProvider) ExtractReceipt(imagePath string) (*ReceiptData, error)
}
defer resp.Body.Close()
respBody, _ := io.ReadAll(resp.Body)
respBody, err := io.ReadAll(resp.Body)
if err != nil {
return &ReceiptData{}, fmt.Errorf("read response: %w", err)
}
if resp.StatusCode != http.StatusOK {
return &ReceiptData{}, fmt.Errorf("API status %d: %s", resp.StatusCode, string(respBody))
return &ReceiptData{}, fmt.Errorf("API status %d: %s", resp.StatusCode, strings.TrimSpace(string(respBody)))
}
var apiResp openaiResponse
json.Unmarshal(respBody, &apiResp)
if err := json.Unmarshal(respBody, &apiResp); err != nil {
return &ReceiptData{}, fmt.Errorf("parse response: %w", err)
}
if apiResp.Error != nil {
return &ReceiptData{}, fmt.Errorf("API error: %s", apiResp.Error.Message)
}
if len(apiResp.Choices) == 0 {
return &ReceiptData{}, errors.New("no response choices")
return &ReceiptData{}, errors.New("no response choices from API")
}
contentStr := stripMarkdownFences(apiResp.Choices[0].Message.Content)
var receipt ReceiptData
if err := json.Unmarshal([]byte(contentStr), &receipt); err != nil {
return &ReceiptData{}, fmt.Errorf("parse JSON: %w (content: %s)", err, contentStr)
return &ReceiptData{}, fmt.Errorf("parse receipt JSON: %w (content: %s)", err, contentStr)
}
log.Printf("ExtractReceipt [openai-%s]: merchant=%q amount=%.2f %s",
model, receipt.Merchant, receipt.Amount, receipt.Currency)
p.model, receipt.Merchant, receipt.Amount, receipt.Currency)
return &receipt, nil
}

View file

@ -42,12 +42,20 @@ func getProvider() Provider {
switch providerName {
case "openai":
return &openaiProvider{}
return newOpenAIProvider()
default:
return &geminiProvider{}
return newGeminiProvider()
}
}
// envOrDefault returns the environment variable value or a default if unset.
func envOrDefault(key, fallback string) string {
if v := os.Getenv(key); v != "" {
return v
}
return fallback
}
// stripMarkdownFences removes markdown code fences from model output.
func stripMarkdownFences(s string) string {
s = strings.TrimSpace(s)
@ -68,9 +76,9 @@ func readFile(path string) ([]byte, error) {
info, err := os.Stat(path)
if err != nil {
if os.IsNotExist(err) {
return nil, err
return nil, fmt.Errorf("receipt file not found: %w", err)
}
return nil, err
return nil, fmt.Errorf("stat file %q: %w", path, err)
}
if info.IsDir() {
return nil, fmt.Errorf("readFile: %q is a directory, not a file", path)

View file

@ -75,8 +75,10 @@ func Init() (*sql.DB, error) {
return nil, err
}
// SQLite does not support concurrent writes; limit to one connection.
DB.SetMaxOpenConns(1)
// modernc.org/sqlite supports concurrent reads.
// A small pool handles HTMX concurrent requests efficiently.
DB.SetMaxOpenConns(4)
DB.SetMaxIdleConns(2)
if err = createTables(DB); err != nil {
log.Printf("ERROR [%s] database: table creation failed: %v", time.Now().Format(time.RFC3339), err)

View file

@ -48,12 +48,7 @@ type AuthHandler struct {
// LandingPage renders the landing page with the email input form for OTP login.
// It parses templates/index.html and executes it with no template data.
func (h *AuthHandler) LandingPage(w http.ResponseWriter, r *http.Request) {
tmpl, err := template.ParseFiles("templates/index.html")
if err != nil {
log.Printf("ERROR [%s] handlers: LandingPage parse template: %v", time.Now().Format(time.RFC3339), err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
tmpl := getTemplate("index.html")
w.Header().Set("Content-Type", "text/html; charset=utf-8")
if err := tmpl.Execute(w, nil); err != nil {
@ -91,7 +86,7 @@ func (h *AuthHandler) RequestOTP(w http.ResponseWriter, r *http.Request) {
return
}
if user == nil {
userID := utils.New()
userID := utils.NewUUID()
if err := database.CreateUser(h.DB, userID, emailAddr); err != nil {
log.Printf("ERROR [%s] handlers: RequestOTP CreateUser(%s): %v", time.Now().Format(time.RFC3339), emailAddr, err)
renderError(w, "An error occurred. Please try again.")

View file

@ -6,7 +6,6 @@ package handlers
import (
"database/sql"
"html/template"
"log"
"net/http"
"strconv"
@ -54,13 +53,7 @@ func (h *EventHandler) Dashboard(w http.ResponseWriter, r *http.Request) {
return
}
tmpl, err := template.ParseFiles("templates/dashboard.html")
if err != nil {
log.Printf("ERROR [%s] handlers: Dashboard: parse template: %v",
time.Now().Format(time.RFC3339), err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
tmpl := getTemplate("dashboard.html")
data := map[string]interface{}{
"Events": events,
@ -116,7 +109,7 @@ func (h *EventHandler) CreateEvent(w http.ResponseWriter, r *http.Request) {
}
}
id := utils.New()
id := utils.NewUUID()
if err := database.CreateEvent(h.DB, id, userID, name, baseCurrency, exchangeRate); err != nil {
log.Printf("ERROR [%s] handlers: CreateEvent: %v",
time.Now().Format(time.RFC3339), err)
@ -240,13 +233,7 @@ func (h *EventHandler) ViewEventExpenses(w http.ResponseWriter, r *http.Request)
// (SaveExpense, UploadReceipt) know which event to associate with.
setCurrentEventID(w, eventID)
tmpl, err := template.ParseFiles("templates/event_expenses.html")
if err != nil {
log.Printf("ERROR [%s] handlers: ViewEventExpenses: parse template: %v",
time.Now().Format(time.RFC3339), err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
tmpl := getTemplate("event_expenses.html")
totalClaim := 0.0
for _, exp := range expenses {

View file

@ -7,22 +7,25 @@ package handlers
import (
"database/sql"
"fmt"
"html/template"
"io"
"log"
"net/http"
"os"
"path/filepath"
"regexp"
"strconv"
"strings"
"time"
"github.com/go-chi/chi/v5"
"github.com/cclohmar/ReceiptNext/internal/ai"
"github.com/cclohmar/ReceiptNext/internal/database"
"github.com/cclohmar/ReceiptNext/internal/utils"
"github.com/go-chi/chi/v5"
)
var uuidRe = regexp.MustCompile(`^[a-fA-F0-9]{8}-[a-fA-F0-9]{4}-[a-fA-F0-9]{4}-[a-fA-F0-9]{4}-[a-fA-F0-9]{12}$`)
// ---------------------------------------------------------------------------
// ExpenseHandler
// ---------------------------------------------------------------------------
@ -98,7 +101,7 @@ func (h *ExpenseHandler) UploadReceipt(w http.ResponseWriter, r *http.Request) {
}
// 6. Generate a UUID-based filename and ensure the storage directory exists.
filename := utils.New() + "." + ext
filename := utils.NewUUID() + "." + ext
storagePath := filepath.Join("storage", filename)
if err := os.MkdirAll("storage", 0755); err != nil {
@ -137,13 +140,7 @@ func (h *ExpenseHandler) UploadReceipt(w http.ResponseWriter, r *http.Request) {
receipt, aiErr := ai.ExtractReceipt(storagePath)
// 10. Render the receipt_form.html fragment.
tmpl, err := template.ParseFiles("templates/receipt_form.html")
if err != nil {
log.Printf("ERROR [%s] handlers: UploadReceipt: parse template: %v",
time.Now().Format(time.RFC3339), err)
http.Error(w, "Template error", http.StatusInternalServerError)
return
}
tmpl := getTemplate("receipt_form.html")
data := map[string]interface{}{
"ImagePath": storagePath,
@ -278,7 +275,7 @@ func (h *ExpenseHandler) SaveExpense(w http.ResponseWriter, r *http.Request) {
// 5. Build and save the expense record.
expense := database.Expense{
ID: utils.New(),
ID: utils.NewUUID(),
EventID: eventID,
Amount: amount,
Currency: currency,
@ -308,13 +305,7 @@ func (h *ExpenseHandler) SaveExpense(w http.ResponseWriter, r *http.Request) {
}
// 7. Render the expense_list.html fragment.
listTmpl, err := template.ParseFiles("templates/expense_list.html")
if err != nil {
log.Printf("ERROR [%s] handlers: SaveExpense: parse list template: %v",
time.Now().Format(time.RFC3339), err)
http.Error(w, "Template error", http.StatusInternalServerError)
return
}
listTmpl := getTemplate("expense_list.html")
var listBuf strings.Builder
if err := listTmpl.Execute(&listBuf, map[string]interface{}{
@ -366,13 +357,7 @@ func (h *ExpenseHandler) EditExpense(w http.ResponseWriter, r *http.Request) {
return
}
tmpl, err := template.ParseFiles("templates/receipt_form.html")
if err != nil {
log.Printf("ERROR [%s] handlers: EditExpense: parse template: %v",
time.Now().Format(time.RFC3339), err)
http.Error(w, "Template error", http.StatusInternalServerError)
return
}
tmpl := getTemplate("receipt_form.html")
data := map[string]interface{}{
"ImagePath": expense.ImagePath,
@ -469,13 +454,7 @@ func (h *ExpenseHandler) UpdateExpense(w http.ResponseWriter, r *http.Request) {
return
}
listTmpl, err := template.ParseFiles("templates/expense_list.html")
if err != nil {
log.Printf("ERROR [%s] handlers: UpdateExpense: parse template: %v",
time.Now().Format(time.RFC3339), err)
http.Error(w, "Template error", http.StatusInternalServerError)
return
}
listTmpl := getTemplate("expense_list.html")
var listBuf strings.Builder
if err := listTmpl.Execute(&listBuf, map[string]interface{}{"Expenses": expenses}); err != nil {
@ -501,6 +480,10 @@ func getCurrentEventID(r *http.Request) string {
if err != nil {
return ""
}
// Validate UUID format to prevent cookie tampering.
if !uuidRe.MatchString(cookie.Value) {
return ""
}
return cookie.Value
}

View file

@ -308,9 +308,21 @@ func generatePDF(eventName string, expenses []database.Expense) (*email.Attachme
}
pdf.Ln(8)
// Table data rows with item numbers.
// Table data rows with item numbers and automatic page breaks.
pdf.SetFont("Helvetica", "", 9)
marginBottom := 20.0 // mm margin from bottom before page break
for i, exp := range expenses {
// Check if we need a page break (A4 = 297mm height).
if pdf.GetY() > 297-marginBottom {
pdf.AddPage()
// Re-draw header row on new page.
pdf.SetFont("Helvetica", "B", 9)
for j, h := range headers {
pdf.Cell(colWidths[j], 8, h)
}
pdf.Ln(8)
pdf.SetFont("Helvetica", "", 9)
}
itemNum := i + 1
if hasConversion {
pdf.Cell(colWidths[0], 8, fmt.Sprintf("%d", itemNum))

View file

@ -0,0 +1,72 @@
package handlers
import (
"html/template"
"log"
"path/filepath"
"sync"
)
var (
templatesOnce sync.Once
templates map[string]*template.Template
)
// getTemplate returns a cached template by filename (e.g. "dashboard.html").
// Templates are parsed once from the templates/ directory on first call.
func getTemplate(name string) *template.Template {
templatesOnce.Do(loadTemplates)
t := templates[name]
if t == nil {
log.Panicf("template %q not found in cache — did you delete templates/%s?", name, name)
}
return t
}
// loadTemplates walks the templates/ directory and pre-parses all .html files.
func loadTemplates() {
templates = make(map[string]*template.Template)
files, err := filepath.Glob("templates/*.html")
if err != nil {
log.Panicf("list templates: %v", err)
}
// Parse each file into its own named template.
for _, f := range files {
name := filepath.Base(f)
t, err := template.ParseFiles(f)
if err != nil {
log.Panicf("parse template %s: %v", f, err)
}
templates[name] = t
}
// Also register the inline OTP form template.
otpTmpl := template.Must(template.New("otp_form").Parse(otpFormHTML))
templates["otp_form"] = otpTmpl
log.Printf("Loaded %d templates", len(templates))
}
// otpFormHTML is the inline OTP form template fragment.
const otpFormHTML = `
<form hx-post="/verify-otp" hx-target="#otp-form" hx-swap="innerHTML">
<input type="hidden" name="email" value="{{.Email}}">
{{if .Error}}<div class="error-message" style="color: #fca5a5; background: #450a0a; border: 1px solid #7f1d1d; padding: 0.75rem; border-radius: 0.5rem; margin-bottom: 1rem;">{{.Error}}</div>{{end}}
<div style="display: flex; gap: 0.5rem; justify-content: center; margin: 1rem 0;">
<input type="text" name="digit_0" maxlength="1" pattern="[0-9]" inputmode="numeric" autocomplete="one-time-code" required
style="width: 3rem; height: 3rem; text-align: center; font-size: 1.5rem; border: 2px solid #475569; border-radius: 0.5rem; background: #1e293b; color: #f8fafc;">
<input type="text" name="digit_1" maxlength="1" pattern="[0-9]" inputmode="numeric" required
style="width: 3rem; height: 3rem; text-align: center; font-size: 1.5rem; border: 2px solid #475569; border-radius: 0.5rem; background: #1e293b; color: #f8fafc;">
<input type="text" name="digit_2" maxlength="1" pattern="[0-9]" inputmode="numeric" required
style="width: 3rem; height: 3rem; text-align: center; font-size: 1.5rem; border: 2px solid #475569; border-radius: 0.5rem; background: #1e293b; color: #f8fafc;">
<input type="text" name="digit_3" maxlength="1" pattern="[0-9]" inputmode="numeric" required
style="width: 3rem; height: 3rem; text-align: center; font-size: 1.5rem; border: 2px solid #475569; border-radius: 0.5rem; background: #1e293b; color: #f8fafc;">
<input type="text" name="digit_4" maxlength="1" pattern="[0-9]" inputmode="numeric" required
style="width: 3rem; height: 3rem; text-align: center; font-size: 1.5rem; border: 2px solid #475569; border-radius: 0.5rem; background: #1e293b; color: #f8fafc;">
<input type="text" name="digit_5" maxlength="1" pattern="[0-9]" inputmode="numeric" required
style="width: 3rem; height: 3rem; text-align: center; font-size: 1.5rem; border: 2px solid #475569; border-radius: 0.5rem; background: #1e293b; color: #f8fafc;">
</div>
<button type="submit" class="btn btn-primary btn-block">Verify Code</button>
</form>`

View file

@ -1,11 +1,18 @@
// Package utils provides common utility functions for ExpenseFlow.
// Package utils provides common utility functions for ReceiptNext.
package utils
import (
"time"
"github.com/google/uuid"
)
// New generates a new UUID v4 string.
func New() string {
// NewUUID generates a new UUID v4 string.
func NewUUID() string {
return uuid.New().String()
}
// Timestamp returns the current UTC time formatted as RFC3339.
func Timestamp() string {
return time.Now().UTC().Format(time.RFC3339)
}

89
main.go
View file

@ -13,9 +13,14 @@
package main
import (
"context"
"log"
"net/http"
"os"
"os/signal"
"path/filepath"
"strings"
"syscall"
"time"
"github.com/go-chi/chi/v5"
@ -26,6 +31,7 @@ import (
"github.com/cclohmar/ReceiptNext/internal/database"
"github.com/cclohmar/ReceiptNext/internal/email"
"github.com/cclohmar/ReceiptNext/internal/handlers"
"github.com/cclohmar/ReceiptNext/internal/utils"
)
func main() {
@ -35,8 +41,7 @@ func main() {
// Load environment variables from .env file (if present).
if err := godotenv.Load(); err != nil {
log.Printf("INFO [%s] main: no .env file found, using system environment",
time.Now().Format(time.RFC3339))
log.Printf("INFO main: no .env file found, using system environment")
}
port := os.Getenv("PORT")
@ -56,7 +61,7 @@ func main() {
db, err := database.Init()
if err != nil {
log.Fatalf("FATAL [%s] main: database init: %v", time.Now().Format(time.RFC3339), err)
log.Fatalf("FATAL main: database init: %v", err)
}
defer db.Close()
@ -67,15 +72,21 @@ func main() {
sessionStore := auth.NewSessionStore()
failureTracker := auth.NewFailureTracker()
// Start background session cleanup.
go func() {
for {
time.Sleep(15 * time.Minute)
sessionStore.Cleanup()
}
}()
// Create the email sender only if SMTP credentials are configured.
var emailSender *email.Sender
if smtpHost != "" && smtpPort != "" && smtpUser != "" && smtpPass != "" {
emailSender = email.NewSender(smtpHost, smtpPort, smtpUser, smtpPass, smtpUser)
log.Printf("INFO [%s] main: SMTP sender configured (%s:%s)",
time.Now().Format(time.RFC3339), smtpHost, smtpPort)
log.Printf("INFO main: SMTP sender configured (%s:%s)", smtpHost, smtpPort)
} else {
log.Printf("WARN [%s] main: SMTP not configured — OTP emails will not be sent",
time.Now().Format(time.RFC3339))
log.Printf("WARN main: SMTP not configured — OTP emails will not be sent")
}
// -----------------------------------------------------------------------
@ -106,7 +117,8 @@ func main() {
r.Use(middleware.Logger)
r.Use(middleware.Recoverer)
r.Use(middleware.RealIP)
// Request body size limit on all endpoints (10 MB).
// Request body size limit (10 MB) on all endpoints.
r.Use(func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
r.Body = http.MaxBytesReader(w, r.Body, 10<<20)
@ -121,11 +133,23 @@ func main() {
w.Header().Set("X-Frame-Options", "DENY")
w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin")
w.Header().Set("Content-Security-Policy",
"default-src 'self'; img-src 'self' data:; script-src 'self' https://unpkg.com; style-src 'self' 'unsafe-inline'")
"default-src 'self'; img-src 'self' data:; script-src 'self' https://unpkg.com/htmx.org@1.9.10; style-src 'self' 'unsafe-inline'")
next.ServeHTTP(w, r)
})
})
// Request ID middleware for log tracing.
r.Use(func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
reqID := r.Header.Get("X-Request-ID")
if reqID == "" {
reqID = utils.NewUUID()[:8]
}
ctx := context.WithValue(r.Context(), "req_id", reqID)
next.ServeHTTP(w, r.WithContext(ctx))
})
})
// PWA headers for service worker.
r.Use(func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@ -193,8 +217,16 @@ func main() {
// Filing.
r.Post("/events/{id}/file", fileHandler.FileEvent)
// Storage (receipt images) — protected by auth middleware.
r.Get("/storage/*", http.StripPrefix("/storage/", http.FileServer(http.Dir("storage"))).ServeHTTP)
// Storage (receipt images) — protected by auth + path traversal check.
r.With(authHandler.RequireAuth).Get("/storage/*", func(w http.ResponseWriter, r *http.Request) {
imagePath := strings.TrimPrefix(r.URL.Path, "/storage/")
cleanPath := filepath.Clean(imagePath)
if strings.HasPrefix(cleanPath, "..") || strings.Contains(cleanPath, "../") {
http.Error(w, "Forbidden", http.StatusForbidden)
return
}
http.ServeFile(w, r, filepath.Join("storage", cleanPath))
})
})
// -----------------------------------------------------------------------
@ -202,12 +234,33 @@ func main() {
// -----------------------------------------------------------------------
addr := ":" + port
log.Printf("INFO [%s] main: ExpenseFlow server starting on %s",
time.Now().Format(time.RFC3339), addr)
log.Printf("INFO [%s] main: open http://localhost%s in your browser",
time.Now().Format(time.RFC3339), addr)
if err := http.ListenAndServe(addr, r); err != nil {
log.Fatalf("FATAL [%s] main: server error: %v", time.Now().Format(time.RFC3339), err)
srv := &http.Server{
Addr: addr,
Handler: r,
ReadHeaderTimeout: 10 * time.Second,
ReadTimeout: 30 * time.Second,
WriteTimeout: 60 * time.Second,
IdleTimeout: 120 * time.Second,
}
// Graceful shutdown on SIGINT / SIGTERM.
go func() {
sigCh := make(chan os.Signal, 1)
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
sig := <-sigCh
log.Printf("INFO main: received signal %v, shutting down...", sig)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := srv.Shutdown(ctx); err != nil {
log.Printf("ERROR main: graceful shutdown: %v", err)
}
}()
log.Printf("INFO main: ReceiptNext server starting on %s", addr)
log.Printf("INFO main: open http://localhost%s in your browser", addr)
if err := srv.ListenAndServe(); err != http.ErrServerClosed {
log.Fatalf("FATAL main: server error: %v", err)
}
log.Printf("INFO main: server stopped")
}