package api import ( "context" "crypto/sha256" "database/sql" "encoding/hex" "net/http" "strings" "time" "git.ryuvia.com/niklas/terdut-server/internal/models" ) type contextKey string const ( ctxUser contextKey = "user" ctxSession contextKey = "session" ) // AuthMiddleware accepts either of the two credentials the server issues: an // API key in an Authorization header (the TUI, scripts) or a session cookie // (the web UI). A request carrying a Bearer header is judged on that alone and // never falls back to the cookie. // // Only the cookie needs a CSRF guard. A browser attaches it to requests other // sites make, whereas an Authorization header is only ever set by the client // that holds the key. func AuthMiddleware(db *sql.DB) func(http.Handler) http.Handler { crossOrigin := http.NewCrossOriginProtection() return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if header := r.Header.Get("Authorization"); header != "" { token, ok := strings.CutPrefix(header, "Bearer ") if !ok || token == "" { respond(w, http.StatusUnauthorized, errResp("unauthorized")) return } userID, ok := apiKeyUser(r.Context(), db, token) if !ok { respond(w, http.StatusUnauthorized, errResp("unauthorized")) return } serveAs(w, r, next, db, userID, 0) return } c, err := r.Cookie(sessionCookie) if err != nil || c.Value == "" { respond(w, http.StatusUnauthorized, errResp("unauthorized")) return } sessionID, userID, ok := sessionUser(r.Context(), db, c.Value) if !ok { respond(w, http.StatusUnauthorized, errResp("unauthorized")) return } if err := crossOrigin.Check(r); err != nil { respond(w, http.StatusForbidden, errResp("cross-origin request rejected")) return } serveAs(w, r, next, db, userID, sessionID) }) } } // apiKeyUser resolves an API key to its user and stamps its last use. func apiKeyUser(ctx context.Context, db *sql.DB, token string) (int64, bool) { var keyID, userID int64 err := db.QueryRowContext(ctx, "SELECT id, user_id FROM api_keys WHERE key_hash = $1", hashToken(token), ).Scan(&keyID, &userID) if err != nil { return 0, false } // best-effort; don't fail the request if this update fails db.ExecContext(ctx, "UPDATE api_keys SET last_used_at = $1 WHERE id = $2", time.Now().Unix(), keyID) return userID, true } // sessionUser resolves a session token to its session and user. The expiry // slides forward with use, but at most once per sessionTouchEvery, so a page // that polls does not write to the database on every request. func sessionUser(ctx context.Context, db *sql.DB, token string) (sessionID, userID int64, ok bool) { now := time.Now() var lastSeen int64 err := db.QueryRowContext(ctx, ` SELECT id, user_id, last_seen_at FROM sessions WHERE token_hash = $1 AND expires_at > $2`, hashToken(token), now.Unix()).Scan(&sessionID, &userID, &lastSeen) if err != nil { return 0, 0, false } if now.Sub(time.Unix(lastSeen, 0)) > sessionTouchEvery { db.ExecContext(ctx, "UPDATE sessions SET last_seen_at = $1, expires_at = $2 WHERE id = $3", now.Unix(), now.Add(sessionTTL).Unix(), sessionID) } return sessionID, userID, true } // serveAs loads the user and hands the request on with it in the context. // sessionID is zero for API-key requests. func serveAs(w http.ResponseWriter, r *http.Request, next http.Handler, db *sql.DB, userID, sessionID int64) { var u models.User var createdUnix int64 if err := db.QueryRowContext(r.Context(), "SELECT id, username, email, created_at FROM users WHERE id = $1", userID, ).Scan(&u.ID, &u.Username, &u.Email, &createdUnix); err != nil { respond(w, http.StatusUnauthorized, errResp("unauthorized")) return } u.CreatedAt = time.Unix(createdUnix, 0).UTC() ctx := context.WithValue(r.Context(), ctxUser, u) if sessionID != 0 { ctx = context.WithValue(ctx, ctxSession, sessionID) } next.ServeHTTP(w, r.WithContext(ctx)) } func hashToken(token string) string { h := sha256.Sum256([]byte(token)) return hex.EncodeToString(h[:]) } func userFromContext(ctx context.Context) (models.User, bool) { u, ok := ctx.Value(ctxUser).(models.User) return u, ok } // sessionFromContext returns the id of the session a request was authenticated // with, or false for an API-key request. func sessionFromContext(ctx context.Context) (int64, bool) { id, ok := ctx.Value(ctxSession).(int64) return id, ok }