feat: 完善后端脚手架基础能力
This commit is contained in:
@@ -0,0 +1,154 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"skeleton/database"
|
||||
"skeleton/middlewares"
|
||||
"skeleton/models"
|
||||
"skeleton/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type Controller struct{}
|
||||
|
||||
func NewController() *Controller { return &Controller{} }
|
||||
|
||||
// Register godoc
|
||||
// @Summary 注册示例用户
|
||||
// @Tags auth
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param body body Credentials true "注册信息"
|
||||
// @Success 201 {object} utils.APIResponse{data=UserResponse}
|
||||
// @Failure 400 {object} utils.APIResponse
|
||||
// @Failure 409 {object} utils.APIResponse
|
||||
// @Router /auth/register [post]
|
||||
func (c *Controller) Register(ctx *gin.Context) {
|
||||
var input Credentials
|
||||
if err := ctx.ShouldBindJSON(&input); err != nil {
|
||||
ctx.JSON(http.StatusBadRequest, utils.Failure(400, err.Error()))
|
||||
return
|
||||
}
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(input.Password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, utils.Failure(500, "密码处理失败"))
|
||||
return
|
||||
}
|
||||
user := models.User{Username: input.Username, Password: string(hash)}
|
||||
if err := database.GetDB().Create(&user).Error; err != nil {
|
||||
ctx.JSON(http.StatusConflict, utils.Failure(409, "用户名已存在或数据无效"))
|
||||
return
|
||||
}
|
||||
ctx.JSON(http.StatusCreated, utils.Success(UserResponse{ID: user.ID, Username: user.Username}))
|
||||
}
|
||||
|
||||
// Login godoc
|
||||
// @Summary 用户登录
|
||||
// @Tags auth
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param body body Credentials true "登录信息"
|
||||
// @Success 200 {object} utils.APIResponse{data=middlewares.TokenPair}
|
||||
// @Failure 401 {object} utils.APIResponse
|
||||
// @Router /auth/login [post]
|
||||
func (c *Controller) Login(ctx *gin.Context) {
|
||||
var input Credentials
|
||||
if err := ctx.ShouldBindJSON(&input); err != nil {
|
||||
ctx.JSON(http.StatusBadRequest, utils.Failure(400, err.Error()))
|
||||
return
|
||||
}
|
||||
var user models.User
|
||||
err := database.GetDB().Where("username = ?", input.Username).First(&user).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
ctx.JSON(http.StatusUnauthorized, utils.Failure(401, "用户名或密码错误"))
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, utils.Failure(500, "查询用户失败"))
|
||||
return
|
||||
}
|
||||
if bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(input.Password)) != nil {
|
||||
ctx.JSON(http.StatusUnauthorized, utils.Failure(401, "用户名或密码错误"))
|
||||
return
|
||||
}
|
||||
pair, err := middlewares.GenerateTokenPair(int(user.ID), user.Username)
|
||||
if err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, utils.Failure(500, "生成 token 失败"))
|
||||
return
|
||||
}
|
||||
ctx.JSON(http.StatusOK, utils.Success(pair))
|
||||
}
|
||||
|
||||
// Refresh godoc
|
||||
// @Summary 刷新并轮换 Token
|
||||
// @Tags auth
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param body body RefreshRequest true "Refresh Token"
|
||||
// @Success 200 {object} utils.APIResponse{data=middlewares.TokenPair}
|
||||
// @Failure 401 {object} utils.APIResponse
|
||||
// @Router /auth/refresh [post]
|
||||
func (c *Controller) Refresh(ctx *gin.Context) {
|
||||
var input RefreshRequest
|
||||
if err := ctx.ShouldBindJSON(&input); err != nil {
|
||||
ctx.JSON(http.StatusBadRequest, utils.Failure(400, err.Error()))
|
||||
return
|
||||
}
|
||||
claims, err := middlewares.ParseToken(input.RefreshToken)
|
||||
if err != nil || claims.TokenType != middlewares.TokenTypeRefresh {
|
||||
ctx.JSON(http.StatusUnauthorized, utils.Failure(401, "refresh token 无效或已过期"))
|
||||
return
|
||||
}
|
||||
revoked, err := database.IsTokenRevoked(ctx.Request.Context(), claims.ID)
|
||||
if err != nil || revoked {
|
||||
ctx.JSON(http.StatusUnauthorized, utils.Failure(401, "refresh token 已撤销"))
|
||||
return
|
||||
}
|
||||
if err := revokeClaims(ctx, claims); err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, utils.Failure(500, "撤销旧 token 失败"))
|
||||
return
|
||||
}
|
||||
pair, err := middlewares.GenerateTokenPair(claims.UserID, claims.Username)
|
||||
if err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, utils.Failure(500, "生成 token 失败"))
|
||||
return
|
||||
}
|
||||
ctx.JSON(http.StatusOK, utils.Success(pair))
|
||||
}
|
||||
|
||||
// Logout godoc
|
||||
// @Summary 注销并撤销 Token
|
||||
// @Tags auth
|
||||
// @Security BearerAuth
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param body body LogoutRequest false "可选 Refresh Token"
|
||||
// @Success 200 {object} utils.APIResponse
|
||||
// @Router /private/auth/logout [post]
|
||||
func (c *Controller) Logout(ctx *gin.Context) {
|
||||
claims, _ := middlewares.GetCurrentClaims(ctx)
|
||||
if err := revokeClaims(ctx, claims); err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, utils.Failure(500, "撤销 access token 失败"))
|
||||
return
|
||||
}
|
||||
var input LogoutRequest
|
||||
if ctx.ShouldBindJSON(&input) == nil && input.RefreshToken != "" {
|
||||
if refreshClaims, err := middlewares.ParseToken(input.RefreshToken); err == nil && refreshClaims.TokenType == middlewares.TokenTypeRefresh {
|
||||
_ = revokeClaims(ctx, refreshClaims)
|
||||
}
|
||||
}
|
||||
ctx.JSON(http.StatusOK, utils.Success(nil))
|
||||
}
|
||||
|
||||
func revokeClaims(ctx *gin.Context, claims *middlewares.JWTClaims) error {
|
||||
if claims == nil || claims.ExpiresAt == nil {
|
||||
return nil
|
||||
}
|
||||
return database.RevokeToken(ctx.Request.Context(), claims.ID, time.Until(claims.ExpiresAt.Time))
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package auth
|
||||
|
||||
type Credentials struct {
|
||||
Username string `json:"username" binding:"required,min=3,max=64" example:"demo"`
|
||||
Password string `json:"password" binding:"required,min=8,max=72" example:"change-me-123"`
|
||||
}
|
||||
|
||||
type RefreshRequest struct {
|
||||
RefreshToken string `json:"refresh_token" binding:"required"`
|
||||
}
|
||||
|
||||
type LogoutRequest struct {
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
}
|
||||
|
||||
type UserResponse struct {
|
||||
ID uint `json:"id"`
|
||||
Username string `json:"username"`
|
||||
}
|
||||
@@ -20,7 +20,12 @@ func NewExampleController() *ExampleController {
|
||||
}
|
||||
}
|
||||
|
||||
// Hello 示例接口
|
||||
// Hello godoc
|
||||
// @Summary 返回示例消息
|
||||
// @Tags example
|
||||
// @Produce json
|
||||
// @Success 200 {object} utils.APIResponse{data=HelloResponse}
|
||||
// @Router /example/hello [get]
|
||||
func (c *ExampleController) Hello(ctx *gin.Context) {
|
||||
ctx.JSON(http.StatusOK, utils.Success(c.service.Hello()))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,96 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
const (
|
||||
writeWait = 10 * time.Second
|
||||
pongWait = 60 * time.Second
|
||||
pingPeriod = pongWait * 9 / 10
|
||||
maxMessageSize = 8 * 1024
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
hub *Hub
|
||||
conn *websocket.Conn
|
||||
send chan OutgoingMessage
|
||||
room string
|
||||
clientID string
|
||||
}
|
||||
|
||||
func NewClient(hub *Hub, conn *websocket.Conn, room, clientID string) *Client {
|
||||
return &Client{
|
||||
hub: hub, conn: conn, send: make(chan OutgoingMessage, 64),
|
||||
room: room, clientID: clientID,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) ReadPump() {
|
||||
defer func() {
|
||||
remaining := c.hub.Unregister(c)
|
||||
_ = c.conn.Close()
|
||||
c.hub.Broadcast(c.room, presenceMessage(c.room, c.clientID, "left", remaining))
|
||||
}()
|
||||
|
||||
c.conn.SetReadLimit(maxMessageSize)
|
||||
_ = c.conn.SetReadDeadline(time.Now().Add(pongWait))
|
||||
c.conn.SetPongHandler(func(string) error {
|
||||
return c.conn.SetReadDeadline(time.Now().Add(pongWait))
|
||||
})
|
||||
|
||||
for {
|
||||
var incoming IncomingMessage
|
||||
if err := c.conn.ReadJSON(&incoming); err != nil {
|
||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) {
|
||||
zap.L().Warn("WebSocket读取失败", zap.String("client_id", c.clientID), zap.Error(err))
|
||||
}
|
||||
return
|
||||
}
|
||||
if incoming.Type != EventMessage || incoming.Data == "" {
|
||||
c.send <- OutgoingMessage{Type: EventError, Data: "仅支持非空 message 事件", Timestamp: time.Now()}
|
||||
continue
|
||||
}
|
||||
c.hub.Broadcast(c.room, OutgoingMessage{
|
||||
Type: EventMessage, Room: c.room, ClientID: c.clientID,
|
||||
Data: incoming.Data, Timestamp: time.Now(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) WritePump() {
|
||||
ticker := time.NewTicker(pingPeriod)
|
||||
defer func() {
|
||||
ticker.Stop()
|
||||
_ = c.conn.Close()
|
||||
}()
|
||||
|
||||
for {
|
||||
select {
|
||||
case message, ok := <-c.send:
|
||||
_ = c.conn.SetWriteDeadline(time.Now().Add(writeWait))
|
||||
if !ok {
|
||||
_ = c.conn.WriteMessage(websocket.CloseMessage, []byte{})
|
||||
return
|
||||
}
|
||||
if err := c.conn.WriteJSON(message); err != nil {
|
||||
return
|
||||
}
|
||||
case <-ticker.C:
|
||||
_ = c.conn.SetWriteDeadline(time.Now().Add(writeWait))
|
||||
if err := c.conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func presenceMessage(room, clientID, action string, online int) OutgoingMessage {
|
||||
return OutgoingMessage{
|
||||
Type: EventPresence, Room: room, ClientID: clientID,
|
||||
Data: action, Online: online, Timestamp: time.Now(),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
type Controller struct {
|
||||
hub *Hub
|
||||
upgrader websocket.Upgrader
|
||||
}
|
||||
|
||||
func NewController(hub *Hub) *Controller {
|
||||
return &Controller{
|
||||
hub: hub,
|
||||
upgrader: websocket.Upgrader{
|
||||
ReadBufferSize: 1024,
|
||||
WriteBufferSize: 1024,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Connect 升级 HTTP 连接并加入指定房间。
|
||||
// 浏览器连接示例:new WebSocket("ws://localhost:8080/api/ws?room=lobby&client_id=demo")
|
||||
func (c *Controller) Connect(ctx *gin.Context) {
|
||||
room := strings.TrimSpace(ctx.DefaultQuery("room", "lobby"))
|
||||
clientID := strings.TrimSpace(ctx.Query("client_id"))
|
||||
if room == "" || len(room) > 64 || len(clientID) > 64 {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"code": 400, "message": "room 或 client_id 不合法", "data": nil})
|
||||
return
|
||||
}
|
||||
if clientID == "" {
|
||||
clientID = randomClientID()
|
||||
}
|
||||
|
||||
conn, err := c.upgrader.Upgrade(ctx.Writer, ctx.Request, nil)
|
||||
if err != nil {
|
||||
zap.L().Warn("WebSocket升级失败", zap.Error(err))
|
||||
return
|
||||
}
|
||||
client := NewClient(c.hub, conn, room, clientID)
|
||||
online := c.hub.Register(client)
|
||||
client.send <- OutgoingMessage{
|
||||
Type: EventWelcome, Room: room, ClientID: clientID,
|
||||
Data: "connected", Online: online, Timestamp: time.Now(),
|
||||
}
|
||||
c.hub.Broadcast(room, presenceMessage(room, clientID, "joined", online))
|
||||
|
||||
go client.WritePump()
|
||||
client.ReadPump()
|
||||
}
|
||||
|
||||
func randomClientID() string {
|
||||
value := make([]byte, 8)
|
||||
if _, err := rand.Read(value); err != nil {
|
||||
return "anonymous"
|
||||
}
|
||||
return hex.EncodeToString(value)
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
func TestControllerConnectAndBroadcast(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
hub := NewHub()
|
||||
controller := NewController(hub)
|
||||
router := gin.New()
|
||||
router.GET("/ws", controller.Connect)
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
|
||||
url := "ws" + strings.TrimPrefix(server.URL, "http") + "/ws?room=test&client_id=tester"
|
||||
conn, _, err := websocket.DefaultDialer.Dial(url, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Dial() error = %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
_ = conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
|
||||
// 连接后依次收到 welcome 和 joined presence。
|
||||
for i := 0; i < 2; i++ {
|
||||
var message OutgoingMessage
|
||||
if err := conn.ReadJSON(&message); err != nil {
|
||||
t.Fatalf("initial ReadJSON() error = %v", err)
|
||||
}
|
||||
}
|
||||
if err := conn.WriteJSON(IncomingMessage{Type: EventMessage, Data: "hello"}); err != nil {
|
||||
t.Fatalf("WriteJSON() error = %v", err)
|
||||
}
|
||||
var message OutgoingMessage
|
||||
if err := conn.ReadJSON(&message); err != nil {
|
||||
t.Fatalf("ReadJSON() error = %v", err)
|
||||
}
|
||||
if message.Type != EventMessage || message.Data != "hello" || message.Room != "test" {
|
||||
t.Fatalf("unexpected broadcast: %+v", message)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package ws
|
||||
|
||||
import "sync"
|
||||
|
||||
// Hub 管理房间及连接。它只负责单进程广播;跨实例广播可在此接入 Redis Pub/Sub。
|
||||
type Hub struct {
|
||||
mu sync.RWMutex
|
||||
rooms map[string]map[*Client]struct{}
|
||||
}
|
||||
|
||||
func NewHub() *Hub {
|
||||
return &Hub{rooms: make(map[string]map[*Client]struct{})}
|
||||
}
|
||||
|
||||
func (h *Hub) Register(client *Client) int {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
if h.rooms[client.room] == nil {
|
||||
h.rooms[client.room] = make(map[*Client]struct{})
|
||||
}
|
||||
h.rooms[client.room][client] = struct{}{}
|
||||
return len(h.rooms[client.room])
|
||||
}
|
||||
|
||||
func (h *Hub) Unregister(client *Client) int {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
clients, ok := h.rooms[client.room]
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
if _, ok := clients[client]; ok {
|
||||
delete(clients, client)
|
||||
close(client.send)
|
||||
}
|
||||
remaining := len(clients)
|
||||
if remaining == 0 {
|
||||
delete(h.rooms, client.room)
|
||||
}
|
||||
return remaining
|
||||
}
|
||||
|
||||
func (h *Hub) Broadcast(room string, message OutgoingMessage) {
|
||||
h.mu.RLock()
|
||||
clients := h.rooms[room]
|
||||
stale := make([]*Client, 0)
|
||||
for client := range clients {
|
||||
select {
|
||||
case client.send <- message:
|
||||
default:
|
||||
stale = append(stale, client)
|
||||
}
|
||||
}
|
||||
h.mu.RUnlock()
|
||||
|
||||
for _, client := range stale {
|
||||
h.Unregister(client)
|
||||
_ = client.conn.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Hub) Count(room string) int {
|
||||
h.mu.RLock()
|
||||
defer h.mu.RUnlock()
|
||||
return len(h.rooms[room])
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package ws
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestHubRegisterBroadcastAndUnregister(t *testing.T) {
|
||||
hub := NewHub()
|
||||
client := &Client{hub: hub, room: "test", clientID: "one", send: make(chan OutgoingMessage, 1)}
|
||||
if online := hub.Register(client); online != 1 {
|
||||
t.Fatalf("Register() online = %d", online)
|
||||
}
|
||||
hub.Broadcast("test", OutgoingMessage{Type: EventMessage, Data: "hello"})
|
||||
message := <-client.send
|
||||
if message.Data != "hello" {
|
||||
t.Fatalf("Broadcast() data = %q", message.Data)
|
||||
}
|
||||
if online := hub.Unregister(client); online != 0 || hub.Count("test") != 0 {
|
||||
t.Fatalf("Unregister() online = %d", online)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
package ws
|
||||
|
||||
import "time"
|
||||
|
||||
const (
|
||||
EventWelcome = "welcome"
|
||||
EventMessage = "message"
|
||||
EventPresence = "presence"
|
||||
EventError = "error"
|
||||
)
|
||||
|
||||
type IncomingMessage struct {
|
||||
Type string `json:"type" binding:"required" example:"message"`
|
||||
Data string `json:"data" binding:"required" example:"hello"`
|
||||
}
|
||||
|
||||
type OutgoingMessage struct {
|
||||
Type string `json:"type"`
|
||||
Room string `json:"room,omitempty"`
|
||||
ClientID string `json:"client_id,omitempty"`
|
||||
Data string `json:"data,omitempty"`
|
||||
Online int `json:"online,omitempty"`
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
}
|
||||
Reference in New Issue
Block a user