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:
parent
6a902f0b85
commit
92f070440f
12 changed files with 289 additions and 138 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.")
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
72
internal/handlers/templates.go
Normal file
72
internal/handlers/templates.go
Normal 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>`
|
||||
|
|
@ -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
89
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")
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue