182 lines
5.1 KiB
Go
182 lines
5.1 KiB
Go
package router
|
||
|
||
import (
|
||
"ai_scheduler/internal/config"
|
||
errorcode "ai_scheduler/internal/data/error"
|
||
errors "ai_scheduler/internal/data/error"
|
||
"ai_scheduler/internal/middleware"
|
||
"ai_scheduler/internal/pkg"
|
||
"ai_scheduler/internal/services/advice"
|
||
"encoding/json"
|
||
"reflect"
|
||
"strings"
|
||
|
||
"github.com/go-playground/validator/v10"
|
||
"github.com/gofiber/fiber/v2"
|
||
"github.com/gofiber/websocket/v2"
|
||
)
|
||
|
||
// SetupRoutes 设置路由
|
||
func SetupRoutes(cfg *config.Config, app *fiber.App, adviceFile *advice.FileService, adviceData *advice.AdvicerService,
|
||
adviceChat *advice.ChatService, adviceProject *advice.ProjectService, adviceTalkSkill *advice.TalkSkillService, adviceClient *advice.ClientService,
|
||
industry *advice.IndustryService, admin *advice.AdminService, modelSup *advice.ModelSupService,
|
||
activity *advice.ActivityService, wxHook *advice.WxHookService,
|
||
smart *advice.SmartService,
|
||
wxProxy *advice.WxProxyService,
|
||
wsHub *advice.WxWsHub,
|
||
adviceCustomer *advice.CustomerService,
|
||
adviceLabel *advice.LabelService,
|
||
) {
|
||
app.Use(func(c *fiber.Ctx) error {
|
||
// 设置 CORS 头
|
||
c.Set("Access-Control-Allow-Origin", "*")
|
||
c.Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
|
||
c.Set("Access-Control-Allow-Headers", "Content-Type, Authorization, X-finder-TOKEN")
|
||
|
||
// AI能力调用路由,设置不同的 CORS 头
|
||
if strings.HasPrefix(c.Path(), "/api/v1/capability") {
|
||
c.Set("Access-Control-Allow-Headers", "Content-Type, X-Source-Key, X-Timestamp")
|
||
}
|
||
|
||
// 如果是预检请求(OPTIONS),直接返回 204
|
||
if c.Method() == "OPTIONS" {
|
||
return c.SendStatus(fiber.StatusNoContent) // 204
|
||
}
|
||
|
||
// 继续处理后续中间件或路由
|
||
return c.Next()
|
||
})
|
||
|
||
r := app.Group("api/v1/")
|
||
registerResponse(r)
|
||
r.Get("/health", func(c *fiber.Ctx) error {
|
||
c.Response().SetBody([]byte("1"))
|
||
return nil
|
||
})
|
||
|
||
// 微信代理路由已移至管理后台路由组内(JWT 鉴权 + 标准响应包装)
|
||
|
||
// WebSocket 实时通信端点(不走 JWT 中间件,Hub 内自行验证 token)
|
||
r.Get("advicer/ws", websocket.New(wsHub.WsHandler))
|
||
|
||
// 个人微信消息回调入口:/api/v1/advicer/wx/callback(上游调用,无 JWT,回调 token 校验在服务内完成)
|
||
r.Post("advicer/wx/callback", wxHook.Callback)
|
||
|
||
// 管理后台路由(需要登录): /api/v1/admin/...
|
||
adminR := r.Group("admin")
|
||
// 登录接口放行,其余接口校验 JWT 登录态
|
||
adminR.Use(middleware.AuthMiddleware(cfg.JwtSecret, "/api/v1/admin/advice/admin/login"))
|
||
AdvicerRouterRegist(
|
||
cfg,
|
||
adminR,
|
||
adviceFile,
|
||
adviceData,
|
||
adviceChat,
|
||
adviceProject,
|
||
adviceTalkSkill,
|
||
adviceClient,
|
||
industry,
|
||
admin,
|
||
modelSup,
|
||
activity,
|
||
wxHook,
|
||
smart,
|
||
wxProxy,
|
||
adviceCustomer,
|
||
adviceLabel,
|
||
)
|
||
|
||
}
|
||
|
||
func registerResponse(router fiber.Router) {
|
||
// 自定义返回
|
||
router.Use(func(c *fiber.Ctx) error {
|
||
err := c.Next()
|
||
return registerCommon(c, err)
|
||
})
|
||
|
||
}
|
||
|
||
func registerCommon(c *fiber.Ctx, err error) error {
|
||
// 调用下一个中间件或路由处理函数
|
||
if c.Path() == "/api/v1/qywx/callback" {
|
||
return nil
|
||
}
|
||
|
||
// 如果有错误发生
|
||
if err != nil {
|
||
bsErr, ok := err.(*errors.BusinessErr)
|
||
if !ok {
|
||
bsErr = errorcode.SysErr(err.Error())
|
||
}
|
||
// 返回自定义错误响应
|
||
return c.JSON(fiber.Map{
|
||
"message": bsErr.Error(),
|
||
"code": bsErr.Code(),
|
||
"data": nil,
|
||
})
|
||
}
|
||
contentType := strings.ToLower(string(c.Response().Header.Peek("Content-Type")))
|
||
if strings.Contains(strings.ToLower(contentType), "text/event-stream") {
|
||
// 是 SSE 请求
|
||
return c.SendString("这是 SSE 请求")
|
||
}
|
||
|
||
body := c.Response().Body()
|
||
if c.Locals("skip_response_wrap") == true {
|
||
// handler 已自行写入原始 JSON(如微信代理透传上游 data),原样输出不做包装
|
||
c.Set(fiber.HeaderContentType, fiber.MIMEApplicationJSONCharsetUTF8)
|
||
return c.Send(body)
|
||
}
|
||
var rawData json.RawMessage
|
||
if len(body) > 0 {
|
||
if err := json.Unmarshal(body, &rawData); err != nil {
|
||
// 解析失败,作为字符串包装成JSON
|
||
rawData = json.RawMessage(`"` + strings.ReplaceAll(string(body), `"`, `\"`) + `"`)
|
||
}
|
||
}
|
||
return c.JSON(fiber.Map{
|
||
"data": rawData,
|
||
"message": errors.Success.Error(),
|
||
"code": errors.Success.Code(),
|
||
})
|
||
}
|
||
|
||
// validateInst 全局复用 validator 实例(线程安全,避免每次请求重复创建)
|
||
var validateInst = func() *validator.Validate {
|
||
v := validator.New()
|
||
v.RegisterTagNameFunc(func(fld reflect.StructField) string {
|
||
name := fld.Tag.Get("zh")
|
||
if name == "" {
|
||
name = fld.Tag.Get("json")
|
||
}
|
||
return name
|
||
})
|
||
return v
|
||
}()
|
||
|
||
func Vali[T any](handler func(*fiber.Ctx, *T) error, _ *T) fiber.Handler {
|
||
return func(c *fiber.Ctx) error {
|
||
var data T
|
||
|
||
// 解析请求
|
||
if err := c.BodyParser(&data); err != nil {
|
||
return errorcode.ParamErr(err.Error())
|
||
}
|
||
|
||
// 验证
|
||
if err := validateInst.Struct(&data); err != nil {
|
||
if ve, ok := err.(validator.ValidationErrors); ok {
|
||
er := make([]string, len(ve))
|
||
for k, e := range ve {
|
||
er[k] = pkg.GetErr(e.Tag(), e.Field(), e.Param())
|
||
}
|
||
return errorcode.ParamErr(strings.Join(er, ","))
|
||
}
|
||
return errorcode.ParamErr(err.Error())
|
||
}
|
||
|
||
return handler(c, &data)
|
||
}
|
||
}
|