ai_scheduler/internal/server/router/router.go

182 lines
5.1 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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