核心实现

本章按依赖方向自底向上实现:存储层 → 业务层 → 认证与限流组件 → HTTP 层 → 入口装配。模块路径为 example.com/shortener,目录结构见设计篇

先初始化项目与依赖:

mkdir shortener && cd shortener
go mod init example.com/shortener
go get gorm.io/gorm gorm.io/driver/postgres
go get github.com/golang-jwt/jwt/v5 golang.org/x/crypto golang.org/x/time
go get github.com/prometheus/client_golang/prometheus

存储层:internal/store

模型定义见设计篇。存储层先将 GORM 错误翻译为领域哨兵错误,上层不再感知 SQL 细节:

package store

import "errors"

var (
    ErrNotFound      = errors.New("record not found")
    ErrAlreadyExists = errors.New("record already exists")
)
package store

import (
    "context"
    "errors"

    "gorm.io/gorm"
)

type GormStore struct {
    db *gorm.DB
}

func New(db *gorm.DB) *GormStore { return &GormStore{db: db} }

func (s *GormStore) AutoMigrate() error {
    return s.db.AutoMigrate(&User{}, &Link{})
}

func (s *GormStore) CreateUser(ctx context.Context, u *User) error {
    err := s.db.WithContext(ctx).Create(u).Error
    if errors.Is(err, gorm.ErrDuplicatedKey) {
        return ErrAlreadyExists
    }
    return err
}

func (s *GormStore) GetUserByUsername(ctx context.Context, username string) (*User, error) {
    var u User
    err := s.db.WithContext(ctx).Where("username = ?", username).First(&u).Error
    switch {
    case errors.Is(err, gorm.ErrRecordNotFound):
        return nil, ErrNotFound
    case err != nil:
        return nil, err
    }
    return &u, nil
}

func (s *GormStore) CreateLink(ctx context.Context, l *Link) error {
    err := s.db.WithContext(ctx).Create(l).Error
    if errors.Is(err, gorm.ErrDuplicatedKey) {
        return ErrAlreadyExists
    }
    return err
}

func (s *GormStore) GetLinkByCode(ctx context.Context, code string) (*Link, error) {
    var l Link
    err := s.db.WithContext(ctx).Where("code = ?", code).First(&l).Error
    switch {
    case errors.Is(err, gorm.ErrRecordNotFound):
        return nil, ErrNotFound
    case err != nil:
        return nil, err
    }
    return &l, nil
}

// IncrementClicks 用 SQL 原子自增,避免读改写竞态(见并发编程章节)
func (s *GormStore) IncrementClicks(ctx context.Context, code string) error {
    return s.db.WithContext(ctx).Model(&Link{}).
        Where("code = ?", code).
        UpdateColumn("clicks", gorm.Expr("clicks + 1")).Error
}
Warning

gorm.ErrDuplicatedKey 的错误翻译需要 gorm.Config{TranslateError: true}(在装配篇的 gorm.Open 中开启),否则唯一索引冲突会以原始驱动错误返回。

业务层:internal/service

服务层定义自己需要的最小接口(接口在消费方),存储实现由装配时注入——这正是测试一章"小接口 + 替身"的前提。

短码生成使用 crypto/randgosec 会拦截 math/rand 生成安全场景随机数):

package service

import (
    "crypto/rand"
    "math/big"
)

const alphabet = "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"

// NewSlug 生成 7 位 Base62 随机短码(62^7 ≈ 3.5e12 种组合)
func NewSlug() (string, error) {
    out := make([]byte, 7)
    for i := range out {
        n, err := rand.Int(rand.Reader, big.NewInt(int64(len(alphabet))))
        if err != nil {
            return "", err
        }
        out[i] = alphabet[n.Int64()]
    }
    return string(out), nil
}
package service

import (
    "context"
    "errors"
    "regexp"

    "example.com/shortener/internal/store"
)

// LinkStore:服务层对存储的最小需求,由 store.GormStore 实现
type LinkStore interface {
    CreateLink(ctx context.Context, l *store.Link) error
    GetLinkByCode(ctx context.Context, code string) (*store.Link, error)
    IncrementClicks(ctx context.Context, code string) error
}

var ErrInvalidCode = errors.New("code must match [0-9a-zA-Z_-]{3,16}")

var codePattern = regexp.MustCompile(`^[0-9a-zA-Z_-]{3,16}$`)

type LinkService struct {
    store    LinkStore
    maxRetry int
}

func NewLinkService(s LinkStore) *LinkService {
    return &LinkService{store: s, maxRetry: 3}
}

func (s *LinkService) Create(ctx context.Context, rawURL, customCode string, userID int64) (*store.Link, error) {
    // 自定义短码:校验格式,冲突即报错交由用户换码
    if customCode != "" {
        if !codePattern.MatchString(customCode) {
            return nil, ErrInvalidCode
        }
        link := &store.Link{Code: customCode, URL: rawURL, UserID: userID}
        if err := s.store.CreateLink(ctx, link); err != nil {
            return nil, err
        }
        return link, nil
    }

    // 随机短码:唯一索引冲突时换码重试
    var lastErr error
    for range s.maxRetry {
        code, err := NewSlug()
        if err != nil {
            return nil, err
        }
        link := &store.Link{Code: code, URL: rawURL, UserID: userID}
        lastErr = s.store.CreateLink(ctx, link)
        switch {
        case lastErr == nil:
            return link, nil
        case !errors.Is(lastErr, store.ErrAlreadyExists):
            return nil, lastErr
        }
    }
    return nil, lastErr
}

func (s *LinkService) Resolve(ctx context.Context, code string) (*store.Link, error) {
    link, err := s.store.GetLinkByCode(ctx, code)
    if err != nil {
        return nil, err
    }
    _ = s.store.IncrementClicks(ctx, code) // 计数失败不影响跳转
    return link, nil
}

用户服务负责注册与登录,登录失败统一报"用户名或密码错误",不泄露用户是否存在:

package service

import (
    "context"
    "errors"
    "time"

    "example.com/shortener/internal/auth"
    "example.com/shortener/internal/store"
)

var ErrInvalidCredentials = errors.New("invalid username or password")

type UserStore interface {
    CreateUser(ctx context.Context, u *store.User) error
    GetUserByUsername(ctx context.Context, username string) (*store.User, error)
}

type UserService struct {
    store  UserStore
    secret []byte
}

func NewUserService(s UserStore, jwtSecret string) *UserService {
    return &UserService{store: s, secret: []byte(jwtSecret)}
}

func (s *UserService) Register(ctx context.Context, username, password string) (*store.User, error) {
    hash, err := auth.HashPassword(password)
    if err != nil {
        return nil, err
    }
    u := &store.User{Username: username, PasswordHash: hash}
    if err := s.store.CreateUser(ctx, u); err != nil {
        return nil, err
    }
    return u, nil
}

func (s *UserService) Login(ctx context.Context, username, password string) (string, error) {
    u, err := s.store.GetUserByUsername(ctx, username)
    switch {
    case errors.Is(err, store.ErrNotFound):
        return "", ErrInvalidCredentials
    case err != nil:
        return "", err
    }
    if !auth.VerifyPassword(u.PasswordHash, password) {
        return "", ErrInvalidCredentials
    }
    return auth.NewToken(u.ID, s.secret, 24*time.Hour)
}

认证组件:internal/auth

密码用 bcrypt 单向哈希(盐内置于哈希串),令牌用 HS256 签名:

package auth

import "golang.org/x/crypto/bcrypt"

func HashPassword(plain string) (string, error) {
    hash, err := bcrypt.GenerateFromPassword([]byte(plain), bcrypt.DefaultCost)
    return string(hash), err
}

func VerifyPassword(hash, plain string) bool {
    return bcrypt.CompareHashAndPassword([]byte(hash), []byte(plain)) == nil
}
package auth

import (
    "errors"
    "time"

    "github.com/golang-jwt/jwt/v5"
)

var ErrInvalidToken = errors.New("invalid token")

func NewToken(userID int64, secret []byte, ttl time.Duration) (string, error) {
    claims := jwt.MapClaims{
        "sub": userID,
        "iat": time.Now().Unix(),
        "exp": time.Now().Add(ttl).Unix(),
    }
    return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(secret)
}

func ParseToken(tokenStr string, secret []byte) (int64, error) {
    token, err := jwt.Parse(tokenStr, func(t *jwt.Token) (any, error) {
        // 锁定签名算法,防止"算法替换"攻击
        if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
            return nil, ErrInvalidToken
        }
        return secret, nil
    })
    if err != nil || !token.Valid {
        return 0, ErrInvalidToken
    }
    claims, ok := token.Claims.(jwt.MapClaims)
    if !ok {
        return 0, ErrInvalidToken
    }
    sub, ok := claims["sub"].(float64) // JSON 数字解码默认为 float64
    if !ok {
        return 0, ErrInvalidToken
    }
    return int64(sub), nil
}
Warning

密钥管理:JWT_SECRET 只从环境变量注入(生产环境用密钥管理服务),泄漏等于任何人可伪造任意用户身份。代码中绝不出现硬编码密钥,gosec 也会扫描此类问题。

限流组件:internal/ratelimit

基于 golang.org/x/time/rate按 IP 令牌桶,限制跳转接口被刷:

package ratelimit

import (
    "net/http"
    "sync"

    "golang.org/x/time/rate"
)

type IPRateLimiter struct {
    mu       sync.Mutex
    visitors map[string]*rate.Limiter
    r        rate.Limit
    burst    int
}

func NewPerIP(r rate.Limit, burst int) *IPRateLimiter {
    return &IPRateLimiter{
        visitors: make(map[string]*rate.Limiter),
        r:        r,
        burst:    burst,
    }
}

func (l *IPRateLimiter) get(ip string) *rate.Limiter {
    l.mu.Lock()
    defer l.mu.Unlock()
    lim, ok := l.visitors[ip]
    if !ok {
        lim = rate.NewLimiter(l.r, l.burst)
        l.visitors[ip] = lim
    }
    return lim
}

// Middleware:每个 IP 每秒 r 个令牌,允许 burst 的突发。
func (l *IPRateLimiter) Middleware(next http.Handler) http.Handler {
    return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
        if !l.get(r.RemoteAddr).Allow() {
            w.Header().Set("Content-Type", "application/json; charset=utf-8")
            w.WriteHeader(http.StatusTooManyRequests)
            _, _ = w.Write([]byte(`{"error":"too many requests"}`))
            return
        }
        next.ServeHTTP(w, r)
    })
}
Tip

visitors 映射会随 IP 数量增长,生产服务应周期性清理长期不活跃的表项(后台 goroutine 定期遍历删除),或使用带淘汰策略的缓存实现。

HTTP 层:internal/handler

先统一响应格式与错误到状态码的映射:

package handler

import (
    "encoding/json"
    "errors"
    "net/http"

    "example.com/shortener/internal/service"
    "example.com/shortener/internal/store"
)

func writeJSON(w http.ResponseWriter, status int, v any) {
    w.Header().Set("Content-Type", "application/json; charset=utf-8")
    w.WriteHeader(status)
    _ = json.NewEncoder(w).Encode(v)
}

func writeErr(w http.ResponseWriter, status int, msg string) {
    writeJSON(w, status, map[string]string{"error": msg})
}

// mapError:领域错误 → HTTP 状态码,收口在一处
func mapError(w http.ResponseWriter, err error) {
    switch {
    case errors.Is(err, store.ErrNotFound):
        writeErr(w, http.StatusNotFound, "not found")
    case errors.Is(err, store.ErrAlreadyExists):
        writeErr(w, http.StatusConflict, "already exists")
    case errors.Is(err, service.ErrInvalidCode):
        writeErr(w, http.StatusBadRequest, err.Error())
    case errors.Is(err, service.ErrInvalidCredentials):
        writeErr(w, http.StatusUnauthorized, err.Error())
    case errors.Is(err, auth.ErrInvalidToken):
        writeErr(w, http.StatusUnauthorized, "invalid token")
    default:
        writeErr(w, http.StatusInternalServerError, "internal error") // 不向客户端泄漏内部细节
    }
}

示例需 import "example.com/shortener/internal/auth"

处理器与路由注册。认证中间件从 Authorization: Bearer 提取并验证令牌,将用户 ID 写入请求上下文:

package handler

import (
    "context"
    "encoding/json"
    "net/http"
    "strings"

    "example.com/shortener/internal/auth"
    "example.com/shortener/internal/ratelimit"
    "example.com/shortener/internal/service"
    "golang.org/x/time/rate"
)

type ctxKey string

const userIDKey ctxKey = "userID"

type Handler struct {
    links  *service.LinkService
    users  *service.UserService
    secret []byte
}

func New(links *service.LinkService, users *service.UserService, jwtSecret string) *Handler {
    return &Handler{links: links, users: users, secret: []byte(jwtSecret)}
}

func (h *Handler) RegisterRoutes(mux *http.ServeMux) {
    mux.HandleFunc("POST /api/register", h.register)
    mux.HandleFunc("POST /api/login", h.login)
    mux.Handle("POST /api/links", h.requireAuth(http.HandlerFunc(h.createLink)))
    // 跳转接口:按 IP 限流(每秒 5 个令牌,突发上限 10)
    mux.Handle("GET /{code}",
        ratelimit.NewPerIP(rate.Limit(5), 10).Middleware(http.HandlerFunc(h.redirect)))
}

func (h *Handler) requireAuth(next http.Handler) http.Handler {
    return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
        const prefix = "Bearer "
        authz := r.Header.Get("Authorization")
        if len(authz) <= len(prefix) || !strings.HasPrefix(authz, prefix) {
            writeErr(w, http.StatusUnauthorized, "missing bearer token")
            return
        }
        userID, err := auth.ParseToken(strings.TrimPrefix(authz, prefix), h.secret)
        if err != nil {
            writeErr(w, http.StatusUnauthorized, "invalid token")
            return
        }
        ctx := context.WithValue(r.Context(), userIDKey, userID)
        next.ServeHTTP(w, r.WithContext(ctx))
    })
}

func (h *Handler) register(w http.ResponseWriter, r *http.Request) {
    var in struct {
        Username string `json:"username"`
        Password string `json:"password"`
    }
    if err := json.NewDecoder(r.Body).Decode(&in); err != nil {
        writeErr(w, http.StatusBadRequest, "invalid json")
        return
    }
    u, err := h.users.Register(r.Context(), in.Username, in.Password)
    if err != nil {
        mapError(w, err)
        return
    }
    writeJSON(w, http.StatusCreated, u)
}

func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
    var in struct {
        Username string `json:"username"`
        Password string `json:"password"`
    }
    if err := json.NewDecoder(r.Body).Decode(&in); err != nil {
        writeErr(w, http.StatusBadRequest, "invalid json")
        return
    }
    token, err := h.users.Login(r.Context(), in.Username, in.Password)
    if err != nil {
        mapError(w, err)
        return
    }
    writeJSON(w, http.StatusOK, map[string]string{"token": token})
}

func (h *Handler) createLink(w http.ResponseWriter, r *http.Request) {
    userID, _ := r.Context().Value(userIDKey).(int64) // requireAuth 保证存在

    var in struct {
        URL  string `json:"url"`
        Code string `json:"code"`
    }
    if err := json.NewDecoder(r.Body).Decode(&in); err != nil {
        writeErr(w, http.StatusBadRequest, "invalid json")
        return
    }
    if in.URL == "" || !(strings.HasPrefix(in.URL, "http://") || strings.HasPrefix(in.URL, "https://")) {
        writeErr(w, http.StatusBadRequest, "url must start with http(s)://")
        return
    }

    link, err := h.links.Create(r.Context(), in.URL, in.Code, userID)
    if err != nil {
        mapError(w, err)
        return
    }
    writeJSON(w, http.StatusCreated, link)
}

func (h *Handler) redirect(w http.ResponseWriter, r *http.Request) {
    link, err := h.links.Resolve(r.Context(), r.PathValue("code"))
    if err != nil {
        mapError(w, err)
        return
    }
    http.Redirect(w, r, link.URL, http.StatusFound)
}

横切中间件(日志、指标、panic 恢复)集中在一个文件:

package handler

import (
    "log/slog"
    "net/http"
    "strconv"
    "time"
)

// Chain 从内向外应用中间件:Chain(h, A, B) == A(B(h))
func Chain(h http.Handler, mws ...func(http.Handler) http.Handler) http.Handler {
    for i := len(mws) - 1; i >= 0; i-- {
        h = mws[i](h)
    }
    return h
}

func Recover(next http.Handler) http.Handler {
    return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
        defer func() {
            if err := recover(); err != nil {
                slog.Error("panic recovered", "err", err, "path", r.URL.Path)
                http.Error(w, "internal error", http.StatusInternalServerError)
            }
        }()
        next.ServeHTTP(w, r)
    })
}

func Logging(logger *slog.Logger, next http.Handler) http.Handler {
    return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
        start := time.Now()
        sw := &statusWriter{ResponseWriter: w, status: http.StatusOK}
        next.ServeHTTP(sw, r)

        level := slog.LevelInfo
        if sw.status >= 500 {
            level = slog.LevelError
        }
        logger.Log(r.Context(), level, "http_request",
            "method", r.Method, "path", r.URL.Path,
            "status", sw.status, "duration_ms", time.Since(start).Milliseconds())
    })
}

func Metrics(next http.Handler) http.Handler {
    return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
        start := time.Now()
        sw := &statusWriter{ResponseWriter: w, status: http.StatusOK}
        next.ServeHTTP(sw, r)

        MetricsRequestsTotal.WithLabelValues(
            r.Method, r.URL.Path, strconv.Itoa(sw.status)).Inc()
        MetricsRequestDuration.WithLabelValues(r.Method, r.URL.Path).
            Observe(time.Since(start).Seconds())
    })
}

type statusWriter struct {
    http.ResponseWriter
    status int
}

func (w *statusWriter) WriteHeader(code int) {
    w.status = code
    w.ResponseWriter.WriteHeader(code)
}
package handler

import (
    "github.com/prometheus/client_golang/prometheus"
)

var (
    MetricsRequestsTotal = prometheus.NewCounterVec(
        prometheus.CounterOpts{Name: "http_requests_total", Help: "Total HTTP requests"},
        []string{"method", "path", "status"},
    )
    MetricsRequestDuration = prometheus.NewHistogramVec(
        prometheus.HistogramOpts{
            Name:    "http_request_duration_seconds",
            Help:    "HTTP request latency",
            Buckets: []float64{0.01, 0.05, 0.1, 0.5, 1, 5},
        },
        []string{"method", "path"},
    )
)

func init() {
    prometheus.MustRegister(MetricsRequestsTotal, MetricsRequestDuration)
}

入口装配:cmd/shortener/main.go

main 只做五件事:读配置 → 建日志 → 连库迁移 → 依赖注入 → 启动与优雅退出:

package main

import (
    "context"
    "errors"
    "log/slog"
    "net/http"
    "os"
    "os/signal"
    "syscall"
    "time"

    "github.com/prometheus/client_golang/prometheus/promhttp"
    "gorm.io/driver/postgres"
    "gorm.io/gorm"

    "example.com/shortener/internal/handler"
    "example.com/shortener/internal/service"
    "example.com/shortener/internal/store"
)

func main() {
    if err := run(); err != nil {
        slog.Error("server exited", "err", err)
        os.Exit(1)
    }
}

func run() error {
    // 1. 配置全部来自环境变量
    dsn := envOr("DATABASE_URL",
        "postgres://shortener:secret@localhost:5432/shortener?sslmode=disable")
    addr := envOr("ADDR", ":8080")
    jwtSecret := os.Getenv("JWT_SECRET")
    if jwtSecret == "" {
        return errors.New("JWT_SECRET is required")
    }

    // 2. 结构化日志
    logger := slog.New(slog.NewJSONHandler(os.Stdout, nil))
    slog.SetDefault(logger)

    // 3. 存储连接与迁移
    db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
        TranslateError: true, // 唯一冲突 → gorm.ErrDuplicatedKey
    })
    if err != nil {
        return err
    }
    st := store.New(db)
    if err := st.AutoMigrate(); err != nil {
        return err
    }

    // 4. 依赖注入:store → service → handler
    links := service.NewLinkService(st)
    users := service.NewUserService(st, jwtSecret)
    h := handler.New(links, users, jwtSecret)

    mux := http.NewServeMux()
    h.RegisterRoutes(mux)
    mux.Handle("GET /metrics", promhttp.Handler()) // 生产环境应限制为内网访问

    var app http.Handler = mux
    app = handler.Metrics(app)
    app = handler.Logging(logger, app)
    app = handler.Recover(app)

    srv := &http.Server{
        Addr:              addr,
        Handler:           app,
        ReadHeaderTimeout: 5 * time.Second,
        ReadTimeout:       10 * time.Second,
        WriteTimeout:      10 * time.Second,
    }

    // 5. 启动 + 优雅退出(信号驱动)
    errCh := make(chan error, 1)
    go func() {
        slog.Info("server listening", "addr", addr)
        if err := srv.ListenAndServe(); !errors.Is(err, http.ErrServerClosed) {
            errCh <- err
        }
    }()

    quit := make(chan os.Signal, 1)
    signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM)

    select {
    case err := <-errCh:
        return err
    case <-quit:
        ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
        defer cancel()
        return srv.Shutdown(ctx)
    }
}

func envOr(key, def string) string {
    if v := os.Getenv(key); v != "" {
        return v
    }
    return def
}
Tip

装配逻辑的自检清单:main 中没有出现任何业务分支;替换数据库只需换 store.New 的实现;errgroup 可在连接多个依赖(如再加缓存)时并行初始化(见工作池与 errgroup)。

至此服务已经完整。下一章为它补上测试、容器化部署与 CI,走完最后一公里。