#核心实现
本章按依赖方向自底向上实现:存储层 → 业务层 → 认证与限流组件 → 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
}
gorm.ErrDuplicatedKey 的错误翻译需要 gorm.Config{TranslateError: true}(在装配篇的 gorm.Open 中开启),否则唯一索引冲突会以原始驱动错误返回。
#业务层:internal/service
服务层定义自己需要的最小接口(接口在消费方),存储实现由装配时注入——这正是测试一章"小接口 + 替身"的前提。
短码生成使用 crypto/rand(gosec 会拦截 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
}
密钥管理: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)
})
}
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
}
装配逻辑的自检清单:main 中没有出现任何业务分支;替换数据库只需换 store.New 的实现;errgroup 可在连接多个依赖(如再加缓存)时并行初始化(见工作池与 errgroup)。
至此服务已经完整。下一章为它补上测试、容器化部署与 CI,走完最后一公里。