Go: 电商 API(中):认证与测试
最后更新:2026-08-26
上篇搭建了电商 API 的骨架。中篇的目标——让它安全、可维护、可测试。
Bob 的电商 API 在安全审计中暴露了严重问题:没有认证、没有权限控制、没有请求限流。现在是时候补齐中间件管线了。
1. 你将学到
- JWT Auth 中间件集成
- RBAC 角色权限控制
- 全局中间件管线(Logging + Recovery + Auth + RBAC + RateLimit)
- SQL 数据库迁移(golang-migrate)
- 集成测试(httptest + test DB)
- 分页中间件
2. 故事:安全审计报告
(1) 痛点:"所有接口都没有认证,任何人都能删除商品"
Bob 的 MVP 上线第一天,安全团队发来审计报告:
"严重漏洞:DELETE /api/v1/products/1 不需要认证。任何人都可以删除商品。中等漏洞:没有请求限流,攻击者可以暴力破解登录接口。建议:JWT 认证 + RBAC 权限 + 速率限制。"
Bob 看了看上个版本留下的后门:
GO
// 坏代码:没有认证,任何人都能下单
func (h *OrderHandler) CreateOrder(w http.ResponseWriter, r *http.Request) {
userID := 1 // 固定值!所有订单都算在 user 1 头上
// 业务逻辑...
}
(2) 本课目标:补齐中间件管线
TEXT
📖 仅展示
Request → Recovery → Logging → CORS → Auth → RBAC → RateLimit → Handler
本周计划:
- JWT 认证中间件(替换硬编码 userID)
- RBAC 权限控制(admin / customer 角色)
- SQL 数据库迁移(版本化管理表结构)
- 集成测试(httptest + 临时数据库)
- 分页中间件
3. 完整实现
▶ 示例:JWT 工具包
GO
📖 仅展示
// internal/auth/jwt.go
package auth
import (
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"strings"
"time"
)
type Claims struct {
UserID int `json:"user_id"`
Role string `json:"role"`
Email string `json:"email"`
}
type TokenPair struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
}
type JWTService struct {
secret []byte
accessTTL time.Duration
refreshTTL time.Duration
}
func NewJWTService(secret string) *JWTService {
return &JWTService{
secret: []byte(secret),
accessTTL: 15 * time.Minute,
refreshTTL: 7 * 24 * time.Hour,
}
}
func (s *JWTService) GenerateToken(claims Claims) (string, error) {
header := map[string]string{"alg": "HS256", "typ": "JWT"}
headerJSON, _ := json.Marshal(header)
payload := map[string]interface{}{
"user_id": claims.UserID,
"role": claims.Role,
"email": claims.Email,
"exp": time.Now().Add(s.accessTTL).Unix(),
"iat": time.Now().Unix(),
}
payloadJSON, _ := json.Marshal(payload)
headerEnc := base64URLEncode(headerJSON)
payloadEnc := base64URLEncode(payloadJSON)
message := headerEnc + "." + payloadEnc
mac := hmac.New(sha256.New, s.secret)
mac.Write([]byte(message))
sig := base64URLEncode(mac.Sum(nil))
return message + "." + sig, nil
}
func (s *JWTService) ValidateToken(token string) (*Claims, error) {
parts := strings.Split(token, ".")
if len(parts) != 3 {
return nil, fmt.Errorf("invalid token")
}
message := parts[0] + "." + parts[1]
mac := hmac.New(sha256.New, s.secret)
mac.Write([]byte(message))
expected := base64URLEncode(mac.Sum(nil))
if !hmac.Equal([]byte(parts[2]), []byte(expected)) {
return nil, fmt.Errorf("invalid signature")
}
payloadJSON, err := base64URLDecode(parts[1])
if err != nil {
return nil, err
}
var payload struct {
UserID int `json:"user_id"`
Role string `json:"role"`
Email string `json:"email"`
Exp float64 `json:"exp"`
}
if err := json.Unmarshal(payloadJSON, &payload); err != nil {
return nil, err
}
if time.Now().Unix() > int64(payload.Exp) {
return nil, fmt.Errorf("token expired")
}
return &Claims{
UserID: payload.UserID,
Role: payload.Role,
Email: payload.Email,
}, nil
}
func base64URLEncode(data []byte) string {
return strings.TrimRight(base64.URLEncoding.EncodeToString(data), "=")
}
func base64URLDecode(s string) ([]byte, error) {
switch len(s) % 4 {
case 2:
s += "=="
case 3:
s += "="
}
return base64.URLEncoding.DecodeString(s)
}
▶ 示例:中间件管线
GO
📖 仅展示
// internal/middleware/middleware.go
package middleware
import (
"context"
"encoding/json"
"fmt"
"log"
"net/http"
"strings"
"sync"
"time"
"ecommerce/internal/auth"
)
type contextKey string
const UserClaimsKey contextKey = "user_claims"
// ---------- Middleware 类型 ----------
type Middleware func(http.Handler) http.Handler
func Chain(h http.Handler, mws ...Middleware) http.Handler {
for i := len(mws) - 1; i >= 0; i-- {
h = mws[i](h)
}
return h
}
// ---------- Recovery ----------
func Recovery(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer func() {
if err := recover(); err != nil {
log.Printf("[PANIC] %v", err)
writeError(w, http.StatusInternalServerError, "internal error")
}
}()
next.ServeHTTP(w, r)
})
}
// ---------- Logging ----------
func Logging(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
log.Printf("[%s] %s %s", r.Method, r.URL.Path, r.RemoteAddr)
next.ServeHTTP(w, r)
log.Printf("[%s] %s → %v", r.Method, r.URL.Path, time.Since(start))
})
}
// ---------- CORS ----------
func CORS(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Access-Control-Allow-Origin", "*")
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization")
if r.Method == http.MethodOptions {
w.WriteHeader(http.StatusNoContent)
return
}
next.ServeHTTP(w, r)
})
}
// ---------- Auth(JWT)----------
func Auth(jwtService *auth.JWTService) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
authHeader := r.Header.Get("Authorization")
if !strings.HasPrefix(authHeader, "Bearer ") {
writeError(w, http.StatusUnauthorized, "missing token")
return
}
token := authHeader[7:]
claims, err := jwtService.ValidateToken(token)
if err != nil {
writeError(w, http.StatusUnauthorized, err.Error())
return
}
ctx := context.WithValue(r.Context(), UserClaimsKey, claims)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
}
// ---------- RBAC ----------
func RequireRole(roles ...string) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
claims, ok := r.Context().Value(UserClaimsKey).(*auth.Claims)
if !ok {
writeError(w, http.StatusUnauthorized, "not authenticated")
return
}
for _, role := range roles {
if claims.Role == role {
next.ServeHTTP(w, r)
return
}
}
writeError(w, http.StatusForbidden, "insufficient permissions")
})
}
}
// ---------- 令牌桶限流 ----------
type TokenBucket struct {
mu sync.Mutex
tokens float64
maxTokens float64
refillRate float64
lastRefill time.Time
}
type IPRateLimiter struct {
mu sync.RWMutex
buckets map[string]*TokenBucket
rate float64
burst int
}
var globalLimiter = NewIPRateLimiter(100, 200)
func NewIPRateLimiter(rate float64, burst int) *IPRateLimiter {
return &IPRateLimiter{
buckets: make(map[string]*TokenBucket),
rate: rate,
burst: burst,
}
}
func (rl *IPRateLimiter) getBucket(ip string) *TokenBucket {
rl.mu.Lock()
defer rl.mu.Unlock()
b, ok := rl.buckets[ip]
if !ok {
b = &TokenBucket{
tokens: float64(rl.burst),
maxTokens: float64(rl.burst),
refillRate: rl.rate,
lastRefill: time.Now(),
}
rl.buckets[ip] = b
}
return b
}
func (b *TokenBucket) allow() bool {
b.mu.Lock()
defer b.mu.Unlock()
now := time.Now()
elapsed := now.Sub(b.lastRefill).Seconds()
b.tokens = minFloat(b.tokens+elapsed*b.refillRate, b.maxTokens)
b.lastRefill = now
if b.tokens >= 1 {
b.tokens--
return true
}
return false
}
func minFloat(a, b float64) float64 {
if a < b {
return a
}
return b
}
func RateLimit(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ip := r.RemoteAddr
if !globalLimiter.getBucket(ip).allow() {
writeError(w, http.StatusTooManyRequests, "rate limit exceeded")
return
}
next.ServeHTTP(w, r)
})
}
// ---------- 分页 ----------
type Pagination struct {
Page int `json:"page"`
PerPage int `json:"per_page"`
Offset int `json:"-"`
}
func ParsePagination(r *http.Request) Pagination {
page := 1
perPage := 20
if p := r.URL.Query().Get("page"); p != "" {
if v, err := parseInt(p); err == nil && v > 0 {
page = v
}
}
if pp := r.URL.Query().Get("per_page"); pp != "" {
if v, err := parseInt(pp); err == nil && v > 0 && v <= 100 {
perPage = v
}
}
return Pagination{
Page: page,
PerPage: perPage,
Offset: (page - 1) * perPage,
}
}
func parseInt(s string) (int, error) {
var n int
for _, c := range s {
if c < '0' || c > '9' {
return 0, fmt.Errorf("not a number")
}
n = n*10 + int(c-'0')
}
return n, nil
}
// ---------- 工具 ----------
func writeError(w http.ResponseWriter, status int, msg string) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
json.NewEncoder(w).Encode(map[string]string{"error": msg})
}
▶ 示例:SQL 数据库迁移
GO
📖 仅展示
// internal/migration/migration.go
package migration
import (
"database/sql"
"fmt"
"log"
)
type Migration struct {
Version int
SQL string
}
var migrations = []Migration{
{
Version: 1,
SQL: `CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
email TEXT UNIQUE NOT NULL,
password TEXT NOT NULL,
name TEXT NOT NULL,
role TEXT NOT NULL DEFAULT 'customer',
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
)`,
},
{
Version: 2,
SQL: `CREATE TABLE IF NOT EXISTS products (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL,
description TEXT,
price REAL NOT NULL,
stock INTEGER NOT NULL DEFAULT 0,
category_id INTEGER,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
FOREIGN KEY (category_id) REFERENCES categories(id)
)`,
},
{
Version: 3,
SQL: `CREATE TABLE IF NOT EXISTS orders (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
product_id INTEGER NOT NULL,
quantity INTEGER NOT NULL,
total_price REAL NOT NULL,
status TEXT NOT NULL DEFAULT 'pending',
created_at DATETIME DEFAULT CURRENT_TIMESTAMP,
FOREIGN KEY (user_id) REFERENCES users(id),
FOREIGN KEY (product_id) REFERENCES products(id)
)`,
},
{
Version: 4,
SQL: `CREATE TABLE IF NOT EXISTS categories (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT UNIQUE NOT NULL,
description TEXT,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
)`,
},
}
func RunMigrations(db *sql.DB) error {
// 创建版本表
_, err := db.Exec(`CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
applied_at DATETIME DEFAULT CURRENT_TIMESTAMP
)`)
if err != nil {
return fmt.Errorf("create migrations table: %w", err)
}
for _, m := range migrations {
// 检查是否已经应用
var exists bool
err := db.QueryRow("SELECT EXISTS(SELECT 1 FROM schema_migrations WHERE version = ?)", m.Version).Scan(&exists)
if err != nil {
return err
}
if exists {
continue
}
// 应用迁移
if _, err := db.Exec(m.SQL); err != nil {
return fmt.Errorf("migration %d: %w", m.Version, err)
}
// 记录版本
if _, err := db.Exec("INSERT INTO schema_migrations (version) VALUES (?)", m.Version); err != nil {
return err
}
log.Printf("[迁移] 版本 %d 已应用", m.Version)
}
return nil
}
▶ 示例:集成测试
GO
📖 仅展示
// internal/handler/integration_test.go
package handler
import (
"bytes"
"database/sql"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"ecommerce/internal/middleware"
"ecommerce/internal/repository"
"ecommerce/internal/service"
_ "github.com/mattn/go-sqlite3"
)
func setupTestDB(t *testing.T) *sql.DB {
db, err := sql.Open("sqlite3", ":memory:")
if err != nil {
t.Fatal(err)
}
// 初始化表
_, err = db.Exec(`
CREATE TABLE users (id INTEGER PRIMARY KEY AUTOINCREMENT, email TEXT UNIQUE NOT NULL, password TEXT NOT NULL, name TEXT NOT NULL, role TEXT DEFAULT 'customer');
CREATE TABLE products (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, description TEXT, price REAL NOT NULL, stock INTEGER DEFAULT 0);
CREATE TABLE orders (id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, product_id INTEGER NOT NULL, quantity INTEGER NOT NULL, total_price REAL NOT NULL, status TEXT DEFAULT 'pending');
`)
if err != nil {
t.Fatal(err)
}
return db
}
func setupTestApp(db *sql.DB) http.Handler {
userRepo := repository.NewUserRepository(db)
productRepo := repository.NewProductRepository(db)
orderRepo := repository.NewOrderRepository(db)
userSvc := service.NewUserService(userRepo)
productSvc := service.NewProductService(productRepo)
orderSvc := service.NewOrderService(orderRepo, productRepo, userRepo)
userHandler := NewUserHandler(userSvc)
productHandler := NewProductHandler(productSvc)
orderHandler := NewOrderHandler(orderSvc)
mux := http.NewServeMux()
userHandler.Register(mux)
productHandler.Register(mux)
orderHandler.Register(mux)
// 应用中间件
return middleware.Chain(mux,
middleware.Recovery,
middleware.Logging,
middleware.CORS,
)
}
func TestRegisterAndLogin(t *testing.T) {
db := setupTestDB(t)
defer db.Close()
app := setupTestApp(db)
// 1. 注册
registerBody := `{"email":"test@example.com","password":"password123","name":"Test User"}`
req := httptest.NewRequest("POST", "/api/v1/register", bytes.NewBufferString(registerBody))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
app.ServeHTTP(resp, req)
if resp.Code != http.StatusCreated {
t.Fatalf("expected 201, got %d", resp.Code)
}
// 2. 登录
loginBody := `{"email":"test@example.com","password":"password123"}`
req = httptest.NewRequest("POST", "/api/v1/login", bytes.NewBufferString(loginBody))
req.Header.Set("Content-Type", "application/json")
resp = httptest.NewRecorder()
app.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d", resp.Code)
}
var result map[string]interface{}
json.NewDecoder(resp.Body).Decode(&result)
data := result["data"].(map[string]interface{})
if data["email"] != "test@example.com" {
t.Errorf("expected test@example.com, got %v", data["email"])
}
}
func TestCreateOrder_RequiresAuth(t *testing.T) {
db := setupTestDB(t)
defer db.Close()
app := setupTestApp(db)
// 没有 Token 的请求应该返回 401
orderBody := `{"product_id":1,"quantity":1}`
req := httptest.NewRequest("POST", "/api/v1/orders", bytes.NewBufferString(orderBody))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
// 注意:当前版本没有 Auth middleware,需要下一课添加
// 这个测试会失败,需要在集成测试环境中添加 Auth middleware
_ = resp
_ = app
}
sequenceDiagram
participant Client
participant MW as Middleware管线
participant AuthS as Auth Service
participant H as Handler
Client->>MW: POST /api/v1/orders (Bearer token)
MW->>MW: Recovery / Logging / CORS
MW->>AuthS: Auth 中间件验证 JWT
AuthS-->>MW: Claims {user_id, role}
MW->>MW: RBAC 检查 customer role
MW->>MW: RateLimit 检查
MW->>H: 通过!注入 user_id 到 Context
H->>H: 创建订单(使用 user_id)
H-->>MW: 201 Created
MW-->>Client: JSON 响应
(4) 请求经过中间件管线
TEXT
📖 仅展示
Request
↓ Recovery(捕获 panic)
↓ Logging(记录 + 耗时)
↓ CORS(跨域)
↓ Auth(JWT 验证 → 注入 Claims)
↓ RBAC(检查角色)
↓ RateLimit(令牌桶)
↓ Handler(业务逻辑)
Response
🔥 易错: Auth 中间件必须在 RBAC 之前。 先验证 Token 获取用户信息,再检查权限。如果先检查权限再验证 Token,未认证用户可能触发权限错误而不是认证错误——安全隐患。
❓ 常见问题
Q JWT 在 middleware 怎么集成?
A Auth 中间件解析 Header 中的 Bearer Token,调用
jwtService.ValidateToken,将 Claims 通过 context.WithValue 注入到请求上下文中。后续 Handler 通过 r.Context().Value(middleware.UserClaimsKey) 获取用户信息。Q RBAC 中级实现?
A 在 Auth 中间件之后,RequireRole 中间件检查 Claims 中的 role 字段是否在允许的角色列表中。更复杂的实现可以引入权限矩阵(每个角色有 permission 列表,中间件检查具体 permission)。
Q SQL migration 工具怎么选?
A (1)
golang-migrate/migrate——最常用,支持多种数据库和迁移源;(2) 手动迁移——本课的做法,适合小型项目;(3) pressly/goose——功能丰富。大项目推荐 golang-migrate。Q Integration test 怎么写?
A 用
httptest.NewServer 或 httptest.NewRecorder 模拟 HTTP 请求,结合临时数据库(:memory: 模式 SQLite)。测试流程:setup DB → 注入依赖 → 创建 Handler → 发送 HTTP 请求 → 验证响应。Q 分页中间件如何实现?
A 解析 query 参数
page 和 per_page,计算 offset。在 Service 层的查询中加上 LIMIT ? OFFSET ?。Handler 可以在响应中添加 X-Total-Count Header 或 meta 信息。中间件只负责解析参数,不修改查询逻辑。📖 小节
- JWT Auth 中间件:解析 Token、注入 Claims
- RBAC:RequireRole 中间件检查角色
- 中间件管线:Recovery → Logging → CORS → Auth → RBAC → RateLimit
- SQL 迁移:版本化管理表结构变更
- 集成测试:httptest + 临时数据库
- 分页:query 参数解析 + LIMIT/OFFSET
📝 作业
-
基础题(难度⭐):将本课的 Auth 中间件集成到上篇的电商 API 中。注册时返回 JWT Token,所有
/api/v1/orders/*端点需要 Bearer Token。验证无 Token 请求返回 401。 -
进阶题(难度⭐⭐):实现完整的中间件管线集成测试。要求:(1) 用
httptest+ 临时 SQLite 数据库;(2) 测试认证流程(注册→登录→获取 Token→用 Token 访问受保护端点);(3) 测试权限拒绝场景(customer 访问 admin 端点);(4) 测试限流场景(短时间大量请求返回 429)。 -
挑战题(难度⭐⭐⭐):实现 golang-migrate 的数据库迁移。要求:(1) 安装
golang-migrate/migrateCLI;(2) 编写 3 个迁移文件(创建 users、products、orders 表);(3) 迁移嵌入 Go 二进制(embed包);(4)cmd/migrate/main.go支持 up/down 命令;(5) 集成测试使用迁移后的数据库结构。