diff --git a/internal/ai/gemini.go b/internal/ai/gemini.go index a9681ff..05e64ac 100644 --- a/internal/ai/gemini.go +++ b/internal/ai/gemini.go @@ -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", diff --git a/internal/ai/llm.go b/internal/ai/llm.go deleted file mode 100644 index 84e858e..0000000 --- a/internal/ai/llm.go +++ /dev/null @@ -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 diff --git a/internal/ai/openai.go b/internal/ai/openai.go index ec1e6f0..4106e3b 100644 --- a/internal/ai/openai.go +++ b/internal/ai/openai.go @@ -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 } diff --git a/internal/ai/receipt.go b/internal/ai/receipt.go index bbcbdc8..2f9739e 100644 --- a/internal/ai/receipt.go +++ b/internal/ai/receipt.go @@ -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) diff --git a/internal/database/db.go b/internal/database/db.go index 4423c32..94dff29 100644 --- a/internal/database/db.go +++ b/internal/database/db.go @@ -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) diff --git a/internal/handlers/auth.go b/internal/handlers/auth.go index 759fb50..c5c512f 100644 --- a/internal/handlers/auth.go +++ b/internal/handlers/auth.go @@ -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.") diff --git a/internal/handlers/events.go b/internal/handlers/events.go index 162a1fb..d02a1e7 100644 --- a/internal/handlers/events.go +++ b/internal/handlers/events.go @@ -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 { diff --git a/internal/handlers/expenses.go b/internal/handlers/expenses.go index 74ccc07..e0c9746 100644 --- a/internal/handlers/expenses.go +++ b/internal/handlers/expenses.go @@ -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 } diff --git a/internal/handlers/file.go b/internal/handlers/file.go index 45af5ec..822f49b 100644 --- a/internal/handlers/file.go +++ b/internal/handlers/file.go @@ -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)) diff --git a/internal/handlers/templates.go b/internal/handlers/templates.go new file mode 100644 index 0000000..acde3e3 --- /dev/null +++ b/internal/handlers/templates.go @@ -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 = ` +
` diff --git a/internal/utils/uuid.go b/internal/utils/uuid.go index 14bfe48..7964886 100644 --- a/internal/utils/uuid.go +++ b/internal/utils/uuid.go @@ -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) +} diff --git a/main.go b/main.go index 1604433..7003d1a 100644 --- a/main.go +++ b/main.go @@ -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") }