hsa-app/internal/web/web.go
Jean-Michel Tremblay 5c21f131fd Add AI auto-fill, tally, recent views, and cheaper-by-default classification
Features (see spec.md v2):
- Wire receipt classification into the upload flow; cheap model (Haiku 4.5)
  is now the default, shown as a footnote with per-scan cost in cents.
- Skip-AI toggle to enter fields by hand.
- Duplicate-transaction warning: live check on date+amount, gated submit.
- Tally tab: person x year totals with margins and grand total.
- Recent uploads / recent receipts tabs with paging and file serving.
- People reconcile on startup: merge stray partial names (e.g. "Jude" ->
  "Jude Tremblay"), reassigning receipts; idempotent seeding.
- scripts/build.sh builds the binary; scripts/run.sh builds and runs with .env.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-18 22:04:42 -04:00

232 lines
7.6 KiB
Go

package web
import (
"context"
"fmt"
"log"
"net/http"
"time"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
"maisym.com/hsa/internal/auth"
"maisym.com/hsa/internal/classify"
"maisym.com/hsa/internal/config"
"maisym.com/hsa/internal/storage"
)
const (
sessionCookie = "hsa_session"
loginCookie = "hsa_login"
)
// Server holds the wired dependencies for the HTTP handlers.
type Server struct {
cfg config.Config
provider *oidc.Provider
oauth *oauth2.Config
verifier *oidc.IDTokenVerifier
store *storage.Store
classifier *classify.Classifier // nil when classification is disabled
}
// NewServer performs OIDC discovery against the issuer and builds the server.
// This makes a network call to the issuer's /.well-known/openid-configuration.
// classifier may be nil to disable receipt auto-classification.
func NewServer(ctx context.Context, cfg config.Config, store *storage.Store, classifier *classify.Classifier) (*Server, error) {
provider, err := oidc.NewProvider(ctx, cfg.IssuerURL)
if err != nil {
return nil, fmt.Errorf("oidc discovery against %s: %w", cfg.IssuerURL, err)
}
oauthCfg := &oauth2.Config{
ClientID: cfg.ClientID,
ClientSecret: cfg.ClientSecret,
RedirectURL: cfg.RedirectURL,
Endpoint: provider.Endpoint(),
Scopes: []string{oidc.ScopeOpenID, "profile", "email", "groups"},
}
return &Server{
cfg: cfg,
provider: provider,
oauth: oauthCfg,
verifier: provider.Verifier(&oidc.Config{ClientID: cfg.ClientID}),
store: store,
classifier: classifier,
}, nil
}
// Routes returns the configured HTTP handler.
func (s *Server) Routes() http.Handler {
mux := http.NewServeMux()
mux.HandleFunc("GET /healthz", s.handleHealthz)
mux.Handle("GET /static/", http.FileServerFS(staticFS))
mux.HandleFunc("GET /login", s.handleLogin)
mux.HandleFunc("GET /callback", s.handleCallback)
mux.HandleFunc("GET /logout", s.handleLogout)
mux.Handle("GET /{$}", s.requireAuth(http.HandlerFunc(s.handleUploadForm)))
mux.Handle("POST /upload", s.requireAuth(http.HandlerFunc(s.handleUpload)))
mux.Handle("POST /classify", s.requireAuth(http.HandlerFunc(s.handleClassify)))
mux.Handle("GET /duplicates", s.requireAuth(http.HandlerFunc(s.handleDuplicates)))
mux.Handle("GET /tally", s.requireAuth(http.HandlerFunc(s.handleTally)))
mux.Handle("GET /recent", s.requireAuth(http.HandlerFunc(s.handleRecentUploads)))
mux.Handle("GET /recent/receipts", s.requireAuth(http.HandlerFunc(s.handleRecentReceipts)))
mux.Handle("GET /receipt/{id}/file", s.requireAuth(http.HandlerFunc(s.handleReceiptFile)))
mux.Handle("GET /manage", s.requireAuth(http.HandlerFunc(s.handleManage)))
mux.Handle("POST /manage/categories", s.requireAuth(http.HandlerFunc(s.handleAddCategory)))
mux.Handle("POST /manage/categories/rename", s.requireAuth(http.HandlerFunc(s.handleRenameCategory)))
mux.Handle("POST /manage/people", s.requireAuth(http.HandlerFunc(s.handleAddPerson)))
mux.Handle("POST /manage/people/rename", s.requireAuth(http.HandlerFunc(s.handleRenamePerson)))
mux.Handle("GET /export", s.requireAuth(http.HandlerFunc(s.handleExportPage)))
mux.Handle("GET /export/db", s.requireAuth(http.HandlerFunc(s.handleExportDB)))
return mux
}
func (s *Server) handleHealthz(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
fmt.Fprintln(w, "ok")
}
// handleLogin starts the OIDC flow: generate PKCE + state + nonce, stash them in
// a short-lived encrypted cookie, and redirect to Authelia's authorize endpoint.
func (s *Server) handleLogin(w http.ResponseWriter, r *http.Request) {
verifier, err := auth.GenerateVerifier()
if err != nil {
s.serverError(w, "generate verifier", err)
return
}
state, err := auth.GenerateState()
if err != nil {
s.serverError(w, "generate state", err)
return
}
nonce, err := auth.GenerateState() // reuse the random generator for the nonce
if err != nil {
s.serverError(w, "generate nonce", err)
return
}
ls := auth.LoginState{State: state, Verifier: verifier, Nonce: nonce}
encoded, err := auth.EncodeLoginState(ls, s.cfg.SessionKey)
if err != nil {
s.serverError(w, "encode login state", err)
return
}
s.setCookie(w, loginCookie, encoded, 10*time.Minute)
challenge := auth.ChallengeS256(verifier)
url := auth.AuthorizeURL(s.oauth, state, challenge, oidc.Nonce(nonce))
http.Redirect(w, r, url, http.StatusFound)
}
// handleCallback completes the OIDC flow: validate state, exchange the code with
// PKCE, verify the ID token, check the group, and set the session cookie.
func (s *Server) handleCallback(w http.ResponseWriter, r *http.Request) {
// Recover the login state from the cookie.
c, err := r.Cookie(loginCookie)
if err != nil {
http.Error(w, "missing login state; start at /login", http.StatusBadRequest)
return
}
ls, err := auth.DecodeLoginState(c.Value, s.cfg.SessionKey)
if err != nil {
http.Error(w, "invalid login state", http.StatusBadRequest)
return
}
s.clearCookie(w, loginCookie)
// CSRF: the state we issued must match the one returned.
if r.URL.Query().Get("state") != ls.State {
http.Error(w, "state mismatch", http.StatusBadRequest)
return
}
// Exchange the authorization code, sending the PKCE verifier.
code := r.URL.Query().Get("code")
token, err := s.oauth.Exchange(r.Context(), code, oauth2.VerifierOption(ls.Verifier))
if err != nil {
s.serverError(w, "token exchange", err)
return
}
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
http.Error(w, "no id_token in response", http.StatusBadGateway)
return
}
idToken, err := s.verifier.Verify(r.Context(), rawIDToken)
if err != nil {
s.serverError(w, "verify id_token", err)
return
}
if idToken.Nonce != ls.Nonce {
http.Error(w, "nonce mismatch", http.StatusBadRequest)
return
}
// Authelia serves groups/email from the UserInfo endpoint rather than the ID
// token by default, so read claims from the ID token and merge in UserInfo.
type oidcClaims struct {
Email string `json:"email"`
PreferredUsername string `json:"preferred_username"`
Groups []string `json:"groups"`
}
var idClaims oidcClaims
if err := idToken.Claims(&idClaims); err != nil {
s.serverError(w, "parse id token claims", err)
return
}
var uiClaims oidcClaims
if userInfo, err := s.provider.UserInfo(r.Context(), oauth2.StaticTokenSource(token)); err == nil {
_ = userInfo.Claims(&uiClaims)
} else {
log.Printf("warning: userinfo fetch failed: %v", err)
}
groups := idClaims.Groups
if len(groups) == 0 {
groups = uiClaims.Groups
}
email := idClaims.Email
if email == "" {
email = uiClaims.Email
}
username := idClaims.PreferredUsername
if username == "" {
username = uiClaims.PreferredUsername
}
subject := email
if subject == "" {
subject = username
}
log.Printf("login attempt: subject=%q groups=%v", subject, groups)
// Authorization gate: must be in the required group.
if !auth.IsAuthorized(groups, s.cfg.RequiredGroup) {
http.Error(w, fmt.Sprintf("403: not a member of %q", s.cfg.RequiredGroup), http.StatusForbidden)
return
}
sess := auth.Session{Subject: subject, Groups: groups, IssuedAt: time.Now().UTC()}
encoded, err := auth.EncodeSession(sess, s.cfg.SessionKey)
if err != nil {
s.serverError(w, "encode session", err)
return
}
s.setCookie(w, sessionCookie, encoded, 12*time.Hour)
http.Redirect(w, r, "/", http.StatusFound)
}
func (s *Server) handleLogout(w http.ResponseWriter, r *http.Request) {
s.clearCookie(w, sessionCookie)
http.Redirect(w, r, "/login", http.StatusFound)
}
func (s *Server) serverError(w http.ResponseWriter, what string, err error) {
log.Printf("error: %s: %v", what, err)
http.Error(w, "internal server error", http.StatusInternalServerError)
}