feat: initial commit — backend API + student cabinet frontend
- Go backend: auth (JWT), points earn/spend, QR token generation, partners, admin grant/stats endpoints with chi router - Next.js 14 frontend: login, student dashboard, transaction history, QR display, partners list - PostgreSQL migrations (4 tables), Redis cache, Docker Compose - CORS middleware, role-based route protection, Zustand auth store Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,8 @@
|
||||
[build]
|
||||
cmd = "go build -o ./tmp/main ./cmd/api"
|
||||
bin = "./tmp/main"
|
||||
include_ext = ["go"]
|
||||
exclude_dir = ["tmp", "vendor"]
|
||||
|
||||
[misc]
|
||||
clean_on_exit = true
|
||||
@@ -0,0 +1,140 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"github.com/cu-points/backend/internal/admin"
|
||||
"github.com/cu-points/backend/internal/auth"
|
||||
"github.com/cu-points/backend/internal/config"
|
||||
"github.com/cu-points/backend/internal/middleware"
|
||||
"github.com/cu-points/backend/internal/partners"
|
||||
"github.com/cu-points/backend/internal/points"
|
||||
"github.com/cu-points/backend/internal/users"
|
||||
cachepkg "github.com/cu-points/backend/pkg/cache"
|
||||
dbpkg "github.com/cu-points/backend/pkg/db"
|
||||
)
|
||||
|
||||
func main() {
|
||||
slog.SetDefault(slog.New(slog.NewJSONHandler(os.Stdout, nil)))
|
||||
|
||||
cfg, err := config.Load()
|
||||
if err != nil {
|
||||
slog.Error("config load failed", "err", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
db, err := dbpkg.NewPool(ctx, cfg.DatabaseURL)
|
||||
if err != nil {
|
||||
slog.Error("postgres connect failed", "err", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
redisClient, err := cachepkg.NewClient(ctx, cfg.RedisURL)
|
||||
if err != nil {
|
||||
slog.Error("redis connect failed", "err", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer redisClient.Close()
|
||||
|
||||
slog.Info("infrastructure connected", "postgres", cfg.DatabaseURL, "redis", cfg.RedisURL)
|
||||
|
||||
// ── Auth ──────────────────────────────────────────────────────────────────
|
||||
jwtManager := auth.NewJWTManager(cfg.JWTSecret, cfg.JWTAccessTTL, cfg.JWTRefreshTTL)
|
||||
authRepo := auth.NewRepository(db)
|
||||
authSvc := auth.NewService(authRepo, jwtManager)
|
||||
authHandler := auth.NewHandler(authSvc)
|
||||
|
||||
// ── Users ─────────────────────────────────────────────────────────────────
|
||||
usersRepo := users.NewRepository(db)
|
||||
usersSvc := users.NewService(usersRepo)
|
||||
usersHandler := users.NewHandler(usersSvc)
|
||||
|
||||
// ── Partners ──────────────────────────────────────────────────────────────
|
||||
partnersRepo := partners.NewRepository(db)
|
||||
partnersSvc := partners.NewService(partnersRepo)
|
||||
partnersHandler := partners.NewHandler(partnersSvc)
|
||||
|
||||
// ── Points ────────────────────────────────────────────────────────────────
|
||||
pointsRepo := points.NewRepository(db)
|
||||
pointsCache := points.NewRedisCache(redisClient)
|
||||
// pointsSvc is also passed to adminSvc so GrantPoints reuses EarnPoints logic.
|
||||
pointsSvc := points.NewService(pointsRepo, pointsCache, cfg.JWTSecret)
|
||||
pointsHandler := points.NewHandler(pointsSvc)
|
||||
|
||||
// ── Admin ─────────────────────────────────────────────────────────────────
|
||||
adminSvc := admin.NewService(db, pointsSvc)
|
||||
adminHandler := admin.NewHandler(adminSvc)
|
||||
|
||||
// ── Router ────────────────────────────────────────────────────────────────
|
||||
r := chi.NewRouter()
|
||||
r.Use(middleware.CORS)
|
||||
r.Use(middleware.Logger)
|
||||
|
||||
r.Route("/api/v1", func(r chi.Router) {
|
||||
// Public — no authentication required.
|
||||
r.Post("/auth/login", authHandler.Login)
|
||||
r.Post("/auth/refresh", authHandler.Refresh)
|
||||
r.Get("/partners", partnersHandler.List)
|
||||
|
||||
// Any authenticated user — profile endpoint used by the login flow for all roles.
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(middleware.Auth(cfg.JWTSecret))
|
||||
|
||||
r.Get("/me", usersHandler.Me)
|
||||
})
|
||||
|
||||
// Student-only endpoints.
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(middleware.Auth(cfg.JWTSecret))
|
||||
r.Use(middleware.RequireRole("student"))
|
||||
|
||||
r.Get("/me/transactions", usersHandler.Transactions)
|
||||
r.Get("/me/qr", pointsHandler.GenerateQR)
|
||||
})
|
||||
|
||||
// Partner-only endpoints.
|
||||
// Rate-limited to 10 spend requests per minute per partner to prevent abuse.
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(middleware.Auth(cfg.JWTSecret))
|
||||
r.Use(middleware.RequireRole("partner"))
|
||||
r.Use(middleware.SpendRateLimit(redisClient, 10))
|
||||
|
||||
r.Post("/partner/spend", pointsHandler.Spend)
|
||||
})
|
||||
|
||||
// Admin-only endpoints.
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(middleware.Auth(cfg.JWTSecret))
|
||||
r.Use(middleware.RequireRole("admin"))
|
||||
|
||||
r.Post("/admin/points/grant", adminHandler.GrantPoints)
|
||||
r.Get("/admin/transactions", adminHandler.ListTransactions)
|
||||
r.Get("/admin/users", adminHandler.ListUsers)
|
||||
r.Get("/admin/stats", adminHandler.Stats)
|
||||
})
|
||||
})
|
||||
|
||||
// Health check — no auth required.
|
||||
r.Get("/health", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
fmt.Fprintln(w, `{"status":"ok"}`)
|
||||
})
|
||||
|
||||
addr := ":" + cfg.Port
|
||||
slog.Info("server starting", "addr", addr)
|
||||
if err := http.ListenAndServe(addr, r); err != nil {
|
||||
slog.Error("server failed", "err", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
// seed inserts development test users into the database.
|
||||
// Run via: make seed (only use against local dev DB, never production)
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
type seedUser struct {
|
||||
email string
|
||||
name string
|
||||
password string
|
||||
role string
|
||||
studentID string // empty string → NULL in DB
|
||||
}
|
||||
|
||||
var devUsers = []seedUser{
|
||||
{
|
||||
email: "student@cu.ru",
|
||||
name: "Иван Студентов",
|
||||
password: "password123",
|
||||
role: "student",
|
||||
studentID: "STU001",
|
||||
},
|
||||
{
|
||||
email: "partner@cu.ru",
|
||||
name: "Кофейня Уют",
|
||||
password: "password123",
|
||||
role: "partner",
|
||||
},
|
||||
{
|
||||
email: "admin@cu.ru",
|
||||
name: "Администратор ЦУ",
|
||||
password: "password123",
|
||||
role: "admin",
|
||||
},
|
||||
}
|
||||
|
||||
func main() {
|
||||
dbURL := os.Getenv("DATABASE_URL")
|
||||
if dbURL == "" {
|
||||
slog.Error("DATABASE_URL is not set")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
db, err := pgxpool.New(ctx, dbURL)
|
||||
if err != nil {
|
||||
slog.Error("connect failed", "err", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
for _, u := range devUsers {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(u.password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
slog.Error("bcrypt failed", "email", u.email, "err", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// NULLIF converts empty string to NULL for student_id
|
||||
_, err = db.Exec(ctx, `
|
||||
INSERT INTO users (email, name, password_hash, role, student_id)
|
||||
VALUES ($1, $2, $3, $4, NULLIF($5, ''))
|
||||
ON CONFLICT (email) DO UPDATE
|
||||
SET password_hash = EXCLUDED.password_hash,
|
||||
name = EXCLUDED.name
|
||||
`, u.email, u.name, string(hash), u.role, u.studentID)
|
||||
if err != nil {
|
||||
slog.Error("seed insert failed", "email", u.email, "err", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Printf("seeded: %-30s role=%-8s password=%s\n", u.email, u.role, u.password)
|
||||
}
|
||||
|
||||
fmt.Println("\nDone. Test credentials above are for local development only.")
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
module github.com/cu-points/backend
|
||||
|
||||
go 1.22
|
||||
|
||||
require (
|
||||
github.com/go-chi/chi/v5 v5.1.0
|
||||
github.com/golang-jwt/jwt/v5 v5.2.1
|
||||
github.com/jackc/pgx/v5 v5.6.0
|
||||
github.com/redis/go-redis/v9 v9.6.1
|
||||
golang.org/x/crypto v0.17.0
|
||||
)
|
||||
|
||||
require (
|
||||
github.com/cespare/xxhash/v2 v2.2.0 // indirect
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
|
||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||
github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a // indirect
|
||||
github.com/jackc/puddle/v2 v2.2.1 // indirect
|
||||
golang.org/x/sync v0.1.0 // indirect
|
||||
golang.org/x/text v0.14.0 // indirect
|
||||
)
|
||||
@@ -0,0 +1,42 @@
|
||||
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
||||
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
||||
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
||||
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
|
||||
github.com/cespare/xxhash/v2 v2.2.0 h1:DC2CZ1Ep5Y4k3ZQ899DldepgrayRUGE6BBZ/cd9Cj44=
|
||||
github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78=
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
|
||||
github.com/go-chi/chi/v5 v5.1.0 h1:acVI1TYaD+hhedDJ3r54HyA6sExp3HfXq7QWEEY/xMw=
|
||||
github.com/go-chi/chi/v5 v5.1.0/go.mod h1:DslCQbL2OYiznFReuXYUmQ2hGd1aDpCnlMNITLSKoi8=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.1 h1:OuVbFODueb089Lh128TAcimifWaLhJwVflnrgM17wHk=
|
||||
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
|
||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||
github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a h1:bbPeKD0xmW/Y25WS6cokEszi5g+S0QxI/d45PkRi7Nk=
|
||||
github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||
github.com/jackc/pgx/v5 v5.6.0 h1:SWJzexBzPL5jb0GEsrPMLIsi/3jOo7RHlzTjcAeDrPY=
|
||||
github.com/jackc/pgx/v5 v5.6.0/go.mod h1:DNZ/vlrUnhWCoFGxHAG8U2ljioxukquj7utPDgtQdTw=
|
||||
github.com/jackc/puddle/v2 v2.2.1 h1:RhxXJtFG022u4ibrCSMSiu5aOq1i77R3OHKNJj77OAk=
|
||||
github.com/jackc/puddle/v2 v2.2.1/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/redis/go-redis/v9 v9.6.1 h1:HHDteefn6ZkTtY5fGUE8tj8uy85AHk6zP7CpzIAM0y4=
|
||||
github.com/redis/go-redis/v9 v9.6.1/go.mod h1:0C0c6ycQsdpVNQpxb1njEQIqkx5UcsM8FJCQLgE9+RA=
|
||||
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||
github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk=
|
||||
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
|
||||
golang.org/x/crypto v0.17.0 h1:r8bRNjWL3GshPW3gkd+RpvzWrZAwPS49OmTGZ/uhM4k=
|
||||
golang.org/x/crypto v0.17.0/go.mod h1:gCAAfMLgwOJRpTjQ2zCCt2OcSfYMTeZVSRtQlPC7Nq4=
|
||||
golang.org/x/sync v0.1.0 h1:wsuoTGHzEhffawBOhz5CYhcrV4IdKZbEyZjBMuTp12o=
|
||||
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/text v0.14.0 h1:ScX5w1eTa3QqT8oi6+ziP7dTV1S2+ALU0bI+0zXKWiQ=
|
||||
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
@@ -0,0 +1,121 @@
|
||||
// Package admin handles administration endpoints: granting points, viewing stats.
|
||||
package admin
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/cu-points/backend/pkg/response"
|
||||
)
|
||||
|
||||
// Handler holds HTTP handler methods for the admin domain.
|
||||
type Handler struct {
|
||||
service *Service
|
||||
}
|
||||
|
||||
// NewHandler creates a new admin Handler.
|
||||
func NewHandler(service *Service) *Handler {
|
||||
return &Handler{service: service}
|
||||
}
|
||||
|
||||
// grantRequest is the expected JSON body for POST /api/v1/admin/points/grant.
|
||||
type grantRequest struct {
|
||||
UserID string `json:"user_id"`
|
||||
Amount int `json:"amount"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// transactionsResponse is the JSON body returned by GET /api/v1/admin/transactions.
|
||||
type transactionsResponse struct {
|
||||
Transactions []AdminTransaction `json:"transactions"`
|
||||
Total int `json:"total"`
|
||||
}
|
||||
|
||||
// GrantPoints handles POST /api/v1/admin/points/grant.
|
||||
// Accepts {user_id, amount, description}; credits the student's balance.
|
||||
// Requires role=admin.
|
||||
func (h *Handler) GrantPoints(w http.ResponseWriter, r *http.Request) {
|
||||
var req grantRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
response.Error(w, http.StatusBadRequest, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if req.UserID == "" {
|
||||
response.Error(w, http.StatusBadRequest, "user_id is required")
|
||||
return
|
||||
}
|
||||
if req.Amount <= 0 {
|
||||
response.Error(w, http.StatusBadRequest, "amount must be positive")
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.service.GrantPoints(r.Context(), req.UserID, req.Amount, req.Description); err != nil {
|
||||
slog.Error("handler.GrantPoints", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
|
||||
response.JSON(w, http.StatusOK, map[string]string{"status": "ok"})
|
||||
}
|
||||
|
||||
// ListTransactions handles GET /api/v1/admin/transactions.
|
||||
// Returns all transactions in the system (paginated), newest first.
|
||||
// Query params: limit (default 50, max 200), offset (default 0).
|
||||
// Response: { "transactions": [...], "total": N }
|
||||
func (h *Handler) ListTransactions(w http.ResponseWriter, r *http.Request) {
|
||||
limit := 50
|
||||
offset := 0
|
||||
if v := r.URL.Query().Get("limit"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil && n > 0 && n <= 200 {
|
||||
limit = n
|
||||
}
|
||||
}
|
||||
if v := r.URL.Query().Get("offset"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil && n >= 0 {
|
||||
offset = n
|
||||
}
|
||||
}
|
||||
|
||||
txs, total, err := h.service.ListTransactions(r.Context(), limit, offset)
|
||||
if err != nil {
|
||||
slog.Error("handler.ListTransactions", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
if txs == nil {
|
||||
txs = []AdminTransaction{}
|
||||
}
|
||||
response.JSON(w, http.StatusOK, transactionsResponse{
|
||||
Transactions: txs,
|
||||
Total: total,
|
||||
})
|
||||
}
|
||||
|
||||
// ListUsers handles GET /api/v1/admin/users.
|
||||
// Returns all students with their current balances.
|
||||
func (h *Handler) ListUsers(w http.ResponseWriter, r *http.Request) {
|
||||
students, err := h.service.ListStudents(r.Context())
|
||||
if err != nil {
|
||||
slog.Error("handler.ListUsers", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
if students == nil {
|
||||
students = []Student{}
|
||||
}
|
||||
response.JSON(w, http.StatusOK, students)
|
||||
}
|
||||
|
||||
// Stats handles GET /api/v1/admin/stats.
|
||||
// Returns aggregated statistics: total students, points issued/spent, active partners.
|
||||
func (h *Handler) Stats(w http.ResponseWriter, r *http.Request) {
|
||||
stats, err := h.service.GetStats(r.Context())
|
||||
if err != nil {
|
||||
slog.Error("handler.Stats", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
response.JSON(w, http.StatusOK, stats)
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/cu-points/backend/internal/points"
|
||||
)
|
||||
|
||||
// Service handles business logic for administrative operations.
|
||||
type Service struct {
|
||||
db *pgxpool.Pool
|
||||
points *points.Service
|
||||
}
|
||||
|
||||
// NewService creates a new admin Service.
|
||||
// pointsSvc is used for all balance mutations so that earn logic is not duplicated here.
|
||||
func NewService(db *pgxpool.Pool, pointsSvc *points.Service) *Service {
|
||||
return &Service{db: db, points: pointsSvc}
|
||||
}
|
||||
|
||||
// AdminTransaction is a transaction record as seen by an administrator.
|
||||
// It includes the associated user email for quick identification.
|
||||
type AdminTransaction struct {
|
||||
ID string `json:"id"`
|
||||
UserID string `json:"user_id"`
|
||||
UserEmail string `json:"user_email"`
|
||||
PartnerID string `json:"partner_id,omitempty"`
|
||||
Amount int `json:"amount"`
|
||||
Type string `json:"type"`
|
||||
Description string `json:"description,omitempty"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
}
|
||||
|
||||
// Student is a user record as seen by an administrator.
|
||||
type Student struct {
|
||||
ID string `json:"id"`
|
||||
Email string `json:"email"`
|
||||
Name string `json:"name"`
|
||||
StudentID string `json:"student_id,omitempty"`
|
||||
Balance int `json:"balance"`
|
||||
}
|
||||
|
||||
// Stats holds aggregated system metrics shown on the admin dashboard.
|
||||
type Stats struct {
|
||||
TotalStudents int `json:"total_students"`
|
||||
TotalPointsIssued int `json:"total_points_issued"`
|
||||
TotalPointsSpent int `json:"total_points_spent"`
|
||||
ActivePartners int `json:"active_partners"`
|
||||
}
|
||||
|
||||
// GrantPoints credits the given amount to the student's balance and records an
|
||||
// admin_grant transaction. Delegates to points.Service.EarnPoints so that all
|
||||
// balance mutation logic lives in one place.
|
||||
func (s *Service) GrantPoints(ctx context.Context, userID string, amount int, description string) error {
|
||||
return s.points.EarnPoints(ctx, points.EarnRequest{
|
||||
UserID: userID,
|
||||
Amount: amount,
|
||||
Type: "admin_grant",
|
||||
Description: description,
|
||||
})
|
||||
}
|
||||
|
||||
// ListTransactions returns a paginated slice of all transactions in the system
|
||||
// (newest first) and the total row count for pagination metadata.
|
||||
func (s *Service) ListTransactions(ctx context.Context, limit, offset int) ([]AdminTransaction, int, error) {
|
||||
var total int
|
||||
err := s.db.QueryRow(ctx, `SELECT COUNT(*) FROM transactions`).Scan(&total)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("service.ListTransactions: count: %w", err)
|
||||
}
|
||||
|
||||
rows, err := s.db.Query(ctx, `
|
||||
SELECT t.id,
|
||||
t.user_id,
|
||||
u.email,
|
||||
COALESCE(t.partner_id::text, ''),
|
||||
t.amount,
|
||||
t.type,
|
||||
COALESCE(t.description, ''),
|
||||
t.created_at
|
||||
FROM transactions t
|
||||
JOIN users u ON u.id = t.user_id
|
||||
ORDER BY t.created_at DESC
|
||||
LIMIT $1 OFFSET $2`,
|
||||
limit, offset,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("service.ListTransactions: query: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var txs []AdminTransaction
|
||||
for rows.Next() {
|
||||
var t AdminTransaction
|
||||
if err := rows.Scan(&t.ID, &t.UserID, &t.UserEmail, &t.PartnerID,
|
||||
&t.Amount, &t.Type, &t.Description, &t.CreatedAt); err != nil {
|
||||
return nil, 0, fmt.Errorf("service.ListTransactions: scan: %w", err)
|
||||
}
|
||||
txs = append(txs, t)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, 0, fmt.Errorf("service.ListTransactions: rows: %w", err)
|
||||
}
|
||||
return txs, total, nil
|
||||
}
|
||||
|
||||
// ListStudents returns all users with role=student, ordered by name.
|
||||
func (s *Service) ListStudents(ctx context.Context) ([]Student, error) {
|
||||
rows, err := s.db.Query(ctx,
|
||||
`SELECT id, email, name, COALESCE(student_id, ''), balance
|
||||
FROM users
|
||||
WHERE role = 'student'
|
||||
ORDER BY name`,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("service.ListStudents: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var students []Student
|
||||
for rows.Next() {
|
||||
var st Student
|
||||
if err := rows.Scan(&st.ID, &st.Email, &st.Name, &st.StudentID, &st.Balance); err != nil {
|
||||
return nil, fmt.Errorf("service.ListStudents: scan: %w", err)
|
||||
}
|
||||
students = append(students, st)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("service.ListStudents: rows: %w", err)
|
||||
}
|
||||
return students, nil
|
||||
}
|
||||
|
||||
// GetStats returns aggregated system statistics for the admin dashboard.
|
||||
func (s *Service) GetStats(ctx context.Context) (*Stats, error) {
|
||||
var stats Stats
|
||||
|
||||
// Single query for all transaction aggregates.
|
||||
err := s.db.QueryRow(ctx, `
|
||||
SELECT
|
||||
COALESCE(SUM(amount) FILTER (WHERE amount > 0), 0),
|
||||
COALESCE(ABS(SUM(amount) FILTER (WHERE amount < 0)), 0)
|
||||
FROM transactions`,
|
||||
).Scan(&stats.TotalPointsIssued, &stats.TotalPointsSpent)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("service.GetStats: transaction aggregates: %w", err)
|
||||
}
|
||||
|
||||
err = s.db.QueryRow(ctx,
|
||||
`SELECT COUNT(*) FROM users WHERE role = 'student'`,
|
||||
).Scan(&stats.TotalStudents)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("service.GetStats: total students: %w", err)
|
||||
}
|
||||
|
||||
err = s.db.QueryRow(ctx,
|
||||
`SELECT COUNT(*) FROM partners WHERE is_active = true`,
|
||||
).Scan(&stats.ActivePartners)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("service.GetStats: active partners: %w", err)
|
||||
}
|
||||
|
||||
return &stats, nil
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
// Package auth handles user authentication: login and token refresh.
|
||||
package auth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/cu-points/backend/pkg/response"
|
||||
)
|
||||
|
||||
// Handler holds HTTP handler methods for the auth domain.
|
||||
// It only parses requests and writes responses — no business logic here.
|
||||
type Handler struct {
|
||||
service *Service
|
||||
}
|
||||
|
||||
// NewHandler creates a new auth Handler backed by the given service.
|
||||
func NewHandler(service *Service) *Handler {
|
||||
return &Handler{service: service}
|
||||
}
|
||||
|
||||
// loginRequest is the expected JSON body for POST /api/v1/auth/login.
|
||||
type loginRequest struct {
|
||||
Email string `json:"email"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
// refreshRequest is the expected JSON body for POST /api/v1/auth/refresh.
|
||||
type refreshRequest struct {
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
}
|
||||
|
||||
// accessTokenResponse is the JSON body returned by a successful refresh.
|
||||
type accessTokenResponse struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
}
|
||||
|
||||
// Login handles POST /api/v1/auth/login.
|
||||
// Accepts {"email": "...", "password": "..."}. Returns a token pair on success.
|
||||
func (h *Handler) Login(w http.ResponseWriter, r *http.Request) {
|
||||
var req loginRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
response.Error(w, http.StatusBadRequest, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if req.Email == "" || req.Password == "" {
|
||||
response.Error(w, http.StatusBadRequest, "email and password are required")
|
||||
return
|
||||
}
|
||||
|
||||
pair, err := h.service.Login(r.Context(), LoginRequest{
|
||||
Email: req.Email,
|
||||
Password: req.Password,
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrInvalidCredentials) {
|
||||
response.Error(w, http.StatusUnauthorized, "invalid email or password")
|
||||
return
|
||||
}
|
||||
slog.Error("handler.Login", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
|
||||
response.JSON(w, http.StatusOK, pair)
|
||||
}
|
||||
|
||||
// Refresh handles POST /api/v1/auth/refresh.
|
||||
// Accepts {"refresh_token": "..."}. Returns a new access token on success.
|
||||
func (h *Handler) Refresh(w http.ResponseWriter, r *http.Request) {
|
||||
var req refreshRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
response.Error(w, http.StatusBadRequest, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if req.RefreshToken == "" {
|
||||
response.Error(w, http.StatusBadRequest, "refresh_token is required")
|
||||
return
|
||||
}
|
||||
|
||||
accessToken, err := h.service.Refresh(r.Context(), req.RefreshToken)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrInvalidCredentials) {
|
||||
response.Error(w, http.StatusUnauthorized, "invalid or expired refresh token")
|
||||
return
|
||||
}
|
||||
slog.Error("handler.Refresh", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
|
||||
response.JSON(w, http.StatusOK, accessTokenResponse{AccessToken: accessToken})
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// Claims are the JWT payload fields used by this service.
|
||||
// Both access and refresh tokens use this struct; the Type field distinguishes them.
|
||||
// Every token includes a unique JWTID (jti) for future revocation support.
|
||||
type Claims struct {
|
||||
jwt.RegisteredClaims // carries sub (user_id), exp, iat, jti
|
||||
Role string `json:"role,omitempty"` // populated only in access tokens
|
||||
Type string `json:"type"` // "access" or "refresh"
|
||||
}
|
||||
|
||||
// JWTManager generates and validates JWT tokens.
|
||||
type JWTManager struct {
|
||||
secret []byte
|
||||
accessTTL time.Duration
|
||||
refreshTTL time.Duration
|
||||
}
|
||||
|
||||
// NewJWTManager creates a JWTManager with the given HMAC secret and TTL durations.
|
||||
func NewJWTManager(secret string, accessTTL, refreshTTL time.Duration) *JWTManager {
|
||||
return &JWTManager{
|
||||
secret: []byte(secret),
|
||||
accessTTL: accessTTL,
|
||||
refreshTTL: refreshTTL,
|
||||
}
|
||||
}
|
||||
|
||||
// GenerateAccessToken creates a signed HS256 access token for the given user.
|
||||
// Claims include: sub (user_id), role, jti (unique ID), iat, exp.
|
||||
func (m *JWTManager) GenerateAccessToken(userID, role string) (string, error) {
|
||||
jti, err := newJTI()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("jwt.GenerateAccessToken: generate jti: %w", err)
|
||||
}
|
||||
|
||||
claims := Claims{
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: userID,
|
||||
ID: jti,
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(m.accessTTL)),
|
||||
},
|
||||
Role: role,
|
||||
Type: "access",
|
||||
}
|
||||
signed, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(m.secret)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("jwt.GenerateAccessToken: sign: %w", err)
|
||||
}
|
||||
return signed, nil
|
||||
}
|
||||
|
||||
// GenerateRefreshToken creates a signed HS256 refresh token for the given user.
|
||||
// Claims include: sub (user_id), jti (unique ID), iat, exp.
|
||||
// The role is intentionally omitted — it is always re-fetched from the DB on use.
|
||||
func (m *JWTManager) GenerateRefreshToken(userID string) (string, error) {
|
||||
jti, err := newJTI()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("jwt.GenerateRefreshToken: generate jti: %w", err)
|
||||
}
|
||||
|
||||
claims := Claims{
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: userID,
|
||||
ID: jti,
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(m.refreshTTL)),
|
||||
},
|
||||
Type: "refresh",
|
||||
}
|
||||
signed, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(m.secret)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("jwt.GenerateRefreshToken: sign: %w", err)
|
||||
}
|
||||
return signed, nil
|
||||
}
|
||||
|
||||
// ParseToken parses and cryptographically validates a JWT string.
|
||||
// It verifies the HMAC signature and token expiry but does NOT check the Type field —
|
||||
// callers are responsible for asserting the expected type ("access" or "refresh").
|
||||
func (m *JWTManager) ParseToken(tokenString string) (*Claims, error) {
|
||||
token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(t *jwt.Token) (interface{}, error) {
|
||||
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, fmt.Errorf("jwt.ParseToken: unexpected signing method: %v", t.Header["alg"])
|
||||
}
|
||||
return m.secret, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("jwt.ParseToken: %w", err)
|
||||
}
|
||||
|
||||
claims, ok := token.Claims.(*Claims)
|
||||
if !ok || !token.Valid {
|
||||
return nil, fmt.Errorf("jwt.ParseToken: invalid token")
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
// newJTI generates a cryptographically random UUID v4 string for use as a JWT ID.
|
||||
func newJTI() (string, error) {
|
||||
b := make([]byte, 16)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
// Set version 4 and variant bits per RFC 4122
|
||||
b[6] = (b[6] & 0x0f) | 0x40
|
||||
b[8] = (b[8] & 0x3f) | 0x80
|
||||
return fmt.Sprintf("%x-%x-%x-%x-%x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:]), nil
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// ErrNotFound is returned when the requested user does not exist in the database.
|
||||
var ErrNotFound = errors.New("not found")
|
||||
|
||||
// UserRecord is the minimal user row fetched from the database during authentication.
|
||||
type UserRecord struct {
|
||||
ID string
|
||||
Email string
|
||||
PasswordHash string
|
||||
Role string
|
||||
}
|
||||
|
||||
// UserRepository defines the database operations the auth service depends on.
|
||||
// Defined as an interface so unit tests can inject a mock without a real database.
|
||||
type UserRepository interface {
|
||||
// GetUserByEmail returns the user row for the given email address.
|
||||
// Returns ErrNotFound if no user exists with that email.
|
||||
GetUserByEmail(ctx context.Context, email string) (*UserRecord, error)
|
||||
|
||||
// GetUserByID returns the user row for the given primary key.
|
||||
// Returns ErrNotFound if the user has been deleted since the token was issued.
|
||||
GetUserByID(ctx context.Context, id string) (*UserRecord, error)
|
||||
}
|
||||
|
||||
// Repository is the PostgreSQL-backed implementation of UserRepository.
|
||||
type Repository struct {
|
||||
db *pgxpool.Pool
|
||||
}
|
||||
|
||||
// NewRepository creates a new PostgreSQL-backed auth Repository.
|
||||
func NewRepository(db *pgxpool.Pool) *Repository {
|
||||
return &Repository{db: db}
|
||||
}
|
||||
|
||||
// GetUserByEmail fetches the user row needed for password verification.
|
||||
// Returns ErrNotFound if no user exists with that email.
|
||||
func (r *Repository) GetUserByEmail(ctx context.Context, email string) (*UserRecord, error) {
|
||||
var u UserRecord
|
||||
err := r.db.QueryRow(ctx,
|
||||
`SELECT id, email, password_hash, role FROM users WHERE email = $1`,
|
||||
email,
|
||||
).Scan(&u.ID, &u.Email, &u.PasswordHash, &u.Role)
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("repository.GetUserByEmail: %w", err)
|
||||
}
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
// GetUserByID fetches the user row needed when refreshing a token.
|
||||
// The role is re-read from the DB so that admin role changes take effect on the next refresh.
|
||||
// Returns ErrNotFound if the user has been deleted since the token was issued.
|
||||
func (r *Repository) GetUserByID(ctx context.Context, id string) (*UserRecord, error) {
|
||||
var u UserRecord
|
||||
err := r.db.QueryRow(ctx,
|
||||
`SELECT id, email, password_hash, role FROM users WHERE id = $1`,
|
||||
id,
|
||||
).Scan(&u.ID, &u.Email, &u.PasswordHash, &u.Role)
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("repository.GetUserByID: %w", err)
|
||||
}
|
||||
return &u, nil
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
// ErrInvalidCredentials is returned for both an unknown email and a wrong password.
|
||||
// Using a single sentinel prevents callers from distinguishing the two cases,
|
||||
// which would otherwise allow email enumeration.
|
||||
var ErrInvalidCredentials = errors.New("invalid email or password")
|
||||
|
||||
// Service contains business logic for authentication.
|
||||
// All password and token operations live here; the handler only parses HTTP.
|
||||
type Service struct {
|
||||
repo UserRepository
|
||||
jwt *JWTManager
|
||||
}
|
||||
|
||||
// NewService creates a new auth Service with the given repository and JWT manager.
|
||||
func NewService(repo UserRepository, jwt *JWTManager) *Service {
|
||||
return &Service{repo: repo, jwt: jwt}
|
||||
}
|
||||
|
||||
// LoginRequest holds credentials submitted by the user on the login form.
|
||||
type LoginRequest struct {
|
||||
Email string
|
||||
Password string
|
||||
}
|
||||
|
||||
// TokenPair holds the access and refresh tokens returned after a successful login.
|
||||
type TokenPair struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
}
|
||||
|
||||
// Login validates credentials and returns a JWT token pair on success.
|
||||
// Returns ErrInvalidCredentials for both an unknown email and a wrong password
|
||||
// so callers cannot distinguish between the two cases (anti-enumeration).
|
||||
func (s *Service) Login(ctx context.Context, req LoginRequest) (*TokenPair, error) {
|
||||
user, err := s.repo.GetUserByEmail(ctx, req.Email)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrNotFound) {
|
||||
// Run a dummy bcrypt comparison so that response time is constant
|
||||
// regardless of whether the email exists in the database.
|
||||
bcrypt.CompareHashAndPassword([]byte("$2a$10$dummyhashpadding000000000000000000000000000000000000000"), []byte(req.Password)) //nolint:errcheck
|
||||
return nil, ErrInvalidCredentials
|
||||
}
|
||||
return nil, fmt.Errorf("service.Login: %w", err)
|
||||
}
|
||||
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.Password)); err != nil {
|
||||
return nil, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
accessToken, err := s.jwt.GenerateAccessToken(user.ID, user.Role)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("service.Login: %w", err)
|
||||
}
|
||||
|
||||
refreshToken, err := s.jwt.GenerateRefreshToken(user.ID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("service.Login: %w", err)
|
||||
}
|
||||
|
||||
slog.Info("user logged in", "user_id", user.ID, "role", user.Role)
|
||||
|
||||
return &TokenPair{
|
||||
AccessToken: accessToken,
|
||||
RefreshToken: refreshToken,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Refresh validates a refresh token and returns a new access token.
|
||||
// The user's role is re-fetched from the database so that role changes take effect immediately
|
||||
// rather than persisting until the old refresh token expires.
|
||||
func (s *Service) Refresh(ctx context.Context, refreshToken string) (string, error) {
|
||||
claims, err := s.jwt.ParseToken(refreshToken)
|
||||
if err != nil {
|
||||
return "", ErrInvalidCredentials
|
||||
}
|
||||
if claims.Type != "refresh" {
|
||||
return "", ErrInvalidCredentials
|
||||
}
|
||||
|
||||
user, err := s.repo.GetUserByID(ctx, claims.Subject)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrNotFound) {
|
||||
return "", ErrInvalidCredentials
|
||||
}
|
||||
return "", fmt.Errorf("service.Refresh: %w", err)
|
||||
}
|
||||
|
||||
accessToken, err := s.jwt.GenerateAccessToken(user.ID, user.Role)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("service.Refresh: %w", err)
|
||||
}
|
||||
|
||||
return accessToken, nil
|
||||
}
|
||||
|
||||
// ValidateToken parses an access token and returns its claims.
|
||||
// Returns ErrInvalidCredentials if the token is invalid, expired, or not an access token.
|
||||
// Used by other services that need to inspect token claims (e.g. extracting user_id).
|
||||
func (s *Service) ValidateToken(token string) (*Claims, error) {
|
||||
claims, err := s.jwt.ParseToken(token)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("service.ValidateToken: %w", err)
|
||||
}
|
||||
if claims.Type != "access" {
|
||||
return nil, fmt.Errorf("service.ValidateToken: %w", ErrInvalidCredentials)
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package auth_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"github.com/cu-points/backend/internal/auth"
|
||||
)
|
||||
|
||||
// mockRepo is a test double for UserRepository.
|
||||
// Populate user and/or repoErr before each test case.
|
||||
type mockRepo struct {
|
||||
user *auth.UserRecord
|
||||
repoErr error
|
||||
}
|
||||
|
||||
func (m *mockRepo) GetUserByEmail(_ context.Context, _ string) (*auth.UserRecord, error) {
|
||||
return m.user, m.repoErr
|
||||
}
|
||||
|
||||
func (m *mockRepo) GetUserByID(_ context.Context, _ string) (*auth.UserRecord, error) {
|
||||
return m.user, m.repoErr
|
||||
}
|
||||
|
||||
// hashPassword hashes the given plain-text password using bcrypt minimum cost for speed.
|
||||
func hashPassword(t *testing.T, password string) string {
|
||||
t.Helper()
|
||||
h, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.MinCost)
|
||||
if err != nil {
|
||||
t.Fatalf("hashPassword: %v", err)
|
||||
}
|
||||
return string(h)
|
||||
}
|
||||
|
||||
// newTestService builds a Service wired to the given mock repo.
|
||||
func newTestService(repo auth.UserRepository) *auth.Service {
|
||||
jwtMgr := auth.NewJWTManager(
|
||||
"test-secret-minimum-32-characters-long",
|
||||
15*time.Minute,
|
||||
168*time.Hour,
|
||||
)
|
||||
return auth.NewService(repo, jwtMgr)
|
||||
}
|
||||
|
||||
func TestService_Login_Success(t *testing.T) {
|
||||
repo := &mockRepo{
|
||||
user: &auth.UserRecord{
|
||||
ID: "a5b66288-4a97-410b-9e30-a7cf61cdabab",
|
||||
Email: "student@cu.ru",
|
||||
PasswordHash: hashPassword(t, "password123"),
|
||||
Role: "student",
|
||||
},
|
||||
}
|
||||
svc := newTestService(repo)
|
||||
|
||||
pair, err := svc.Login(context.Background(), auth.LoginRequest{
|
||||
Email: "student@cu.ru",
|
||||
Password: "password123",
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got: %v", err)
|
||||
}
|
||||
if pair.AccessToken == "" {
|
||||
t.Error("expected non-empty access token")
|
||||
}
|
||||
if pair.RefreshToken == "" {
|
||||
t.Error("expected non-empty refresh token")
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_Login_WrongPassword(t *testing.T) {
|
||||
repo := &mockRepo{
|
||||
user: &auth.UserRecord{
|
||||
ID: "a5b66288-4a97-410b-9e30-a7cf61cdabab",
|
||||
Email: "student@cu.ru",
|
||||
PasswordHash: hashPassword(t, "password123"),
|
||||
Role: "student",
|
||||
},
|
||||
}
|
||||
svc := newTestService(repo)
|
||||
|
||||
_, err := svc.Login(context.Background(), auth.LoginRequest{
|
||||
Email: "student@cu.ru",
|
||||
Password: "wrongpassword",
|
||||
})
|
||||
|
||||
if !errors.Is(err, auth.ErrInvalidCredentials) {
|
||||
t.Errorf("expected ErrInvalidCredentials, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_Login_UserNotFound(t *testing.T) {
|
||||
repo := &mockRepo{repoErr: auth.ErrNotFound}
|
||||
svc := newTestService(repo)
|
||||
|
||||
_, err := svc.Login(context.Background(), auth.LoginRequest{
|
||||
Email: "nobody@cu.ru",
|
||||
Password: "password123",
|
||||
})
|
||||
|
||||
if !errors.Is(err, auth.ErrInvalidCredentials) {
|
||||
t.Errorf("expected ErrInvalidCredentials, got: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
// Package config loads all runtime configuration from environment variables.
|
||||
// All other packages must read settings through this package — no magic strings elsewhere.
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Config holds all application settings read from environment variables.
|
||||
type Config struct {
|
||||
DatabaseURL string
|
||||
RedisURL string
|
||||
JWTSecret string
|
||||
JWTAccessTTL time.Duration
|
||||
JWTRefreshTTL time.Duration
|
||||
Port string
|
||||
Env string
|
||||
}
|
||||
|
||||
// Load reads environment variables and returns a populated Config.
|
||||
// Returns an error if any required variable is missing or cannot be parsed.
|
||||
func Load() (*Config, error) {
|
||||
cfg := &Config{
|
||||
DatabaseURL: os.Getenv("DATABASE_URL"),
|
||||
RedisURL: os.Getenv("REDIS_URL"),
|
||||
JWTSecret: os.Getenv("JWT_SECRET"),
|
||||
Port: envOrDefault("PORT", "8080"),
|
||||
Env: envOrDefault("ENV", "development"),
|
||||
}
|
||||
|
||||
for _, req := range []struct{ name, val string }{
|
||||
{"DATABASE_URL", cfg.DatabaseURL},
|
||||
{"REDIS_URL", cfg.RedisURL},
|
||||
{"JWT_SECRET", cfg.JWTSecret},
|
||||
} {
|
||||
if req.val == "" {
|
||||
return nil, fmt.Errorf("config.Load: required env var %s is not set", req.name)
|
||||
}
|
||||
}
|
||||
|
||||
var err error
|
||||
cfg.JWTAccessTTL, err = time.ParseDuration(envOrDefault("JWT_ACCESS_TTL", "15m"))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("config.Load: invalid JWT_ACCESS_TTL: %w", err)
|
||||
}
|
||||
cfg.JWTRefreshTTL, err = time.ParseDuration(envOrDefault("JWT_REFRESH_TTL", "168h"))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("config.Load: invalid JWT_REFRESH_TTL: %w", err)
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func envOrDefault(key, fallback string) string {
|
||||
if v := os.Getenv(key); v != "" {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
// Package middleware provides HTTP middleware: JWT verification, role guard, request logging.
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// contextKey is an unexported type for context keys in this package,
|
||||
// preventing collisions with keys set by other packages.
|
||||
type contextKey string
|
||||
|
||||
const (
|
||||
userIDKey contextKey = "user_id"
|
||||
userRoleKey contextKey = "user_role"
|
||||
jtiKey contextKey = "jti"
|
||||
)
|
||||
|
||||
// tokenClaims mirrors the JWT payload fields the middleware needs to inspect.
|
||||
// Defined locally so the middleware does not import the auth package.
|
||||
type tokenClaims struct {
|
||||
jwt.RegisteredClaims // provides Subject (user_id), JWTID (jti), ExpiresAt
|
||||
Role string `json:"role"`
|
||||
Type string `json:"type"`
|
||||
}
|
||||
|
||||
// UserIDFromContext retrieves the authenticated user's ID stored by Auth middleware.
|
||||
func UserIDFromContext(ctx context.Context) string {
|
||||
v, _ := ctx.Value(userIDKey).(string)
|
||||
return v
|
||||
}
|
||||
|
||||
// UserRoleFromContext retrieves the authenticated user's role stored by Auth middleware.
|
||||
func UserRoleFromContext(ctx context.Context) string {
|
||||
v, _ := ctx.Value(userRoleKey).(string)
|
||||
return v
|
||||
}
|
||||
|
||||
// JTIFromContext retrieves the JWT ID (jti) stored by Auth middleware.
|
||||
// Useful for token revocation checks in downstream handlers.
|
||||
func JTIFromContext(ctx context.Context) string {
|
||||
v, _ := ctx.Value(jtiKey).(string)
|
||||
return v
|
||||
}
|
||||
|
||||
// Auth returns middleware that validates the Bearer JWT in the Authorization header.
|
||||
// On success it injects user_id, role, and jti into the request context.
|
||||
// Rejects tokens that are expired, have a bad signature, or are not of type "access"
|
||||
// (prevents refresh tokens from being used on protected endpoints).
|
||||
func Auth(secret string) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
tokenStr, err := bearerToken(r)
|
||||
if err != nil {
|
||||
http.Error(w, `{"error":"missing or invalid Authorization header"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
claims := &tokenClaims{}
|
||||
token, err := jwt.ParseWithClaims(tokenStr, claims, func(t *jwt.Token) (interface{}, error) {
|
||||
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, fmt.Errorf("middleware.Auth: unexpected signing method: %v", t.Header["alg"])
|
||||
}
|
||||
return []byte(secret), nil
|
||||
})
|
||||
if err != nil || !token.Valid {
|
||||
http.Error(w, `{"error":"invalid or expired token"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
// Explicitly block refresh tokens from reaching protected endpoints.
|
||||
if claims.Type != "access" {
|
||||
http.Error(w, `{"error":"access token required"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
ctx := context.WithValue(r.Context(), userIDKey, claims.Subject)
|
||||
ctx = context.WithValue(ctx, userRoleKey, claims.Role)
|
||||
ctx = context.WithValue(ctx, jtiKey, claims.ID)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// bearerToken extracts the token string from the Authorization: Bearer <token> header.
|
||||
func bearerToken(r *http.Request) (string, error) {
|
||||
h := r.Header.Get("Authorization")
|
||||
if !strings.HasPrefix(h, "Bearer ") {
|
||||
return "", fmt.Errorf("middleware.bearerToken: missing Bearer prefix")
|
||||
}
|
||||
return strings.TrimPrefix(h, "Bearer "), nil
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package middleware
|
||||
|
||||
import "net/http"
|
||||
|
||||
// CORS adds permissive CORS headers for local development.
|
||||
// Handles the browser preflight OPTIONS request so chi doesn't return 405.
|
||||
func CORS(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Access-Control-Allow-Origin", r.Header.Get("Origin"))
|
||||
w.Header().Set("Access-Control-Allow-Credentials", "true")
|
||||
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, PATCH, DELETE, OPTIONS")
|
||||
w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type")
|
||||
|
||||
if r.Method == http.MethodOptions {
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Logger is structured request-logging middleware using log/slog.
|
||||
// It records method, path, status code, and response duration for every request.
|
||||
func Logger(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
start := time.Now()
|
||||
wrapped := &responseWriter{ResponseWriter: w, status: http.StatusOK}
|
||||
next.ServeHTTP(wrapped, r)
|
||||
slog.Info("request",
|
||||
"method", r.Method,
|
||||
"path", r.URL.Path,
|
||||
"status", wrapped.status,
|
||||
"duration_ms", time.Since(start).Milliseconds(),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
// responseWriter wraps http.ResponseWriter to capture the status code.
|
||||
type responseWriter struct {
|
||||
http.ResponseWriter
|
||||
status int
|
||||
}
|
||||
|
||||
func (rw *responseWriter) WriteHeader(code int) {
|
||||
rw.status = code
|
||||
rw.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// SpendRateLimit returns middleware that limits requests to maxPerMinute per
|
||||
// authenticated partner (identified by user_id in context). It must be placed
|
||||
// after Auth + RequireRole("partner") so that UserIDFromContext is populated.
|
||||
//
|
||||
// Implementation: Redis INCR + EXPIRE sliding-window counter.
|
||||
// Key: rate_limit:spend:<partner_id> — expires after 1 minute.
|
||||
// On Redis failure the middleware fails open (lets the request through) so that
|
||||
// a Redis outage does not take down point transactions.
|
||||
func SpendRateLimit(rdb *redis.Client, maxPerMinute int) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
partnerID := UserIDFromContext(r.Context())
|
||||
key := "rate_limit:spend:" + partnerID
|
||||
|
||||
// Use a short-lived context for Redis so a slow Redis doesn't stall the request.
|
||||
rCtx, cancel := context.WithTimeout(r.Context(), 200*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
count, err := rdb.Incr(rCtx, key).Result()
|
||||
if err != nil {
|
||||
// Fail open: Redis unavailable should not block transactions.
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
// Set the expiry only on the first increment so the window resets each minute.
|
||||
if count == 1 {
|
||||
rdb.Expire(rCtx, key, time.Minute) //nolint:errcheck
|
||||
}
|
||||
if count > int64(maxPerMinute) {
|
||||
http.Error(w, `{"error":"rate limit exceeded, max 10 requests per minute"}`, http.StatusTooManyRequests)
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package middleware
|
||||
|
||||
import "net/http"
|
||||
|
||||
// RequireRole returns middleware that allows only requests whose authenticated user
|
||||
// holds one of the permitted roles. Must be chained after Auth middleware.
|
||||
func RequireRole(allowed ...string) func(http.Handler) http.Handler {
|
||||
set := make(map[string]struct{}, len(allowed))
|
||||
for _, r := range allowed {
|
||||
set[r] = struct{}{}
|
||||
}
|
||||
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if _, ok := set[UserRoleFromContext(r.Context())]; !ok {
|
||||
http.Error(w, `{"error":"forbidden"}`, http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
// Package partners handles the public partner listing endpoint.
|
||||
package partners
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/cu-points/backend/pkg/response"
|
||||
)
|
||||
|
||||
// Handler holds HTTP handler methods for the partners domain.
|
||||
type Handler struct {
|
||||
service *Service
|
||||
}
|
||||
|
||||
// NewHandler creates a new partners Handler.
|
||||
func NewHandler(service *Service) *Handler {
|
||||
return &Handler{service: service}
|
||||
}
|
||||
|
||||
// List handles GET /api/v1/partners.
|
||||
// Returns all active partners. This endpoint is publicly accessible (no auth required).
|
||||
func (h *Handler) List(w http.ResponseWriter, r *http.Request) {
|
||||
partners, err := h.service.ListActive(r.Context())
|
||||
if err != nil {
|
||||
slog.Error("handler.List", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
// Return an empty array rather than null when there are no partners.
|
||||
if partners == nil {
|
||||
partners = []Partner{}
|
||||
}
|
||||
response.JSON(w, http.StatusOK, partners)
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package partners
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// ErrNotFound is returned when the requested partner does not exist.
|
||||
var ErrNotFound = errors.New("partner not found")
|
||||
|
||||
// Repository handles database access for the partners domain.
|
||||
type Repository struct {
|
||||
db *pgxpool.Pool
|
||||
}
|
||||
|
||||
// NewRepository creates a new partners Repository.
|
||||
func NewRepository(db *pgxpool.Pool) *Repository {
|
||||
return &Repository{db: db}
|
||||
}
|
||||
|
||||
// ListActive fetches all partners where is_active = true, ordered alphabetically.
|
||||
func (r *Repository) ListActive(ctx context.Context) ([]Partner, error) {
|
||||
rows, err := r.db.Query(ctx,
|
||||
`SELECT id, name, address, max_spend_pct
|
||||
FROM partners
|
||||
WHERE is_active = true
|
||||
ORDER BY name`,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("repository.ListActive: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var ps []Partner
|
||||
for rows.Next() {
|
||||
var p Partner
|
||||
if err := rows.Scan(&p.ID, &p.Name, &p.Address, &p.MaxSpendPct); err != nil {
|
||||
return nil, fmt.Errorf("repository.ListActive: scan: %w", err)
|
||||
}
|
||||
ps = append(ps, p)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("repository.ListActive: rows: %w", err)
|
||||
}
|
||||
return ps, nil
|
||||
}
|
||||
|
||||
// GetByID fetches a single partner by its primary key.
|
||||
// Returns ErrNotFound if no partner exists with that ID.
|
||||
func (r *Repository) GetByID(ctx context.Context, id string) (*Partner, error) {
|
||||
var p Partner
|
||||
err := r.db.QueryRow(ctx,
|
||||
`SELECT id, name, address, max_spend_pct
|
||||
FROM partners WHERE id = $1`,
|
||||
id,
|
||||
).Scan(&p.ID, &p.Name, &p.Address, &p.MaxSpendPct)
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("repository.GetByID: %w", err)
|
||||
}
|
||||
return &p, nil
|
||||
}
|
||||
|
||||
// GetByUserID fetches the partner record associated with the given cashier user account.
|
||||
// Returns ErrNotFound if no partner is linked to that user.
|
||||
func (r *Repository) GetByUserID(ctx context.Context, userID string) (*Partner, error) {
|
||||
var p Partner
|
||||
err := r.db.QueryRow(ctx,
|
||||
`SELECT id, name, address, max_spend_pct
|
||||
FROM partners WHERE user_id = $1`,
|
||||
userID,
|
||||
).Scan(&p.ID, &p.Name, &p.Address, &p.MaxSpendPct)
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("repository.GetByUserID: %w", err)
|
||||
}
|
||||
return &p, nil
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package partners
|
||||
|
||||
import "context"
|
||||
|
||||
// Service handles business logic for the partners domain.
|
||||
type Service struct {
|
||||
repo *Repository
|
||||
}
|
||||
|
||||
// NewService creates a new partners Service.
|
||||
func NewService(repo *Repository) *Service {
|
||||
return &Service{repo: repo}
|
||||
}
|
||||
|
||||
// Partner represents a participating business.
|
||||
type Partner struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Address string `json:"address"`
|
||||
MaxSpendPct int `json:"max_spend_pct"`
|
||||
}
|
||||
|
||||
// ListActive returns all partners with is_active = true.
|
||||
func (s *Service) ListActive(ctx context.Context) ([]Partner, error) {
|
||||
return s.repo.ListActive(ctx)
|
||||
}
|
||||
|
||||
// GetByID returns a partner by its primary key.
|
||||
// Returns ErrNotFound (from repository) if the partner does not exist.
|
||||
func (s *Service) GetByID(ctx context.Context, id string) (*Partner, error) {
|
||||
return s.repo.GetByID(ctx, id)
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
// Package points handles earning and spending of loyalty points.
|
||||
package points
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/cu-points/backend/internal/middleware"
|
||||
"github.com/cu-points/backend/pkg/response"
|
||||
)
|
||||
|
||||
// Handler holds HTTP handler methods for the points domain.
|
||||
type Handler struct {
|
||||
service *Service
|
||||
}
|
||||
|
||||
// NewHandler creates a new points Handler.
|
||||
func NewHandler(service *Service) *Handler {
|
||||
return &Handler{service: service}
|
||||
}
|
||||
|
||||
// qrResponse is the JSON body returned by GenerateQR.
|
||||
type qrResponse struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
|
||||
// spendRequest is the expected JSON body for POST /api/v1/partner/spend.
|
||||
type spendRequest struct {
|
||||
QRToken string `json:"qr_token"`
|
||||
Amount int `json:"amount"`
|
||||
}
|
||||
|
||||
// GenerateQR handles GET /api/v1/me/qr.
|
||||
// Returns a one-time QR JWT token with 5-minute TTL for the authenticated student.
|
||||
func (h *Handler) GenerateQR(w http.ResponseWriter, r *http.Request) {
|
||||
userID := middleware.UserIDFromContext(r.Context())
|
||||
|
||||
token, err := h.service.GenerateQRToken(r.Context(), userID)
|
||||
if err != nil {
|
||||
slog.Error("handler.GenerateQR", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
|
||||
response.JSON(w, http.StatusOK, qrResponse{Token: token})
|
||||
}
|
||||
|
||||
// Spend handles POST /api/v1/partner/spend.
|
||||
// Accepts {qr_token, amount}; debits student balance atomically.
|
||||
// Requires role=partner (enforced by the router's RequireRole middleware).
|
||||
func (h *Handler) Spend(w http.ResponseWriter, r *http.Request) {
|
||||
var req spendRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
response.Error(w, http.StatusBadRequest, "invalid JSON body")
|
||||
return
|
||||
}
|
||||
if req.QRToken == "" {
|
||||
response.Error(w, http.StatusBadRequest, "qr_token is required")
|
||||
return
|
||||
}
|
||||
if req.Amount <= 0 {
|
||||
response.Error(w, http.StatusBadRequest, "amount must be positive")
|
||||
return
|
||||
}
|
||||
|
||||
partnerID := middleware.UserIDFromContext(r.Context())
|
||||
|
||||
err := h.service.SpendPoints(r.Context(), SpendRequest{
|
||||
QRToken: req.QRToken,
|
||||
Amount: req.Amount,
|
||||
PartnerID: partnerID,
|
||||
})
|
||||
if err != nil {
|
||||
switch {
|
||||
case errors.Is(err, ErrInvalidQRToken):
|
||||
response.Error(w, http.StatusUnauthorized, "invalid or expired QR token")
|
||||
case errors.Is(err, ErrQRAlreadyUsed):
|
||||
response.Error(w, http.StatusConflict, "QR token has already been used")
|
||||
case errors.Is(err, ErrInsufficientBalance):
|
||||
response.Error(w, http.StatusUnprocessableEntity, "insufficient balance")
|
||||
default:
|
||||
slog.Error("handler.Spend", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
response.JSON(w, http.StatusOK, map[string]string{"status": "ok"})
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
package points
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// Repository defines the database operations needed by the points service.
|
||||
// Using an interface allows the service to be unit-tested with a mock implementation.
|
||||
type Repository interface {
|
||||
// GetBalance fetches the current balance for the given user.
|
||||
GetBalance(ctx context.Context, userID string) (int, error)
|
||||
// EarnAtomic credits amount to the user's balance and inserts a transaction row
|
||||
// in a single DB transaction. txType must be "earn" or "admin_grant".
|
||||
EarnAtomic(ctx context.Context, userID string, amount int, txType, description string) error
|
||||
// SpendAtomic debits amount from user balance and inserts a spend transaction
|
||||
// in a single database transaction. The DB CHECK (balance >= 0) is the
|
||||
// authoritative guard; the service also pre-checks to return ErrInsufficientBalance early.
|
||||
SpendAtomic(ctx context.Context, userID, partnerID string, amount int) error
|
||||
}
|
||||
|
||||
// CacheClient defines the Redis operations needed by the points service.
|
||||
type CacheClient interface {
|
||||
// IsQRUsed returns true if the given jti has already been redeemed.
|
||||
IsQRUsed(ctx context.Context, jti string) (bool, error)
|
||||
// MarkQRUsed records the jti as used with a TTL of 5 minutes.
|
||||
MarkQRUsed(ctx context.Context, jti string) error
|
||||
}
|
||||
|
||||
// pgRepository is the PostgreSQL-backed implementation of Repository.
|
||||
type pgRepository struct {
|
||||
db *pgxpool.Pool
|
||||
}
|
||||
|
||||
// NewRepository creates a new PostgreSQL-backed points Repository.
|
||||
func NewRepository(db *pgxpool.Pool) Repository {
|
||||
return &pgRepository{db: db}
|
||||
}
|
||||
|
||||
// GetBalance returns the current point balance for the given user.
|
||||
func (r *pgRepository) GetBalance(ctx context.Context, userID string) (int, error) {
|
||||
var balance int
|
||||
err := r.db.QueryRow(ctx,
|
||||
`SELECT balance FROM users WHERE id = $1`,
|
||||
userID,
|
||||
).Scan(&balance)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("repository.GetBalance: %w", err)
|
||||
}
|
||||
return balance, nil
|
||||
}
|
||||
|
||||
// EarnAtomic credits amount to the user's balance and records a transaction,
|
||||
// all within a single DB transaction.
|
||||
// txType must be a value accepted by the transactions.type CHECK constraint ("earn" or "admin_grant").
|
||||
func (r *pgRepository) EarnAtomic(ctx context.Context, userID string, amount int, txType, description string) error {
|
||||
tx, err := r.db.Begin(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("repository.EarnAtomic: begin: %w", err)
|
||||
}
|
||||
defer tx.Rollback(ctx) //nolint:errcheck
|
||||
|
||||
_, err = tx.Exec(ctx,
|
||||
`UPDATE users SET balance = balance + $1 WHERE id = $2`,
|
||||
amount, userID,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("repository.EarnAtomic: update balance: %w", err)
|
||||
}
|
||||
|
||||
_, err = tx.Exec(ctx,
|
||||
`INSERT INTO transactions (user_id, amount, type, description)
|
||||
VALUES ($1, $2, $3, $4)`,
|
||||
userID, amount, txType, description,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("repository.EarnAtomic: insert transaction: %w", err)
|
||||
}
|
||||
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
return fmt.Errorf("repository.EarnAtomic: commit: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SpendAtomic deducts amount from the student's balance and records a spend
|
||||
// transaction, all within a single DB transaction.
|
||||
// The negative amount stored in transactions follows the ledger convention:
|
||||
// positive = earn, negative = spend.
|
||||
func (r *pgRepository) SpendAtomic(ctx context.Context, userID, partnerID string, amount int) error {
|
||||
tx, err := r.db.Begin(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("repository.SpendAtomic: begin: %w", err)
|
||||
}
|
||||
defer tx.Rollback(ctx) //nolint:errcheck
|
||||
|
||||
var newBalance int
|
||||
err = tx.QueryRow(ctx,
|
||||
`UPDATE users SET balance = balance - $1 WHERE id = $2 RETURNING balance`,
|
||||
amount, userID,
|
||||
).Scan(&newBalance)
|
||||
if err != nil {
|
||||
return fmt.Errorf("repository.SpendAtomic: update balance: %w", err)
|
||||
}
|
||||
|
||||
_, err = tx.Exec(ctx,
|
||||
`INSERT INTO transactions (user_id, partner_id, amount, type)
|
||||
VALUES ($1, $2, $3, 'spend')`,
|
||||
userID, partnerID, -amount,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("repository.SpendAtomic: insert transaction: %w", err)
|
||||
}
|
||||
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
return fmt.Errorf("repository.SpendAtomic: commit: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// redisCache is the Redis-backed implementation of CacheClient.
|
||||
type redisCache struct {
|
||||
client *redis.Client
|
||||
}
|
||||
|
||||
// NewRedisCache creates a Redis-backed CacheClient for QR token one-time-use tracking.
|
||||
func NewRedisCache(client *redis.Client) CacheClient {
|
||||
return &redisCache{client: client}
|
||||
}
|
||||
|
||||
// IsQRUsed returns true if the given jti key already exists in Redis.
|
||||
func (c *redisCache) IsQRUsed(ctx context.Context, jti string) (bool, error) {
|
||||
count, err := c.client.Exists(ctx, "used_qr:"+jti).Result()
|
||||
return count > 0, err
|
||||
}
|
||||
|
||||
// MarkQRUsed sets used_qr:<jti> = "1" with a 5-minute TTL.
|
||||
// TTL matches the QR token expiry so the key is automatically cleaned up.
|
||||
func (c *redisCache) MarkQRUsed(ctx context.Context, jti string) error {
|
||||
return c.client.Set(ctx, "used_qr:"+jti, "1", 5*time.Minute).Err()
|
||||
}
|
||||
@@ -0,0 +1,185 @@
|
||||
package points
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// ErrInsufficientBalance is returned when a student's balance is lower than the requested spend amount.
|
||||
var ErrInsufficientBalance = errors.New("insufficient balance")
|
||||
|
||||
// ErrQRAlreadyUsed is returned when a QR token has already been redeemed.
|
||||
var ErrQRAlreadyUsed = errors.New("QR token already used")
|
||||
|
||||
// ErrInvalidQRToken is returned when the QR JWT is malformed, expired, or has the wrong type.
|
||||
var ErrInvalidQRToken = errors.New("invalid or expired QR token")
|
||||
|
||||
const qrTokenTTL = 5 * time.Minute
|
||||
|
||||
// qrClaims are the JWT payload fields for one-time QR spend tokens.
|
||||
type qrClaims struct {
|
||||
jwt.RegisteredClaims
|
||||
Type string `json:"type"` // always "qr"
|
||||
}
|
||||
|
||||
// Service contains the critical business logic for points operations.
|
||||
// This is the most important file in the project — all balance mutations live here.
|
||||
type Service struct {
|
||||
repo Repository
|
||||
cache CacheClient
|
||||
secret []byte
|
||||
}
|
||||
|
||||
// NewService creates a new points Service.
|
||||
// secret must be the same HMAC secret used for all JWTs in this application.
|
||||
func NewService(repo Repository, cache CacheClient, secret string) *Service {
|
||||
return &Service{repo: repo, cache: cache, secret: []byte(secret)}
|
||||
}
|
||||
|
||||
// EarnRequest holds the data needed to credit a student's balance.
|
||||
// Type must be "earn" or "admin_grant" — enforced by the DB CHECK constraint.
|
||||
type EarnRequest struct {
|
||||
UserID string
|
||||
Amount int
|
||||
Type string // "earn" or "admin_grant"
|
||||
Description string
|
||||
}
|
||||
|
||||
// EarnPoints credits the given amount to the student's balance and records a
|
||||
// transaction of the specified type atomically. This is the single place where
|
||||
// all credit logic lives — admin grants, future LMS integrations, etc. must
|
||||
// call this method rather than touching the DB directly.
|
||||
func (s *Service) EarnPoints(ctx context.Context, req EarnRequest) error {
|
||||
if req.Amount <= 0 {
|
||||
return fmt.Errorf("service.EarnPoints: amount must be positive")
|
||||
}
|
||||
if err := s.repo.EarnAtomic(ctx, req.UserID, req.Amount, req.Type, req.Description); err != nil {
|
||||
return fmt.Errorf("service.EarnPoints: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SpendRequest holds the data needed to debit a student's balance at a partner.
|
||||
type SpendRequest struct {
|
||||
QRToken string
|
||||
Amount int
|
||||
PartnerID string
|
||||
}
|
||||
|
||||
// SpendPoints debits the given amount from the student's balance
|
||||
// and records a spend transaction atomically in a single DB transaction.
|
||||
// Returns ErrInvalidQRToken if the token is malformed or expired.
|
||||
// Returns ErrQRAlreadyUsed if the QR token has been redeemed before.
|
||||
// Returns ErrInsufficientBalance if balance < amount.
|
||||
func (s *Service) SpendPoints(ctx context.Context, req SpendRequest) error {
|
||||
// Step 1: validate QR JWT and extract student_id and jti.
|
||||
claims, err := s.parseQRToken(req.QRToken)
|
||||
if err != nil {
|
||||
return ErrInvalidQRToken
|
||||
}
|
||||
studentID := claims.Subject
|
||||
jti := claims.ID
|
||||
|
||||
// Step 2: one-time-use check — reject if already redeemed.
|
||||
used, err := s.cache.IsQRUsed(ctx, jti)
|
||||
if err != nil {
|
||||
return fmt.Errorf("service.SpendPoints: cache check: %w", err)
|
||||
}
|
||||
if used {
|
||||
return ErrQRAlreadyUsed
|
||||
}
|
||||
|
||||
// Step 3: pre-check balance for a clear error message before hitting the DB.
|
||||
// The DB CHECK (balance >= 0) is the authoritative guard; this is a fast-fail.
|
||||
balance, err := s.repo.GetBalance(ctx, studentID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("service.SpendPoints: get balance: %w", err)
|
||||
}
|
||||
if balance < req.Amount {
|
||||
return ErrInsufficientBalance
|
||||
}
|
||||
|
||||
// Step 4–5: debit balance and insert spend transaction atomically.
|
||||
// SpendAtomic uses a DB transaction; the balance CHECK constraint is the
|
||||
// last line of defence against concurrent overdrafts.
|
||||
if err := s.repo.SpendAtomic(ctx, studentID, req.PartnerID, req.Amount); err != nil {
|
||||
// Propagate balance constraint violation with a domain error.
|
||||
if isConstraintError(err) {
|
||||
return ErrInsufficientBalance
|
||||
}
|
||||
return fmt.Errorf("service.SpendPoints: spend atomic: %w", err)
|
||||
}
|
||||
|
||||
// Step 6: mark token as used only after the DB commit succeeds.
|
||||
// If MarkQRUsed fails, the spend already committed — log but don't rollback.
|
||||
if err := s.cache.MarkQRUsed(ctx, jti); err != nil {
|
||||
return fmt.Errorf("service.SpendPoints: mark qr used: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GenerateQRToken creates a one-time JWT for the student to present at a partner terminal.
|
||||
// The token encodes the student's user_id and a unique jti; TTL is 5 minutes.
|
||||
func (s *Service) GenerateQRToken(ctx context.Context, userID string) (string, error) {
|
||||
jti, err := newJTI()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("service.GenerateQRToken: generate jti: %w", err)
|
||||
}
|
||||
|
||||
claims := qrClaims{
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: userID,
|
||||
ID: jti,
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(qrTokenTTL)),
|
||||
},
|
||||
Type: "qr",
|
||||
}
|
||||
|
||||
token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(s.secret)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("service.GenerateQRToken: sign: %w", err)
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// parseQRToken validates the JWT signature/expiry and asserts type="qr".
|
||||
func (s *Service) parseQRToken(tokenStr string) (*qrClaims, error) {
|
||||
token, err := jwt.ParseWithClaims(tokenStr, &qrClaims{}, func(t *jwt.Token) (interface{}, error) {
|
||||
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, fmt.Errorf("unexpected signing method: %v", t.Header["alg"])
|
||||
}
|
||||
return s.secret, nil
|
||||
})
|
||||
if err != nil || !token.Valid {
|
||||
return nil, fmt.Errorf("invalid token")
|
||||
}
|
||||
claims, ok := token.Claims.(*qrClaims)
|
||||
if !ok || claims.Type != "qr" {
|
||||
return nil, fmt.Errorf("not a QR token")
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
// newJTI generates a cryptographically random UUID v4 string for use as a JWT ID.
|
||||
func newJTI() (string, error) {
|
||||
b := make([]byte, 16)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
b[6] = (b[6] & 0x0f) | 0x40
|
||||
b[8] = (b[8] & 0x3f) | 0x80
|
||||
return fmt.Sprintf("%x-%x-%x-%x-%x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:]), nil
|
||||
}
|
||||
|
||||
// isConstraintError reports whether err contains a PostgreSQL balance CHECK violation.
|
||||
func isConstraintError(err error) bool {
|
||||
return err != nil && strings.Contains(err.Error(), "check")
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
// Package users handles student profile and transaction history endpoints.
|
||||
package users
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/cu-points/backend/internal/middleware"
|
||||
"github.com/cu-points/backend/pkg/response"
|
||||
)
|
||||
|
||||
// Handler holds HTTP handler methods for the users domain.
|
||||
type Handler struct {
|
||||
service *Service
|
||||
}
|
||||
|
||||
// NewHandler creates a new users Handler.
|
||||
func NewHandler(service *Service) *Handler {
|
||||
return &Handler{service: service}
|
||||
}
|
||||
|
||||
// transactionsResponse is the JSON body returned by GET /me/transactions.
|
||||
type transactionsResponse struct {
|
||||
Transactions []Transaction `json:"transactions"`
|
||||
Total int `json:"total"`
|
||||
}
|
||||
|
||||
// Me handles GET /api/v1/me.
|
||||
// Returns the authenticated student's profile and current balance.
|
||||
// Requires role=student (enforced by the router's RequireRole middleware).
|
||||
func (h *Handler) Me(w http.ResponseWriter, r *http.Request) {
|
||||
userID := middleware.UserIDFromContext(r.Context())
|
||||
|
||||
profile, err := h.service.GetProfile(r.Context(), userID)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrNotFound) {
|
||||
response.Error(w, http.StatusNotFound, "user not found")
|
||||
return
|
||||
}
|
||||
slog.Error("handler.Me", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
|
||||
response.JSON(w, http.StatusOK, profile)
|
||||
}
|
||||
|
||||
// Transactions handles GET /api/v1/me/transactions.
|
||||
// Returns paginated transaction history for the authenticated student.
|
||||
// Query params: limit (default 20, max 100), offset (default 0).
|
||||
// Response: { "transactions": [...], "total": N }
|
||||
func (h *Handler) Transactions(w http.ResponseWriter, r *http.Request) {
|
||||
userID := middleware.UserIDFromContext(r.Context())
|
||||
|
||||
limit := 20
|
||||
offset := 0
|
||||
if v := r.URL.Query().Get("limit"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil && n > 0 && n <= 100 {
|
||||
limit = n
|
||||
}
|
||||
}
|
||||
if v := r.URL.Query().Get("offset"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil && n >= 0 {
|
||||
offset = n
|
||||
}
|
||||
}
|
||||
|
||||
txs, total, err := h.service.GetTransactions(r.Context(), userID, limit, offset)
|
||||
if err != nil {
|
||||
slog.Error("handler.Transactions", "err", err)
|
||||
response.Error(w, http.StatusInternalServerError, "internal server error")
|
||||
return
|
||||
}
|
||||
// Return an empty array rather than null when there are no transactions.
|
||||
if txs == nil {
|
||||
txs = []Transaction{}
|
||||
}
|
||||
response.JSON(w, http.StatusOK, transactionsResponse{
|
||||
Transactions: txs,
|
||||
Total: total,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package users
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// ErrNotFound is returned when the requested user does not exist.
|
||||
var ErrNotFound = errors.New("not found")
|
||||
|
||||
// Repository handles all database access for the users domain.
|
||||
type Repository struct {
|
||||
db *pgxpool.Pool
|
||||
}
|
||||
|
||||
// NewRepository creates a new users Repository.
|
||||
func NewRepository(db *pgxpool.Pool) *Repository {
|
||||
return &Repository{db: db}
|
||||
}
|
||||
|
||||
// GetByID fetches a user's profile by primary key.
|
||||
// Returns ErrNotFound if no user exists with that ID.
|
||||
func (r *Repository) GetByID(ctx context.Context, id string) (*Profile, error) {
|
||||
var p Profile
|
||||
err := r.db.QueryRow(ctx,
|
||||
`SELECT id, email, name, COALESCE(student_id, ''), balance
|
||||
FROM users WHERE id = $1`,
|
||||
id,
|
||||
).Scan(&p.ID, &p.Email, &p.Name, &p.StudentID, &p.Balance)
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("repository.GetByID: %w", err)
|
||||
}
|
||||
return &p, nil
|
||||
}
|
||||
|
||||
// UpdateBalance adds delta to the user's balance within an existing pgx transaction.
|
||||
// delta is positive when earning points, negative when spending.
|
||||
// Returns the new balance after the update via RETURNING, so the service can include
|
||||
// it in the API response without a second query.
|
||||
// The database CHECK (balance >= 0) acts as the last line of defense against overdrafts;
|
||||
// this function will return an error if the constraint fires.
|
||||
func (r *Repository) UpdateBalance(ctx context.Context, tx pgx.Tx, id string, delta int) (int, error) {
|
||||
var newBalance int
|
||||
err := tx.QueryRow(ctx,
|
||||
`UPDATE users SET balance = balance + $1 WHERE id = $2 RETURNING balance`,
|
||||
delta, id,
|
||||
).Scan(&newBalance)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("repository.UpdateBalance: %w", err)
|
||||
}
|
||||
return newBalance, nil
|
||||
}
|
||||
|
||||
// CountTransactions returns the total number of transactions for the given user.
|
||||
// Used alongside ListTransactions to populate pagination metadata.
|
||||
func (r *Repository) CountTransactions(ctx context.Context, userID string) (int, error) {
|
||||
var count int
|
||||
err := r.db.QueryRow(ctx,
|
||||
`SELECT COUNT(*) FROM transactions WHERE user_id = $1`,
|
||||
userID,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("repository.CountTransactions: %w", err)
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// ListTransactions returns paginated transactions for the given user, ordered newest first.
|
||||
func (r *Repository) ListTransactions(ctx context.Context, userID string, limit, offset int) ([]Transaction, error) {
|
||||
rows, err := r.db.Query(ctx, `
|
||||
SELECT id,
|
||||
amount,
|
||||
type,
|
||||
COALESCE(description, ''),
|
||||
COALESCE(partner_id::text, ''),
|
||||
created_at
|
||||
FROM transactions
|
||||
WHERE user_id = $1
|
||||
ORDER BY created_at DESC
|
||||
LIMIT $2 OFFSET $3`,
|
||||
userID, limit, offset,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("repository.ListTransactions: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var txs []Transaction
|
||||
for rows.Next() {
|
||||
var t Transaction
|
||||
if err := rows.Scan(&t.ID, &t.Amount, &t.Type, &t.Description, &t.PartnerID, &t.CreatedAt); err != nil {
|
||||
return nil, fmt.Errorf("repository.ListTransactions: scan: %w", err)
|
||||
}
|
||||
txs = append(txs, t)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("repository.ListTransactions: rows: %w", err)
|
||||
}
|
||||
return txs, nil
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package users
|
||||
|
||||
import "context"
|
||||
|
||||
// Service handles business logic for the users domain.
|
||||
type Service struct {
|
||||
repo *Repository
|
||||
}
|
||||
|
||||
// NewService creates a new users Service.
|
||||
func NewService(repo *Repository) *Service {
|
||||
return &Service{repo: repo}
|
||||
}
|
||||
|
||||
// Profile represents a student's public profile and current balance.
|
||||
type Profile struct {
|
||||
ID string `json:"id"`
|
||||
Email string `json:"email"`
|
||||
Name string `json:"name"`
|
||||
StudentID string `json:"student_id,omitempty"`
|
||||
Balance int `json:"balance"`
|
||||
}
|
||||
|
||||
// Transaction represents a single point-earning or point-spending event.
|
||||
type Transaction struct {
|
||||
ID string `json:"id"`
|
||||
Amount int `json:"amount"`
|
||||
Type string `json:"type"`
|
||||
Description string `json:"description,omitempty"`
|
||||
PartnerID string `json:"partner_id,omitempty"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
}
|
||||
|
||||
// GetProfile returns the profile and current balance for the given user.
|
||||
func (s *Service) GetProfile(ctx context.Context, userID string) (*Profile, error) {
|
||||
return s.repo.GetByID(ctx, userID)
|
||||
}
|
||||
|
||||
// GetTransactions returns a paginated list of transactions for the given user
|
||||
// (newest first) together with the total row count for pagination metadata.
|
||||
func (s *Service) GetTransactions(ctx context.Context, userID string, limit, offset int) ([]Transaction, int, error) {
|
||||
total, err := s.repo.CountTransactions(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
txs, err := s.repo.ListTransactions(ctx, userID, limit, offset)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return txs, total, nil
|
||||
}
|
||||
Vendored
+23
@@ -0,0 +1,23 @@
|
||||
// Package cache provides Redis client initialization.
|
||||
package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// NewClient creates and validates a Redis client using the given REDIS_URL.
|
||||
// Returns an error if the URL cannot be parsed or if the initial PING fails.
|
||||
func NewClient(ctx context.Context, redisURL string) (*redis.Client, error) {
|
||||
opts, err := redis.ParseURL(redisURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cache.NewClient: parse URL: %w", err)
|
||||
}
|
||||
client := redis.NewClient(opts)
|
||||
if err := client.Ping(ctx).Err(); err != nil {
|
||||
return nil, fmt.Errorf("cache.NewClient: ping: %w", err)
|
||||
}
|
||||
return client, nil
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
// Package db provides PostgreSQL connection pool initialization.
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// NewPool creates and validates a pgx connection pool using the given DATABASE_URL.
|
||||
// Returns an error if the pool cannot be created or if the initial ping fails.
|
||||
func NewPool(ctx context.Context, databaseURL string) (*pgxpool.Pool, error) {
|
||||
pool, err := pgxpool.New(ctx, databaseURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("db.NewPool: create pool: %w", err)
|
||||
}
|
||||
if err := pool.Ping(ctx); err != nil {
|
||||
return nil, fmt.Errorf("db.NewPool: ping: %w", err)
|
||||
}
|
||||
return pool, nil
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
// Package response provides helpers for writing consistent JSON API responses.
|
||||
// Every handler must use these helpers — never call json.Encode directly.
|
||||
package response
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
type successBody struct {
|
||||
Data interface{} `json:"data"`
|
||||
}
|
||||
|
||||
type errorBody struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
|
||||
// JSON writes a successful JSON response with the given HTTP status code and data payload.
|
||||
// The payload is wrapped in {"data": ...} to match the API envelope convention.
|
||||
func JSON(w http.ResponseWriter, status int, data interface{}) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
json.NewEncoder(w).Encode(successBody{Data: data}) //nolint:errcheck
|
||||
}
|
||||
|
||||
// Error writes a JSON error response with the given HTTP status code and human-readable message.
|
||||
// The message is wrapped in {"error": ...}.
|
||||
func Error(w http.ResponseWriter, status int, message string) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
json.NewEncoder(w).Encode(errorBody{Error: message}) //nolint:errcheck
|
||||
}
|
||||
Reference in New Issue
Block a user