红魔咖啡馆

头发越掉越多,头发越掉越少

0%

【go-zero】HTTP服务

服务端

配置

对HTTP服务主机、端口、整数进行控制

这些字段都可以在yaml配置文件中定义和修改

源码定义:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
RestConf struct {
    service.ServiceConf
    Host     string `json:",default=0.0.0.0"`
    Port     int
    CertFile string `json:",optional"`
    KeyFile  string `json:",optional"`
    Verbose  bool   `json:",optional"`
    MaxConns int    `json:",default=10000"`
    MaxBytes int64  `json:",default=1048576"`
    // milliseconds
    Timeout      int64         `json:",default=3000"`
    CpuThreshold int64         `json:",default=900,range=[0:1000]"`
    Signature    SignatureConf `json:",optional"`
    Middlewares MiddlewaresConf
  }

配置表格:

名称 类型 含义 默认值 是否必选
Host string 监听地址 0.0.0.0
Port int 监听端口
CertFile string https证书文件
KeyFile string https私钥文件
Verbose bool 是否打印详细日志
MaxConns int 并发请求数 10000
MaxBytes Int64 最大ContentLength 1048576
Timeout int64 超时时间(ms) 3000
CpuThreshold int64 降载阈值,默认900(90%),可允许设置范围0到1000 900
Signature SignatureConf 签名配置
Middlewares MiddlewaresConf 启用中间件
  • service.ServiceConf:通用服务配置,包括服务名、日志配置和运行模式等,是所有服务共有的基本设置
  • HostPort:监听的主机IP和端口
  • CertFileKeyFile:SSL证书文件路径和私钥文件路径,提供这两个文件可以启动HTTPS服务
  • Verbose:开启后会打印更详细的调试信息
  • MaxConns:最大并发连接数,防止请求过多服务器过载
  • MaxBytes:请求体的最大字节,默认1mb,防止请求过大
  • CpuThreshold:CPU使用率阈值,默认90%,当超过后服务会拒绝新请求,起到熔断保护
  • Signature:请求签名验证,包括启用状态,密钥等,用于身份校验
  • Middlewares:全局中间件的配置,如启用/禁用某个中间件,设置中间件参数等

中间件

go-zero中内置了如下中间件,均默认启用:

  • 鉴权管理中间件 AuthorizeHandler
  • 熔断中间件 BreakerHandler
  • 内容安全中间件 ContentSecurityHandler
  • 解密中间件 CryptionHandler
  • 压缩管理中间件 GunzipHandler
  • 日志中间件 LogHandler
  • ContentLength 管理中间件 MaxBytesHandler
  • 限流中间件 MaxConnsHandler
  • 指标统计中间件 MetricHandler
  • 普罗米修斯指标中间件 PrometheusHandler
  • panic 恢复中间件 RecoverHandler
  • 负载监控中间件 SheddingHandler
  • 超时中间件 TimeoutHandler
  • 链路追踪中间件 TraceHandler

配置如下:

1
2
3
4
5
6
7
8
9
10
11
12
13
type MiddlewaresConf struct {
    Trace      bool `json:",default=true"`
    Log        bool `json:",default=true"`
    Prometheus bool `json:",default=true"`
    MaxConns   bool `json:",default=true"`
    Breaker    bool `json:",default=true"`
    Shedding   bool `json:",default=true"`
    Timeout    bool `json:",default=true"`
    Recover    bool `json:",default=true"`
    Metrics    bool `json:",default=true"`
    MaxBytes   bool `json:",default=true"`
    Gunzip     bool `json:",default=true"`
}

例,如果想禁用指标收集,在etc/config.yaml配置:

1
2
3
4
5
Name: HelloWorld.api
Host: 127.0.0.1
Port: 8080
Middlewares:
  Metrics: false

main.go应用配置:

1
2
3
4
5
6
7
func main() {
    var restConf rest.RestConf
    conf.MustLoad("etc/config.yaml", &restConf)
    srv := rest.MustNewServer(restConf)
    defer srv.Stop()
    ...
}

TraceHandler

用于链路追踪,go-zero继承了Opentelemetry来标准化链路追踪,如果想把追踪信息上报jaeger,可以做如下配置:

1
2
3
4
5
6
7
8
9
10
Name: HelloWorld.api
Host: 127.0.0.1
Port: 8080
Middlewares:
  Metrics: true
Telemetry:
  Name: hello
  Endpoint: localhost:4317
  Batcher: otlpgrpc
  Sampler: 1.0

应用配置后,打开Jaeger就可以查看链路请求信息,默认携带http.hosthttp.methodhttp.routehttp.status_code等属性

LogHandler

默认每次http请求都会输出对应请求日志,格式如下,可以通过配置禁用该中间件:

1
2
3
4
5
6
7
{
  "@timestamp": "2023-03-05T21:53:44.244+08:00",
  "caller": "handler/loghandler.go:160",
  "content": "[HTTP] 200 - GET /hello - 127.0.0.1:51499 - Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/110.0.0.0 Safari/537.36",
  "duration": "0.0ms",
  "level": "info"
}

PrometheusHandler

httpserver默认集成Prometheus指标监控,用于自动收集每个HTTP请求的性能指标,默认启用

  • 请求耗时统计:指标类型为Histogram,指标名为http_server_requests_duration_ms

    用于记录每个请求花了多少毫秒,按照path(接口路径)分组,默认buckets定义分别为 5, 10, 25, 50, 100, 250, 500, 1000(毫秒)

  • 错误请求计数:指标类型为Counter,指标名为http_server_requests_code_total

    记录不同HTTP状态码的请求总数,按path分组

MaxConnsHandler

用于限制http最大并发请求数,当超过设置值会返回http.StatusServiceUnavailable状态码

自定义中间件

写函数自定义中间件,使用server.Use()应用自定义中间件

1
2
3
4
5
6
7
8
9
10
11
12
server := rest.MustNewServer(rest.RestConf{})
defer server.Stop()

server.Use(middleware)

// 自定义的中间件
func middleware(next http.HandlerFunc) http.HandlerFunc {
  return func(w http.ResponseWriter, r *http.Request) {
    w.Header().Add("X-Middleware", "static-middleware")
    next(w, r)
  }
}

错误处理

自定义业务错误类型

内部定义的错误类型errors.CodeMsg,包含两个字段:

  • Code:业务状态码
  • Msg:提示信息

使用errors.New()创建

1
2
3
import "github.com/zeromicro/x/errors"
// 在需要的地方创建带错误码和消息的错误
httpx.Error(w, errors.New(400, "参数错误"))

统一错误处理器

调用httpx.Error处理响应时会执行SetErrorHandler

  • 如果错误是*errors.CodeMsg,说明是预期的业务错误,返回200,但JSON中携带错误代码和消息
  • 其他错误统一视为服务端错误,返回500
1
2
3
4
5
6
7
8
9
10
11
12
13
httpx.SetErrorHandler(func(err error) (int, any) {
    switch e := err.(type) {
    case *errors.CodeMsg:
        // 业务错误 → HTTP 200,通过 Body 返回 Code+Msg
        return http.StatusOK, xhttp.BaseResponse[struct{}]{
            Code: e.Code,
            Msg:  e.Msg,
        }
    default:
        // 未知错误 → HTTP 500,Body 为空或通用错误
        return http.StatusInternalServerError, nil
    }
})

在handler中使用

  • 成功的响应直接使用httpx.OKJson
  • 若失败统一用httpx.Error,让errorhandler判断
1
2
3
4
5
6
7
8
9
10
11
12
func handle(w http.ResponseWriter, r *http.Request) {
    var req HelloRequest
    if err := httpx.Parse(r, &req); err != nil {
        httpx.Error(w, err)  // 解析失败,触发 errorHandler,返回 500
        return
    }
    if req.Name == "error" {
        httpx.Error(w, errors.New(400, "参数错误")) // 业务错误,返回 200 {code:400}
        return
    }
    httpx.OkJson(w, HelloResponse{Msg: "hello " + req.Name})
}

请求体

gozero中支持:

  • 通过http.Request的Body字段获取请求参数,只接受application/json格式
  • 通过http.Request的Form字段获取请求参数,只接受application/x-www-form-urlencoded格式
  • 通过httpx.Parse方法获取path参数和请求头参数

Form表单请求参数

可以通过结构体form tag定义参数名称,支持content-type为application/x-www-form-urlencoded的请求参数,也支持get方法的query参数

1
2
3
4
5
6
7
type Request struct {
    Name    string  `form:"name"` // 必填参数
    Age     int     `form:"age,optional"` // optional定义非必填参数
}

var req Request
err := httpx.Parse(r, &req) // 解析参数

x-www-form-urlencoded是一种HTTP请求格式,用于将表单数据发送到服务器,他会把键值对转换成类似URL query字符串的样子,特殊字符会被百分号编码,如key1=value1&key2=value2

gozero的表单tag可以直接解析这种格式

JSON请求参数

JSON参数请求接收来自POST请求,content-type为application/json的请求参数,通过结构体json tag定义参数名称

1
2
3
4
5
6
7
type Request struct {
    Name string `json:"name"`
    Age  int    `json:"age"`
}

var req Request
err := httpx.Parse(r, &req) // 解析参数

Path请求参数

Path参数请求支持在路径上定义参数,用于获取特定的路径参数,通过结构体path tag定义参数变量名

1
2
3
4
5
6
7
8
9
10
11
12
13
type Request struct {
    Name string `path:"name"`
}

// Path定义
rest.Route{
    Method:  http.MethodGet,
    Path:    "/user/:name",
    Handler: handle,
}

var req Request
err := httpx.Parse(r, &req) // 解析参数

Header参数获取

用于获取特定的请求头的值,通过结构体header tag定义请求头的key值

1
2
3
4
5
6
type Request struct {
    Age int `form:"age,optional"`
}

var req Request
err := httpx.Parse(r, &req) // 解析参数

参数校验

optional:

如果参数不是必填的,可以在参数后面加上optional关键字

对于必填参数,若没有传递参数,会返回如下错误field age is not set

range:

可以定义参数区间,在结构体参数后面加上range关键字,区间说明见上面

default:

通过default来定义参数默认值,在结构体参数后面加上default关键字,默认值为零值

响应体

gozero中支持:

  • 通过http.ResopnseWriterWrite方法返回相应参数
  • 通过http.ResopnseWriterHeader方法返回响应头
  • 默认支持application/json格式
1
2
3
4
5
6
7
type Response struct {
  Name   string `json:"name"`
  Age    int    `json:"age"`
}

resp := &Response{Name: "jack", Age: 18}
httpx.OkJson(w, resp)
  • httpx.OKJson(w, resp)返回一个200状态的Json
  • httpx.WriteJson(w, statusCode, data)返回自定义状态码的Json

code-data统一响应

为了和前端达成统一的响应格式,我们通常会封装一层业务code,msg以及业务数据,格式如下:

1
2
3
4
5
6
7
{
  "code": 0,
  "msg": "ok",
  "data": {
    ...
  }
}

我们可以通过zeromicro/x扩展包代替默认响应写法

  • errors.New(code int, msg string):创建一个带业务状态码的错误,而不仅仅是string
  • xhttp.JsonBaseResponseCtx(ctx, w, v):将响应自动包装为统一格式,写入ResponseWriter

Server-Sent Events

用于实时数据更新,允许服务器通过持久化的HTTP连接向客户端单向推送事件,比起WebSocket它更轻量,适合简单的实时更新场景

核心特点:

  • 单向通信:服务器主动推送,客户端被动接收
  • 简单协议:基于text/event-stream格式,易于实现
  • 自动重连:浏览器内置重连机制,断开后可自动尝试恢复

服务端代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
package main

import (
  "fmt"
  "net/http"
  "time"

  "github.com/zeromicro/go-zero/core/logx"
  "github.com/zeromicro/go-zero/rest"
)

type SseHandler struct {
  clients map[chan string]bool
}

func NewSseHandler() *SseHandler {
  return &SseHandler{
    clients: make(map[chan string]bool),
  }
}

// Serve 处理 SSE 连接
func (h *SseHandler) Serve(w http.ResponseWriter, r *http.Request) {
  // 设置 SSE 必需的 HTTP 头
  // for versions > v1.8.1, no need to add 3 lines below
  w.Header().Add("Content-Type", "text/event-stream")
  w.Header().Add("Cache-Control", "no-cache")
  w.Header().Add("Connection", "keep-alive")

  // 为每个客户端创建一个 channel
  clientChan := make(chan string)
  h.clients[clientChan] = true

  // 客户端断开时清理
  defer func() {
    delete(h.clients, clientChan)
    close(clientChan)
  }()

  // 持续监听并推送事件
  for {
    select {
    case msg := <-clientChan:
      // 发送事件数据
      fmt.Fprintf(w, "data: %s\n\n", msg)
      w.(http.Flusher).Flush()
    case <-r.Context().Done():
      // 客户端断开连接
      return
    }
  }
}

// SimulateEvents 模拟周期性事件
func (h *SseHandler) SimulateEvents() {
  ticker := time.NewTicker(time.Second)
  defer ticker.Stop()

  for range ticker.C {
    message := fmt.Sprintf("Server time: %s", time.Now().Format(time.RFC3339))
    // 广播给所有客户端
    for clientChan := range h.clients {
      select {
      case clientChan <- message:
      default:
        // 跳过阻塞的 channel
      }
    }
  }
}

func main() {
  // 创建 go-zero REST 服务,集成静态文件服务
  server := rest.MustNewServer(rest.RestConf{
    Host: "0.0.0.0",
    Port: 8080,
  }, rest.WithFileServer("/static", http.Dir("static")))
  defer server.Stop()

  // 初始化 SSE 处理
  sseHandler := NewSseHandler()

  // 注册 SSE 路由
  // for go-zero versions > v1.8.1
  server.AddRoute(rest.Route{
    Method:  http.MethodGet,
    Path:    "/sse",
    Handler: sseHandler.Serve,
  }, rest.WithSSE())

  // 在单独的 goroutine 中模拟事件
  go sseHandler.SimulateEvents()

  logx.Info("Server starting on :8080")
  server.Start()
}
  • SSE服务以一个SSEHandler结构,其中定义了一个存储channel的映射,维护所有客户端channel,方便广播消息
  • Serve方法用于处理channel连接:
    • 首先设置SSE必须的HTTP头(gozero1.8.1往上不需要)
    • 为每个连接创建一个channel,存入clients
    • 使用select监听channel消息或客户端断开信号
    • 收到消息时格式化为SSE协议,通过Flush()推送
    • 最后要关闭所有channel
  • SimulateEvents方法模拟发送持续请求:
    • 通过time.Ticker每秒生成一个事件
    • 遍历所有channel,将消息广播到所有连接的客户端
    • 使用select+default,避免某个客户端阻塞影响整体
  • main函数:
    • 使用rest.MustNewServer创建服务,监听端口
    • 通过rest.WithFileServer配置静态文件服务,映射/static到本地static目录
    • 注册/sse路由,绑定SseHandler.Server并禁用超时,若在api文件中定义SSE路由,需要加上timeout: 0s

代码生成

@server语句块下声明sse:true

其中,sse的handler模板中是一个以chan来传递事件的channel,而不是一个响应体

客户端

HTTP client是一个用于发送HTTP请求的库,支持如下功能

  • content-type自动识别,仅支持 application/jsonapplication/x-www-form-urlencoded 两种格式
  • 支持将 path 参数自动填充到 url
  • 支持将结构体中 header 填充到 http 请求头

Form请求

get和post使用方式相同,只需要将结构体中tag标记form即可

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
type Request struct {
    Node   string `path:"node"`
    ID     int    `form:"id"`
    Header string `header:"X-Header"`
}

var domain = flag.String("domain", "http://localhost:3333", "the domain to request")

func main() {
    flag.Parse()

    req := Request{
        Node:   "foo",
        ID:     1024,
        Header: "foo-header",
    }
    resp, err := httpc.Do(context.Background(), http.MethodGet, *domain+"/nodes/:node", req)
    // resp, err := httpc.Do(context.Background(), http.MethodPost, *domain+"/nodes/:node", req)
    if err != nil {
        fmt.Println(err)
        return
    }

    io.Copy(os.Stdout, resp.Body)
}

Json请求

用法相同,只需要修改结构体tag为json即可

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
type Request struct {
    Node   string `path:"node"`
    Foo    string `json:"foo"`
    Bar    string `json:"bar"`
    Header string `header:"X-Header"`
}

var domain = flag.String("domain", "http://localhost:3333", "the domain to request")

func main() {
    flag.Parse()

    req := Request{
        Node:   "foo",
        Header: "foo-header",
        Foo: "foo",
        Bar: "bar",
    }
    resp, err := httpc.Do(context.Background(), http.MethodPost, *domain+"/nodes/:node", req)
    if err != nil {
        fmt.Println(err)
        return
    }

    io.Copy(os.Stdout, resp.Body)
}

JWT认证

api文件中的@server语句块中声明jwt配置后,框架会自动注入token验证中间件

API规范

在api文件中定义公开接口,启用jwt中间件,当客户端访问这下面的接口是,必须在请求头中携带Authorization: Bearer <token>

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
service user-api {
    // 公开接口
    @handler Login
    post /user/login (LoginReq) returns (LoginResp)
    @handler RefreshToken
    post /user/refresh (RefreshReq) returns (LoginResp)
}

@server (
    jwt: Auth
)
service user-api {
    @handler GetProfile
    get /user/profile (ProfileReq) returns (ProfileResp)
    @handler UpdateProfile
    put /user/profile (UpdateProfileReq) returns (UpdateProfileResp)
}

配置文件

etc/user-api.yaml中配置

1
2
3
4
5
Auth:
  AccessSecret: "your-256-bit-secret"
  AccessExpire: 86400        # Access Token 有效期:24 小时
RefreshSecret: "another-random-secret"
RefreshExpire: 604800        # Refresh Token 有效期:7 天

生成access token

1
2
3
4
5
6
7
8
9
10
11
func generateAccessToken(secret string, userId int64, role string) (string, error) {
    now := time.Now()
    claims := jwt.MapClaims{
        "userId": userId,
        "role":   role,
        "iat":    now.Unix(),
        "exp":    now.Add(24 * time.Hour).Unix(),
    }
    return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).
        SignedString([]byte(secret))
}
  • claim是一个MapClaims,可以随意添加自定义字段和标准字段
  • 定义签名算法为HS256,使用SignedString获取最终的JWT字符串

登录逻辑

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
func (l *LoginLogic) Login(req *types.LoginReq) (*types.LoginResp, error) {
    user, err := l.svcCtx.UserModel.FindOneByUsername(l.ctx, req.Username)
    if err != nil {
        return nil, errorx.NewCodeError(401, "用户名或密码错误")
    }
    if !checkPassword(req.Password, user.Password) {
        return nil, errorx.NewCodeError(401, "用户名或密码错误")
    }

    accessToken, _ := generateAccessToken(
        l.svcCtx.Config.Auth.AccessSecret, user.Id, user.Role)
    refreshToken, _ := generateRefreshToken(
        l.svcCtx.Config.RefreshSecret, user.Id)

    return &types.LoginResp{
        AccessToken:  accessToken,
        RefreshToken: refreshToken,
        ExpiresIn:    l.svcCtx.Config.Auth.AccessExpire,
    }, nil
}
  • 根据用户名查询数据库,若不存在返回401
  • 对比密码,不匹配返回401
  • 认证通过后,调用 generateAccessTokengenerateRefreshToken 分别生成两个 JWT
  • 返回LoginResp,包括两个jwt与accessToken的过期时间

刷新Token

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
func (l *RefreshTokenLogic) RefreshToken(req *types.RefreshReq) (*types.LoginResp, error) {
    token, err := jwt.ParseWithClaims(req.RefreshToken, jwt.MapClaims{},
        func(t *jwt.Token) (any, error) {
            return []byte(l.svcCtx.Config.RefreshSecret), nil
        })
    if err != nil || !token.Valid {
        return nil, errorx.NewCodeError(401, "Refresh Token 已失效或非法")
    }
    claims := token.Claims.(jwt.MapClaims)
    userId := int64(claims["userId"].(float64))

    user, _ := l.svcCtx.UserModel.FindOne(l.ctx, userId)
    newAccess, _ := generateAccessToken(
        l.svcCtx.Config.Auth.AccessSecret, userId, user.Role)
    newRefresh, _ := generateRefreshToken(
        l.svcCtx.Config.RefreshSecret, userId)

    return &types.LoginResp{
        AccessToken:  newAccess,
        RefreshToken: newRefresh,
        ExpiresIn:    l.svcCtx.Config.Auth.AccessExpire,
    }, nil
}
  • 客户端发送refreshReq,用RefreshSecret解析RefreshToken
  • 若解析失败或token过期,返回401
  • 从Token的claim中取出userId,查询数据库相关信息
  • 重新生成一份新的token返回
  • 旧的refresh token就此失效

自动验证

只要在接口声明了jwt: Auth,gozero生成的handler就会自动包含校验代码,类似于:

1
2
3
4
// 生成代码中会包含类似这样的调用
jwtMiddleware := jwt.NewMiddleware(svcCtx.Config.Auth)
// 然后将 handler 包裹在中间件内部
handler = jwtMiddleware(handler)

无需再手动解析token,只需要通过l.ctx.Value("UserId")等上下文值就可以获取用户信息

文件上传

gozero的每个handler都可以访问原始http.Request,所以go的multipart库可以直接用

API定义

1
2
3
4
5
6
7
8
9
10
service upload-api {
    @handler UploadFile
    post /upload/file returns (UploadResp)
}

type UploadResp {
    Filename string `json:"filename"`
    Size     int64  `json:"size"`
    URL      string `json:"url"`
}
  • goctl生成后,UploadFileHandler是一个http.HandlerFunc,因此可以访问http.Request
  • 返回结构包括原始文件名、文件大小和访问路径

Hanlder层

用于解析请求与校验文件类型

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
func UploadFileHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
    return func(w http.ResponseWriter, r *http.Request) {
        // 最多 32 MB 不落盘,超出部分自动落盘
        if err := r.ParseMultipartForm(32 << 20); err != nil {
            httpx.Error(w, err)
            return
        }

        file, header, err := r.FormFile("file")
        if err != nil {
            httpx.Error(w, err)
            return
        }
        defer file.Close()

        // 通过读取前 512 字节检测 MIME 类型
        buf := make([]byte, 512)
        n, _ := file.Read(buf)
        mimeType := http.DetectContentType(buf[:n])
        if !isAllowedType(mimeType) {
            httpx.Error(w, fmt.Errorf("不支持的文件类型: %s", mimeType),
                http.StatusUnsupportedMediaType)
            return
        }
        file.Seek(0, io.SeekStart)

        l := logic.NewUploadFileLogic(r.Context(), svcCtx)
        resp, err := l.UploadFile(file, header)
        if err != nil {
            httpx.Error(w, err)
            return
        }
        httpx.OkJson(w, resp)
    }
}
  • 设定最大缓存为32mb,超过了自动写入临时文件

  • FormFile返回multipart.File和文件头信息

  • 使用http.DetectContentType读取文件头512字节检测MIME类型

    注意:这里不要相信客户端声明的content-type或后缀名,需要自行判断,并实现isAllowedType来设置配型白名单

  • 读取文件后将指针Seek送回开头,以便后续逻辑完整写入

Logic层

用于在服务端校验与本地文件存储

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
func (l *UploadFileLogic) UploadFile(file multipart.File, header *multipart.FileHeader) (*types.UploadResp, error) {
    const maxSize = 10 << 20 // 10 MB
    if header.Size > maxSize {
        return nil, errorx.NewCodeError(400, "文件超过 10 MB 限制")
    }

    safeFilename := fmt.Sprintf("%d_%s", time.Now().UnixNano(),
        filepath.Base(filepath.Clean(header.Filename)))

    dst, err := os.Create(filepath.Join(l.svcCtx.Config.UploadDir, safeFilename))
    if err != nil {
        return nil, err
    }
    defer dst.Close()

    size, err := io.Copy(dst, file)
    if err != nil {
        return nil, err
    }

    return &types.UploadResp{
        Filename: header.Filename,
        Size:     size,
        URL:      "/files/" + safeFilename,
    }, nil
}
  • 首先校验大小,对比设定的业务上限,进行二次检查

  • filepath.Base(filepath.Clean(header.Filename)))用于文件名安全处理,防止攻击,前面再接入纳秒时间戳,确保文件名唯一

  • l.svcCtx.Config.UploadDir设定存储目录,一般放在服务配置的etc/app.yaml里:

    1
    2
    MaxBytes: 67108864   # 整个请求 body 大小限制 64 MB
    UploadDir: ./uploads
  • 最后返回的URL生成相对路径,前面需要另外的路由映射到该目录

云存储

一般生产环境需要上传到云存储,如AWS,阿里云等平台,下面的示例是将文件存入AWS S3

1
2
3
4
5
6
7
8
9
10
11
12
func (l *UploadFileLogic) uploadToS3(file multipart.File, key string) (string, error) {
    _, err := l.svcCtx.S3.PutObject(l.ctx, &s3.PutObjectInput{
        Bucket: aws.String(l.svcCtx.Config.S3Bucket),
        Key:    aws.String(key),
        Body:   file,
    })
    if err != nil {
        return "", err
    }
    return fmt.Sprintf("https://%s.s3.amazonaws.com/%s",
        l.svcCtx.Config.S3Bucket, key), nil
}