u
This commit is contained in:
1 parent
7c8e2906bc
commit
72a1a9184a
25 files changed
+744
-125
No files matched your search
@@ -1,27 +1,63 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
appConfig "base-framework/configs/app"
|
||||
"base-framework/internal/app"
|
||||
_ "github.com/lib/pq"
|
||||
"log"
|
||||
"net/http"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
func helloHandler(w http.ResponseWriter, r *http.Request) {
|
||||
_, err := fmt.Fprintf(w, "Hello, world!")
|
||||
func main() {
|
||||
_, err := appConfig.InitAppConfig("./configs/app/config.yaml")
|
||||
if err != nil {
|
||||
log.Printf("write response error: %v", err)
|
||||
log.Fatalf("加载配置失败: %v", err)
|
||||
}
|
||||
|
||||
log.Printf("服务器启动在端口 %d", appConfig.Server.Port)
|
||||
|
||||
r := app.InitAppRouter()
|
||||
|
||||
err = appConfig.InitDBConfig("D:\\db.yaml")
|
||||
if err != nil {
|
||||
log.Fatalf("加载数据库配置失败: %v", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
func main() {
|
||||
http.HandleFunc("/", helloHandler)
|
||||
|
||||
log.Printf("启动成功...")
|
||||
|
||||
err := http.ListenAndServe(":8082", nil)
|
||||
|
||||
err = http.ListenAndServe(":"+strconv.Itoa(appConfig.Server.Port), r)
|
||||
if err != nil {
|
||||
log.Fatalf("server start failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
//// 连接字符串格式
|
||||
//connStr := "host=117.72.182.135 port=5432 user=postgres password=zhang520.. dbname=test sslmode=disable"
|
||||
//
|
||||
//db, err := sql.Open("postgres", connStr)
|
||||
//if err != nil {
|
||||
// log.Fatal("打开数据库失败:", err)
|
||||
//}
|
||||
//defer db.Close()
|
||||
//
|
||||
//// 设置连接池参数(可选)
|
||||
//db.SetMaxOpenConns(20)
|
||||
//db.SetMaxIdleConns(5)
|
||||
//db.SetConnMaxLifetime(0)
|
||||
//
|
||||
//// 测试连接
|
||||
//err = db.Ping()
|
||||
//if err != nil {
|
||||
// log.Fatal("数据库连接失败:", err)
|
||||
//}
|
||||
//
|
||||
//fmt.Println("数据库连接成功!")
|
||||
//
|
||||
//// 查询示例
|
||||
//var version string
|
||||
//err = db.QueryRow("SELECT version()").Scan(&version)
|
||||
//if err != nil {
|
||||
// log.Fatal(err)
|
||||
//}
|
||||
//fmt.Println("PostgreSQL version:", version)
|
||||
File renamed without changes.
@@ -0,0 +1,40 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"github.com/spf13/viper"
|
||||
"time"
|
||||
)
|
||||
|
||||
type ServerConfig struct {
|
||||
Port int
|
||||
}
|
||||
|
||||
type JWTConfig struct {
|
||||
Secret string
|
||||
Expiry time.Duration
|
||||
}
|
||||
|
||||
var (
|
||||
Server ServerConfig
|
||||
JWT JWTConfig
|
||||
)
|
||||
|
||||
func InitAppConfig(configPath string) (error, error) {
|
||||
v := viper.New()
|
||||
v.SetConfigFile(configPath)
|
||||
v.SetConfigType("yaml")
|
||||
|
||||
if err := v.ReadInConfig(); err != nil {
|
||||
return err, nil
|
||||
}
|
||||
|
||||
if err := v.UnmarshalKey("server", &Server); err != nil {
|
||||
return err, nil
|
||||
}
|
||||
|
||||
if err := v.UnmarshalKey("jwt", &JWT); err != nil {
|
||||
return err, nil
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
server:
|
||||
port: 8082
|
||||
jwt:
|
||||
secret: 3Bde3BGEbYqtqyEUzW3ry8jKFcaPH17fRmTmqE7MDr05Lwj95uruRKrrkb44TJ4s
|
||||
expiry: 43200 # 12 * 60 * 60 秒过期
|
||||
@@ -0,0 +1,80 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"sync"
|
||||
|
||||
"github.com/fsnotify/fsnotify"
|
||||
"github.com/spf13/viper"
|
||||
)
|
||||
|
||||
// DBConfig 数据库单个配置结构
|
||||
type DBConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
Password string
|
||||
Dbname string
|
||||
}
|
||||
|
||||
// dbConfigs 全局变量,存放所有数据库配置,key为配置名,如 test1, test2
|
||||
var (
|
||||
dbConfigs map[string]DBConfig
|
||||
mu sync.RWMutex
|
||||
v *viper.Viper
|
||||
)
|
||||
|
||||
// InitDBConfig 初始化并加载 db.yaml 配置,同时启动监听
|
||||
func InitDBConfig(configPath string) error {
|
||||
v = viper.New()
|
||||
v.SetConfigFile(configPath)
|
||||
v.SetConfigType("yaml")
|
||||
|
||||
if err := v.ReadInConfig(); err != nil {
|
||||
return fmt.Errorf("读取数据库配置失败: %w", err)
|
||||
}
|
||||
|
||||
if err := unmarshalConfigs(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 监听配置文件变化
|
||||
v.WatchConfig()
|
||||
v.OnConfigChange(func(e fsnotify.Event) {
|
||||
log.Printf("数据库配置文件发生变化: %s\n", e.Name)
|
||||
if err := unmarshalConfigs(); err != nil {
|
||||
log.Printf("重新加载数据库配置失败: %v\n", err)
|
||||
} else {
|
||||
log.Println("数据库配置已更新")
|
||||
}
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// unmarshalConfigs 解析配置到全局变量,内部加锁保证并发安全
|
||||
func unmarshalConfigs() error {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
temp := make(map[string]DBConfig)
|
||||
if err := v.Unmarshal(&temp); err != nil {
|
||||
return fmt.Errorf("解析数据库配置失败: %w", err)
|
||||
}
|
||||
dbConfigs = temp
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetDBConfigs 并发安全地返回当前的所有数据库配置副本
|
||||
func GetDBConfigs() map[string]DBConfig {
|
||||
mu.RLock()
|
||||
defer mu.RUnlock()
|
||||
|
||||
// 返回一份拷贝,避免外部修改内部数据
|
||||
res := make(map[string]DBConfig, len(dbConfigs))
|
||||
for k, v := range dbConfigs {
|
||||
res[k] = v
|
||||
}
|
||||
return res
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
test1:
|
||||
host: 127.0.0.1
|
||||
port: 5432
|
||||
username: admin
|
||||
password: 123456
|
||||
dbname: test1
|
||||
test2:
|
||||
host: 127.0.0.1
|
||||
port: 5432
|
||||
username: admin
|
||||
password: 123456
|
||||
dbname: test1
|
||||
Whitespace-only changes.
@@ -1,3 +1,23 @@
|
||||
module base-framework
|
||||
|
||||
go 1.24
|
||||
|
||||
require (
|
||||
github.com/fsnotify/fsnotify v1.8.0 // indirect
|
||||
github.com/go-viper/mapstructure/v2 v2.2.1 // indirect
|
||||
github.com/golang-jwt/jwt v3.2.2+incompatible // indirect
|
||||
github.com/lib/pq v1.10.9 // indirect
|
||||
github.com/pelletier/go-toml/v2 v2.2.3 // indirect
|
||||
github.com/sagikazarmark/locafero v0.7.0 // indirect
|
||||
github.com/sourcegraph/conc v0.3.0 // indirect
|
||||
github.com/spf13/afero v1.12.0 // indirect
|
||||
github.com/spf13/cast v1.7.1 // indirect
|
||||
github.com/spf13/pflag v1.0.6 // indirect
|
||||
github.com/spf13/viper v1.20.1 // indirect
|
||||
github.com/subosito/gotenv v1.6.0 // indirect
|
||||
go.uber.org/atomic v1.9.0 // indirect
|
||||
go.uber.org/multierr v1.9.0 // indirect
|
||||
golang.org/x/sys v0.29.0 // indirect
|
||||
golang.org/x/text v0.21.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
)
|
||||
@@ -0,0 +1,40 @@
|
||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/fsnotify/fsnotify v1.8.0 h1:dAwr6QBTBZIkG8roQaJjGof0pp0EeF+tNV7YBP3F/8M=
|
||||
github.com/fsnotify/fsnotify v1.8.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
|
||||
github.com/go-viper/mapstructure/v2 v2.2.1 h1:ZAaOCxANMuZx5RCeg0mBdEZk7DZasvvZIxtHqx8aGss=
|
||||
github.com/go-viper/mapstructure/v2 v2.2.1/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
||||
github.com/golang-jwt/jwt v3.2.2+incompatible h1:IfV12K8xAKAnZqdXVzCZ+TOjboZ2keLg81eXfW3O+oY=
|
||||
github.com/golang-jwt/jwt v3.2.2+incompatible/go.mod h1:8pz2t5EyA70fFQQSrl6XZXzqecmYZeUEB8OUGHkxJ+I=
|
||||
github.com/lib/pq v1.10.9 h1:YXG7RB+JIjhP29X+OtkiDnYaXQwpS4JEWq7dtCCRUEw=
|
||||
github.com/lib/pq v1.10.9/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o=
|
||||
github.com/pelletier/go-toml/v2 v2.2.3 h1:YmeHyLY8mFWbdkNWwpr+qIL2bEqT0o95WSdkNHvL12M=
|
||||
github.com/pelletier/go-toml/v2 v2.2.3/go.mod h1:MfCQTFTvCcUyyvvwm1+G6H/jORL20Xlb6rzQu9GuUkc=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/sagikazarmark/locafero v0.7.0 h1:5MqpDsTGNDhY8sGp0Aowyf0qKsPrhewaLSsFaodPcyo=
|
||||
github.com/sagikazarmark/locafero v0.7.0/go.mod h1:2za3Cg5rMaTMoG/2Ulr9AwtFaIppKXTRYnozin4aB5k=
|
||||
github.com/sourcegraph/conc v0.3.0 h1:OQTbbt6P72L20UqAkXXuLOj79LfEanQ+YQFNpLA9ySo=
|
||||
github.com/sourcegraph/conc v0.3.0/go.mod h1:Sdozi7LEKbFPqYX2/J+iBAM6HpqSLTASQIKqDmF7Mt0=
|
||||
github.com/spf13/afero v1.12.0 h1:UcOPyRBYczmFn6yvphxkn9ZEOY65cpwGKb5mL36mrqs=
|
||||
github.com/spf13/afero v1.12.0/go.mod h1:ZTlWwG4/ahT8W7T0WQ5uYmjI9duaLQGy3Q2OAl4sk/4=
|
||||
github.com/spf13/cast v1.7.1 h1:cuNEagBQEHWN1FnbGEjCXL2szYEXqfJPbP2HNUaca9Y=
|
||||
github.com/spf13/cast v1.7.1/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo=
|
||||
github.com/spf13/pflag v1.0.6 h1:jFzHGLGAlb3ruxLB8MhbI6A8+AQX/2eW4qeyNZXNp2o=
|
||||
github.com/spf13/pflag v1.0.6/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
github.com/spf13/viper v1.20.1 h1:ZMi+z/lvLyPSCoNtFCpqjy0S4kPbirhpTMwl8BkW9X4=
|
||||
github.com/spf13/viper v1.20.1/go.mod h1:P9Mdzt1zoHIG8m2eZQinpiBjo6kCmZSKBClNNqjJvu4=
|
||||
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||
github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8=
|
||||
github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU=
|
||||
go.uber.org/atomic v1.9.0 h1:ECmE8Bn/WFTYwEW/bpKD3M8VtR/zQVbavAoalC1PYyE=
|
||||
go.uber.org/atomic v1.9.0/go.mod h1:fEN4uk6kAWBTFdckzkM89CLk9XfWZrxpCo0nPH17wJc=
|
||||
go.uber.org/multierr v1.9.0 h1:7fIwc/ZtS0q++VgcfqFDxSBZVv/Xo49/SYnDFupUwlI=
|
||||
go.uber.org/multierr v1.9.0/go.mod h1:X2jQV1h+kxSjClGpnseKVIxpmcjrj7MNnI0bnlfKTVQ=
|
||||
golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU=
|
||||
golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/text v0.21.0 h1:zyQAAkrwaneQ066sspRyJaG9VNi/YJ1NfzcGB3hZ/qo=
|
||||
golang.org/x/text v0.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
@@ -0,0 +1,10 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"base-framework/internal/app/handle"
|
||||
"base-framework/pkg/router"
|
||||
)
|
||||
|
||||
func InitTest(r *router.Router) {
|
||||
r.GET("/test", handle.Test)
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
package handle
|
||||
|
||||
import (
|
||||
"base-framework/pkg/router"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
func Test(c *router.Context) {
|
||||
c.JSON(http.StatusOK, map[string]string{
|
||||
"user_id": c.Param("id"),
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"base-framework/internal/app/api"
|
||||
"base-framework/pkg/middleware"
|
||||
"base-framework/pkg/router"
|
||||
)
|
||||
|
||||
func InitAppRouter() *router.Router {
|
||||
r := router.NewRouter()
|
||||
|
||||
r.Use(middleware.Recover())
|
||||
r.Use(middleware.Auth()).ExcludePaths("/login")
|
||||
r.Use(middleware.Logger())
|
||||
|
||||
initApi(r)
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
func initApi(r *router.Router) {
|
||||
api.InitTest(r)
|
||||
}
|
||||
@@ -1,4 +0,0 @@
|
||||
package app
|
||||
|
||||
func RunServer() {
|
||||
}
|
||||
@@ -1,11 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// LoggingMiddleware logs incoming requests
|
||||
func LoggingMiddleware(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Printf("Request received: %s %s\n", r.Method, r.URL.Path)
|
||||
}
|
||||
@@ -1,98 +0,0 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// HandlerFunc defines the signature of a handler function
|
||||
type HandlerFunc func(w http.ResponseWriter, r *http.Request)
|
||||
|
||||
// Router struct holds all the routes, middlewares, and groups
|
||||
type Router struct {
|
||||
routes map[string]map[string]HandlerFunc
|
||||
middleware []HandlerFunc
|
||||
groups map[string]*Group
|
||||
}
|
||||
|
||||
// NewRouter creates a new Router instance
|
||||
func NewRouter() *Router {
|
||||
return &Router{
|
||||
routes: make(map[string]map[string]HandlerFunc),
|
||||
groups: make(map[string]*Group),
|
||||
}
|
||||
}
|
||||
|
||||
// Use adds a global middleware to the router
|
||||
func (r *Router) Use(middleware HandlerFunc) {
|
||||
r.middleware = append(r.middleware, middleware)
|
||||
}
|
||||
|
||||
// Handle registers a route with a method and a handler
|
||||
func (r *Router) Handle(method, path string, handler HandlerFunc) {
|
||||
if _, exists := r.routes[method]; !exists {
|
||||
r.routes[method] = make(map[string]HandlerFunc)
|
||||
}
|
||||
r.routes[method][path] = handler
|
||||
}
|
||||
|
||||
// Group creates a new group with a prefix
|
||||
func (r *Router) Group(prefix string) *Group {
|
||||
group := &Group{
|
||||
router: r,
|
||||
prefix: prefix,
|
||||
}
|
||||
r.groups[prefix] = group
|
||||
return group
|
||||
}
|
||||
|
||||
// Group struct represents a group of routes with a common prefix
|
||||
type Group struct {
|
||||
router *Router
|
||||
prefix string
|
||||
middleware []HandlerFunc
|
||||
}
|
||||
|
||||
// Use adds a middleware to a group
|
||||
func (g *Group) Use(middleware HandlerFunc) {
|
||||
g.middleware = append(g.middleware, middleware)
|
||||
}
|
||||
|
||||
// Handle registers a route in the group
|
||||
func (g *Group) Handle(method, path string, handler HandlerFunc) {
|
||||
fullPath := g.prefix + path
|
||||
|
||||
// Apply group middleware first
|
||||
for _, m := range g.middleware {
|
||||
g.router.Use(m)
|
||||
}
|
||||
|
||||
// Register the route
|
||||
g.router.Handle(method, fullPath, handler)
|
||||
}
|
||||
|
||||
// Get registers a GET route in the group
|
||||
func (g *Group) Get(path string, handler HandlerFunc) {
|
||||
g.Handle("GET", path, handler)
|
||||
}
|
||||
|
||||
// Post registers a POST route in the group
|
||||
func (g *Group) Post(path string, handler HandlerFunc) {
|
||||
g.Handle("POST", path, handler)
|
||||
}
|
||||
|
||||
// ServeHTTP is the main entry point for handling HTTP requests
|
||||
func (r *Router) ServeHTTP(w http.ResponseWriter, req *http.Request) {
|
||||
// Apply global middlewares
|
||||
for _, m := range r.middleware {
|
||||
m(w, req)
|
||||
}
|
||||
|
||||
// Route matching
|
||||
if handlers, ok := r.routes[req.Method]; ok {
|
||||
if handler, ok := handlers[req.URL.Path]; ok {
|
||||
handler(w, req)
|
||||
return
|
||||
}
|
||||
}
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
File renamed without changes.
@@ -0,0 +1,55 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"base-framework/pkg/router"
|
||||
"base-framework/pkg/utils"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func Auth() router.HandlerFunc {
|
||||
return func(c *router.Context) {
|
||||
tokenHeader := c.Header("Authorization")
|
||||
userIdHeader := c.Header("user_id")
|
||||
orgIDHeader := c.Header("org_id")
|
||||
|
||||
if tokenHeader == "" {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "Authorization header missing"})
|
||||
return
|
||||
}
|
||||
if userIdHeader == "" {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "user_id header missing"})
|
||||
return
|
||||
}
|
||||
if orgIDHeader == "" {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "org_id header missing"})
|
||||
return
|
||||
}
|
||||
|
||||
parts := strings.Fields(tokenHeader)
|
||||
if len(parts) != 2 || strings.ToLower(parts[0]) != "bearer" {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "Authorization header format must be Bearer {token}"})
|
||||
return
|
||||
}
|
||||
|
||||
tokenStr := parts[1]
|
||||
claims, err := utils.VerifyToken(tokenStr)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "Invalid token: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
if claims.UserID != userIdHeader {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "user_id does not match token"})
|
||||
return
|
||||
}
|
||||
|
||||
if claims.OrgID != orgIDHeader {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "org_id does not match token"})
|
||||
return
|
||||
}
|
||||
|
||||
// 认证通过,继续执行后续中间件或处理器
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"base-framework/pkg/router"
|
||||
"log"
|
||||
"time"
|
||||
)
|
||||
|
||||
func Logger() router.HandlerFunc {
|
||||
return func(c *router.Context) {
|
||||
start := time.Now()
|
||||
c.Next()
|
||||
log.Printf("[%s] %s in %v", c.Request.Method, c.Request.URL.Path, time.Since(start))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"base-framework/pkg/router"
|
||||
"log"
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
)
|
||||
|
||||
func Recover() router.HandlerFunc {
|
||||
return func(c *router.Context) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
log.Printf("[PANIC] %v\n%s", err, debug.Stack())
|
||||
c.JSON(http.StatusInternalServerError, map[string]string{
|
||||
"error": "Internal Server Error",
|
||||
})
|
||||
}
|
||||
}()
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
File renamed without changes.
@@ -0,0 +1,42 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
type Context struct {
|
||||
Writer http.ResponseWriter
|
||||
Request *http.Request
|
||||
Params map[string]string
|
||||
index int
|
||||
handlers []HandlerFunc
|
||||
}
|
||||
|
||||
// Next 执行下一个中间件或处理函数
|
||||
func (c *Context) Next() {
|
||||
c.index++
|
||||
for c.index < len(c.handlers) {
|
||||
c.handlers[c.index](c)
|
||||
c.index++
|
||||
}
|
||||
}
|
||||
|
||||
// Param 获取路由参数
|
||||
func (c *Context) Param(key string) string {
|
||||
return c.Params[key]
|
||||
}
|
||||
|
||||
// Header 获取请求头
|
||||
func (c *Context) Header(key string) string {
|
||||
return c.Request.Header.Get(key)
|
||||
}
|
||||
|
||||
// JSON 返回JSON格式响应
|
||||
func (c *Context) JSON(statusCode int, data interface{}) {
|
||||
c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
c.Writer.WriteHeader(statusCode)
|
||||
if err := json.NewEncoder(c.Writer).Encode(data); err != nil {
|
||||
http.Error(c.Writer, err.Error(), http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,257 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// HandlerFunc 路由和中间件处理函数签名
|
||||
type HandlerFunc func(*Context)
|
||||
|
||||
// node 路由树节点
|
||||
type node struct {
|
||||
pattern string
|
||||
part string
|
||||
children []*node
|
||||
isParam bool
|
||||
handler HandlerFunc
|
||||
}
|
||||
|
||||
// 匹配子节点,返回第一个匹配的
|
||||
func (n *node) matchChild(part string) *node {
|
||||
for _, child := range n.children {
|
||||
if child.part == part || child.isParam {
|
||||
return child
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// 匹配所有匹配的子节点
|
||||
func (n *node) matchChildren(part string) []*node {
|
||||
nodes := make([]*node, 0)
|
||||
for _, child := range n.children {
|
||||
if child.part == part || child.isParam {
|
||||
nodes = append(nodes, child)
|
||||
}
|
||||
}
|
||||
return nodes
|
||||
}
|
||||
|
||||
// 插入路由节点
|
||||
func (n *node) insert(pattern string, parts []string, height int, handler HandlerFunc) {
|
||||
if len(parts) == height {
|
||||
n.pattern = pattern
|
||||
n.handler = handler
|
||||
return
|
||||
}
|
||||
part := parts[height]
|
||||
child := n.matchChild(part)
|
||||
if child == nil {
|
||||
child = &node{
|
||||
part: part,
|
||||
isParam: len(part) > 0 && part[0] == ':',
|
||||
}
|
||||
n.children = append(n.children, child)
|
||||
}
|
||||
child.insert(pattern, parts, height+1, handler)
|
||||
}
|
||||
|
||||
// 搜索路由节点,并收集参数
|
||||
func (n *node) search(parts []string, height int, params map[string]string) *node {
|
||||
if len(parts) == height || n.part == "*" {
|
||||
if n.pattern == "" {
|
||||
return nil
|
||||
}
|
||||
return n
|
||||
}
|
||||
part := parts[height]
|
||||
children := n.matchChildren(part)
|
||||
|
||||
for _, child := range children {
|
||||
if child.isParam {
|
||||
params[child.part[1:]] = part
|
||||
}
|
||||
res := child.search(parts, height+1, params)
|
||||
if res != nil {
|
||||
return res
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// 中间件条目,包含排除路径集合
|
||||
type middlewareEntry struct {
|
||||
handler HandlerFunc
|
||||
excludePaths map[string]struct{}
|
||||
}
|
||||
|
||||
// Router 路由器结构
|
||||
type Router struct {
|
||||
roots map[string]*node
|
||||
middleware []middlewareEntry
|
||||
basePath string
|
||||
}
|
||||
|
||||
// NewRouter 创建路由器实例
|
||||
func NewRouter() *Router {
|
||||
return &Router{
|
||||
roots: make(map[string]*node),
|
||||
}
|
||||
}
|
||||
|
||||
// MiddlewareHandle 用于链式配置中间件排除路径
|
||||
type MiddlewareHandle struct {
|
||||
router *Router
|
||||
entryIdx int
|
||||
}
|
||||
|
||||
// Use 注册中间件,返回链式配置句柄
|
||||
func (r *Router) Use(handler HandlerFunc) *MiddlewareHandle {
|
||||
entry := middlewareEntry{
|
||||
handler: handler,
|
||||
excludePaths: nil,
|
||||
}
|
||||
r.middleware = append(r.middleware, entry)
|
||||
return &MiddlewareHandle{
|
||||
router: r,
|
||||
entryIdx: len(r.middleware) - 1,
|
||||
}
|
||||
}
|
||||
|
||||
// ExcludePaths 设置排除路径
|
||||
func (mh *MiddlewareHandle) ExcludePaths(paths ...string) *Router {
|
||||
excludeMap := make(map[string]struct{}, len(paths))
|
||||
for _, p := range paths {
|
||||
excludeMap[p] = struct{}{}
|
||||
}
|
||||
mh.router.middleware[mh.entryIdx].excludePaths = excludeMap
|
||||
return mh.router
|
||||
}
|
||||
|
||||
// Group 路由分组,返回新的 Router,继承根节点和中间件,basePath拼接
|
||||
func (r *Router) Group(prefix string, m ...HandlerFunc) *Router {
|
||||
// 复制原middleware
|
||||
newMiddleware := make([]middlewareEntry, len(r.middleware))
|
||||
copy(newMiddleware, r.middleware)
|
||||
|
||||
// 将新增的 HandlerFunc 转成 middlewareEntry
|
||||
for _, handler := range m {
|
||||
entry := middlewareEntry{
|
||||
handler: handler,
|
||||
excludePaths: nil,
|
||||
}
|
||||
newMiddleware = append(newMiddleware, entry)
|
||||
}
|
||||
|
||||
return &Router{
|
||||
roots: r.roots,
|
||||
middleware: newMiddleware,
|
||||
basePath: joinPaths(r.basePath, prefix),
|
||||
}
|
||||
}
|
||||
|
||||
// 辅助函数,提取middlewareEntry中的handler为HandlerFunc slice
|
||||
func (r *Router) middlewareHandlers() []HandlerFunc {
|
||||
handlers := make([]HandlerFunc, 0, len(r.middleware))
|
||||
for _, m := range r.middleware {
|
||||
handlers = append(handlers, m.handler)
|
||||
}
|
||||
return handlers
|
||||
}
|
||||
|
||||
// GET 注册GET请求路由
|
||||
func (r *Router) GET(path string, handler HandlerFunc) *Router {
|
||||
return r.handle("GET", path, handler)
|
||||
}
|
||||
|
||||
// POST 注册POST请求路由
|
||||
func (r *Router) POST(path string, handler HandlerFunc) *Router {
|
||||
return r.handle("POST", path, handler)
|
||||
}
|
||||
|
||||
// handle 注册具体方法和路径的处理函数
|
||||
func (r *Router) handle(method, path string, handler HandlerFunc) *Router {
|
||||
fullPath := joinPaths(r.basePath, path)
|
||||
parts := parsePattern(fullPath)
|
||||
if r.roots[method] == nil {
|
||||
r.roots[method] = &node{}
|
||||
}
|
||||
r.roots[method].insert(fullPath, parts, 0, handler)
|
||||
return r
|
||||
}
|
||||
|
||||
// ServeHTTP 实现 http.Handler,执行路由匹配和中间件
|
||||
func (r *Router) ServeHTTP(w http.ResponseWriter, req *http.Request) {
|
||||
root := r.roots[req.Method]
|
||||
if root == nil {
|
||||
http.NotFound(w, req)
|
||||
return
|
||||
}
|
||||
searchParts := parsePattern(req.URL.Path)
|
||||
params := make(map[string]string)
|
||||
n := root.search(searchParts, 0, params)
|
||||
if n == nil || n.handler == nil {
|
||||
http.NotFound(w, req)
|
||||
return
|
||||
}
|
||||
|
||||
// 按排除路径过滤中间件
|
||||
handlers := make([]HandlerFunc, 0)
|
||||
for _, m := range r.middleware {
|
||||
if m.excludePaths != nil {
|
||||
if _, excluded := m.excludePaths[req.URL.Path]; excluded {
|
||||
continue
|
||||
}
|
||||
}
|
||||
handlers = append(handlers, m.handler)
|
||||
}
|
||||
handlers = append(handlers, n.handler)
|
||||
|
||||
c := &Context{
|
||||
Writer: w,
|
||||
Request: req,
|
||||
Params: params,
|
||||
index: -1,
|
||||
handlers: handlers,
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
|
||||
// 解析路径,去除空字符串
|
||||
func parsePattern(pattern string) []string {
|
||||
vs := strings.Split(strings.Trim(pattern, "/"), "/")
|
||||
parts := make([]string, 0, len(vs))
|
||||
for _, item := range vs {
|
||||
if item != "" {
|
||||
parts = append(parts, item)
|
||||
}
|
||||
}
|
||||
return parts
|
||||
}
|
||||
|
||||
// 拼接两个路径字符串,保证中间只有一个 '/'
|
||||
func joinPaths(a, b string) string {
|
||||
if a == "" {
|
||||
if b == "" {
|
||||
return "/"
|
||||
}
|
||||
if !strings.HasPrefix(b, "/") {
|
||||
return "/" + b
|
||||
}
|
||||
return b
|
||||
}
|
||||
if b == "" {
|
||||
return a
|
||||
}
|
||||
aslash := strings.HasSuffix(a, "/")
|
||||
bslash := strings.HasPrefix(b, "/")
|
||||
switch {
|
||||
case aslash && bslash:
|
||||
return a + b[1:]
|
||||
case !aslash && !bslash:
|
||||
return a + "/" + b
|
||||
default:
|
||||
return a + b
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
appConfig "base-framework/configs/app"
|
||||
"errors"
|
||||
"github.com/golang-jwt/jwt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// CustomClaims 定义自己的 payload 结构,可以根据需要扩展
|
||||
type CustomClaims struct {
|
||||
OrgID string `json:"org_id"`
|
||||
UserID string `json:"user_id"`
|
||||
jwt.StandardClaims
|
||||
}
|
||||
|
||||
// CreateToken 创建一个JWT token
|
||||
func CreateToken(orgID, userID string) (string, error) {
|
||||
expireTime := time.Now().Add(appConfig.JWT.Expiry).Unix()
|
||||
|
||||
claims := CustomClaims{
|
||||
OrgID: orgID,
|
||||
UserID: userID,
|
||||
StandardClaims: jwt.StandardClaims{
|
||||
ExpiresAt: expireTime,
|
||||
IssuedAt: time.Now().Unix(),
|
||||
Issuer: "your-app-name", // 可以自定义
|
||||
},
|
||||
}
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
signKey := []byte(appConfig.JWT.Secret)
|
||||
return token.SignedString(signKey)
|
||||
}
|
||||
|
||||
// VerifyToken 验证并解析 JWT token
|
||||
func VerifyToken(tokenString string) (*CustomClaims, error) {
|
||||
signKey := []byte(appConfig.JWT.Secret)
|
||||
|
||||
token, err := jwt.ParseWithClaims(tokenString, &CustomClaims{}, func(token *jwt.Token) (interface{}, error) {
|
||||
// 校验签名算法
|
||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, errors.New("unexpected signing method")
|
||||
}
|
||||
return signKey, nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if claims, ok := token.Claims.(*CustomClaims); ok && token.Valid {
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
return nil, errors.New("invalid token")
|
||||
}
|
||||
Reference in new issue
Block a user