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, ) { 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, ) } 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) } }