BLOG / POST
Dependency injection
依赖注入是一种管理组件间依赖关系的设计模式,用于实现对象间的松耦合。
- di
- design-pattern
依赖注入是一种设计模式,主要用于实现对线之间的松耦合,管理代码中各个组件之间的依赖关系。
在传统的软件设计中,对象通常自己负责创建和管理它们所依赖的其他对象。而依赖注入则是将对象的创建和管理职责从对象本身转移到外部的容器或框架中。这个外部的实体负责创建被依赖的对象,并将其注入到需要它的对象中,从而实现对象之间的解耦。
应用场景
一个后端单体服务,大致会分为如下三层:
- handler(controller) 层:用于处理HTTP请求,调用具体业务的服务,然后返回响应
- service 层:处理具体的业务逻辑
- repository 层:调用数据库CURD等
实现一个登录接口
- UserHandler
// handler/user.go
type UserHandler interface {
LoginPwd(c fiber.Ctx) error
}
type userHandler struct {
userService service.UserService // 注入 user service 获取业务处理逻辑
}
// 创建userHandler
func NewUserHandler(userService service.UserService) UserHandler {
return &userHandler{
userService: userService,
}
}
// 登录接口 handler
func (u *userHandler) LoginPwd(c fiber.Ctx) error {
params := new(req.LoginPwdReq)
if err := c.Bind().JSON(params); err != nil {
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
"message": err.Error(),
})
}
// 调用具体的登录逻辑方法
result, err := u.userService.LoginPwd(c, params)
if err != nil {
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
"message": err.Error(),
})
}
return c.Status(fiber.StatusOK).JSON(result)
}
- UserService
// service/user.go
type UserService interface {
LoginPwd(ctx fiber.Ctx, req *req.LoginPwdReq) (*res.LoginRes, error)
}
type userService struct {
Repository *Repository
}
func NewUserService(repo *Repository) UserService {
return &userService{
Repository: repo ,
}
}
func (s *userService) LoginPwd(ctx fiber.Ctx, req *req.LoginPwdReq) {
// 具体登录逻辑代码
...
// 调用 repository
s.Repository.GetUserById()
}
- main.go
func main() {
// 伪代码
repo := NewRepo()
userService := service.NewUserService(repo)
userHandler := handler.NewUserHandler(userService)
app := server.NewHttpServer(userHandler)
return app, func() {
}, nil
defer cleanup()
app.Listen(fmt.Sprintf(":%d", config.GetInt("http.port")))
}
上面代码中:
- userHandler 通过依赖注入获取到了userService的具体业务逻辑
- userService 通过依赖注入获取了数据库相关的操作方法
通过依赖注入的设计模式,实现分层之间的代码解藕。各层之间的逻辑也变得更加清晰明了,方便维护。
wire
上面代码中,目前只有一个 user的服务。但是随着项目逐渐庞大起来,依赖注入的过程就会变得重复和繁琐。wire 框架的作用就是来简化这一过程。
wire 框架只需要你创建一个 wire.go 文件。文件内容大致如下:
//go:build wireinject
// +build wireinject
package wire
import (
"github.com/gofiber/fiber/v3"
"github.com/google/wire"
"github.com/spf13/viper"
"github.com/vinoMamba/AiDoc/internal/handler"
"github.com/vinoMamba/AiDoc/internal/repository"
"github.com/vinoMamba/AiDoc/internal/server"
"github.com/vinoMamba/AiDoc/internal/service"
"github.com/vinoMamba/AiDoc/pkg/jwt"
"github.com/vinoMamba/AiDoc/pkg/mail"
"github.com/vinoMamba/AiDoc/pkg/redis"
"github.com/vinoMamba/AiDoc/pkg/sid"
)
var serverSet = wire.NewSet(
server.NewHttpServer,
)
var handlerSet = wire.NewSet(
handler.NewUserHandler,
)
var serviceSet = wire.NewSet(
service.NewService,
service.NewUserService,
repository.New,
repository.NewConn,
redis.NewRedisConn,
)
func NewApp(*viper.Viper) (*fiber.App, func(), error) {
panic(wire.Build(
serverSet,
handlerSet,
serviceSet,
sid.NewSid,
jwt.NewJWT,
mail.NewMail,
))
}
然后你只需要执行命令:
wire cmd/server/wire # 具体路径
就会生成一个 wire_gen.go 的代码:
// Code generated by Wire. DO NOT EDIT.
//go:generate go run -mod=mod github.com/google/wire/cmd/wire
//go:build !wireinject
// +build !wireinject
package wire
import (
"github.com/gofiber/fiber/v3"
"github.com/google/wire"
"github.com/spf13/viper"
"github.com/vinoMamba/AiDoc/internal/handler"
"github.com/vinoMamba/AiDoc/internal/repository"
"github.com/vinoMamba/AiDoc/internal/server"
"github.com/vinoMamba/AiDoc/internal/service"
"github.com/vinoMamba/AiDoc/pkg/jwt"
"github.com/vinoMamba/AiDoc/pkg/mail"
"github.com/vinoMamba/AiDoc/pkg/redis"
"github.com/vinoMamba/AiDoc/pkg/sid"
)
// Injectors from wire.go:
func NewApp(viperViper *viper.Viper) (*fiber.App, func(), error) {
dbtx := repository.NewConn(viperViper)
queries := repository.New(dbtx)
sidSid := sid.NewSid()
jwtJWT := jwt.NewJWT(viperViper)
mailMail := mail.NewMail(viperViper)
redisInternal := redis.NewRedisConn(viperViper)
serviceService := service.NewService(queries, sidSid, jwtJWT, viperViper, mailMail, redisInternal)
userService := service.NewUserService(serviceService)
userHandler := handler.NewUserHandler(userService)
spaceService := service.NewSpaceService(serviceService)
spaceHandler := handler.NewSapceHandler(spaceService)
memberService := service.NewMemberService(serviceService)
memberHandler := handler.NewMemberHandler(memberService)
noteService := service.NewNoteService(serviceService)
notificationHandler := handler.NewNotificationHandler(noteService)
projectService := service.NewProjectService(serviceService)
projectHandler := handler.NewProjectHandler(projectService)
app := server.NewHttpServer(userHandler, spaceHandler, memberHandler, notificationHandler, projectHandler, jwtJWT)
return app, func() {
}, nil
}
// wire.go:
var serverSet = wire.NewSet(server.NewHttpServer)
var handlerSet = wire.NewSet(handler.NewUserHandler, handler.NewSapceHandler, handler.NewMemberHandler, handler.NewNotificationHandler, handler.NewProjectHandler)
var serviceSet = wire.NewSet(service.NewService, service.NewUserService, service.NewSpaceService, service.NewMemberService, service.NewNoteService, service.NewProjectService, repository.New, repository.NewConn, redis.NewRedisConn)
最后,只需要在 main.go 中调用 NewApp
package main
import (
"flag"
"fmt"
"io"
"os"
"github.com/gofiber/fiber/v3/log"
"github.com/vinoMamba/AiDoc/cmd/server/wire"
"github.com/vinoMamba/AiDoc/pkg/config"
)
func main() {
envConf := flag.String("config", "./config/local.yaml", "config file path")
flag.Parse()
config := config.NewConfig(*envConf)
file, _ := os.OpenFile(config.GetString("log.log_file_name"), os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0666)
iw := io.MultiWriter(os.Stdout, file)
log.SetOutput(iw)
// 调用 NewApp
app, cleanup, err := wire.NewApp(config)
if err != nil {
panic(err)
}
defer cleanup()
app.Listen(fmt.Sprintf(":%d", config.GetInt("http.port")))
}
具体介绍看官方文档 https://github.com/google/wire
完。