添加 WebSocket VPN 接口及用户认证
- 移动 config 包到 internal/config,新增 DatabaseConfig - 使用 GORM + SQLite 管理用户数据,预留 MySQL 支持 - 新增 internal/model/user.go GORM 用户模型 - 新增 internal/db/db.go 数据库初始化及默认管理员 - 新增 internal/vpn/ WebSocket 鉴权与隧道骨架 - /ws 路径支持用户名+密码鉴权 (bcrypt) - 心跳保活 + echo 隧道模式
This commit is contained in:
@@ -0,0 +1,76 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/goccy/go-yaml"
|
||||
)
|
||||
|
||||
type WebConfig struct {
|
||||
Port int `yaml:"port"`
|
||||
Sock string `yaml:"sock"`
|
||||
}
|
||||
|
||||
type DatabaseConfig struct {
|
||||
Type string `yaml:"type"`
|
||||
Path string `yaml:"path"`
|
||||
DSN string `yaml:"dsn"`
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
Web WebConfig `yaml:"web"`
|
||||
Database DatabaseConfig `yaml:"database"`
|
||||
}
|
||||
|
||||
func defaultConfig() *Config {
|
||||
return &Config{
|
||||
Web: WebConfig{
|
||||
Port: 8080,
|
||||
Sock: "/run/lmvpnweb.sock",
|
||||
},
|
||||
Database: DatabaseConfig{
|
||||
Type: "sqlite",
|
||||
Path: "data/lmvpn.db",
|
||||
DSN: "",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func Load(path string) (*Config, error) {
|
||||
dir := filepath.Dir(path)
|
||||
if err := os.MkdirAll(dir, 0755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cfg := defaultConfig()
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
if !os.IsNotExist(err) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := saveConfig(path, cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
if err := yaml.Unmarshal(data, cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := saveConfig(path, cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func saveConfig(path string, cfg *Config) error {
|
||||
data, err := yaml.Marshal(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(path, data, 0644)
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
"lmvpn/internal/config"
|
||||
"lmvpn/internal/model"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"gorm.io/driver/mysql"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var DB *gorm.DB
|
||||
|
||||
func Init(cfg *config.DatabaseConfig) error {
|
||||
var d gorm.Dialector
|
||||
|
||||
switch cfg.Type {
|
||||
case "sqlite":
|
||||
d = sqlite.Open(cfg.Path)
|
||||
case "mysql":
|
||||
if cfg.DSN == "" {
|
||||
return fmt.Errorf("mysql DSN 不能为空")
|
||||
}
|
||||
d = mysql.Open(cfg.DSN)
|
||||
default:
|
||||
return fmt.Errorf("不支持的数据库类型: %s", cfg.Type)
|
||||
}
|
||||
|
||||
var err error
|
||||
DB, err = gorm.Open(d, &gorm.Config{})
|
||||
if err != nil {
|
||||
return fmt.Errorf("数据库连接失败: %w", err)
|
||||
}
|
||||
|
||||
if err := DB.AutoMigrate(&model.User{}); err != nil {
|
||||
return fmt.Errorf("数据库迁移失败: %w", err)
|
||||
}
|
||||
|
||||
if err := seedDefaultAdmin(); err != nil {
|
||||
return fmt.Errorf("创建默认管理员失败: %w", err)
|
||||
}
|
||||
|
||||
log.Printf("数据库初始化完成: %s", cfg.Type)
|
||||
return nil
|
||||
}
|
||||
|
||||
func seedDefaultAdmin() error {
|
||||
var count int64
|
||||
DB.Model(&model.User{}).Count(&count)
|
||||
if count > 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte("admin123"), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
admin := &model.User{
|
||||
Username: "admin",
|
||||
Password: string(hash),
|
||||
Role: "admin",
|
||||
Status: 1,
|
||||
}
|
||||
|
||||
if err := DB.Create(admin).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
log.Println("已创建默认管理员: admin / admin123")
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type User struct {
|
||||
ID uint `gorm:"primaryKey;autoIncrement"`
|
||||
Username string `gorm:"uniqueIndex;size:64;not null"`
|
||||
Password string `gorm:"size:128;not null"`
|
||||
Role string `gorm:"size:16;default:user"`
|
||||
Status int `gorm:"default:1"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (User) TableName() string {
|
||||
return "users"
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package vpn
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"lmvpn/internal/model"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type authMessage struct {
|
||||
Type string `json:"type"`
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
type authResponse struct {
|
||||
Type string `json:"type"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
func authenticate(conn *websocket.Conn, db *gorm.DB) (*model.User, error) {
|
||||
_, msgBytes, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var msg authMessage
|
||||
if err := json.Unmarshal(msgBytes, &msg); err != nil || msg.Type != "auth" {
|
||||
resp := authResponse{Type: "auth_err", Message: "消息格式错误"}
|
||||
sendJSON(conn, resp)
|
||||
conn.Close()
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
var user model.User
|
||||
if err := db.Where("username = ? AND status = 1", msg.Username).First(&user).Error; err != nil {
|
||||
resp := authResponse{Type: "auth_err", Message: "用户名或密码错误"}
|
||||
sendJSON(conn, resp)
|
||||
conn.Close()
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(msg.Password)); err != nil {
|
||||
resp := authResponse{Type: "auth_err", Message: "用户名或密码错误"}
|
||||
sendJSON(conn, resp)
|
||||
conn.Close()
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
resp := authResponse{Type: "auth_ok"}
|
||||
if err := sendJSON(conn, resp); err != nil {
|
||||
conn.Close()
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func sendJSON(conn *websocket.Conn, v interface{}) error {
|
||||
data, _ := json.Marshal(v)
|
||||
return conn.WriteMessage(websocket.TextMessage, data)
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package vpn
|
||||
|
||||
import (
|
||||
"log"
|
||||
"net/http"
|
||||
|
||||
"lmvpn/internal/db"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
var upgrader = websocket.Upgrader{
|
||||
ReadBufferSize: 4096,
|
||||
WriteBufferSize: 4096,
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true
|
||||
},
|
||||
}
|
||||
|
||||
func HandleWS(c *gin.Context) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
log.Printf("WebSocket 升级失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
user, err := authenticate(conn, db.DB)
|
||||
if err != nil {
|
||||
log.Printf("认证读取失败: %v", err)
|
||||
conn.Close()
|
||||
return
|
||||
}
|
||||
if user == nil {
|
||||
return
|
||||
}
|
||||
|
||||
runTunnel(conn, user)
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
package vpn
|
||||
|
||||
import (
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"lmvpn/internal/model"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
const (
|
||||
readTimeout = 60 * time.Second
|
||||
writeTimeout = 10 * time.Second
|
||||
pingPeriod = 30 * time.Second
|
||||
)
|
||||
|
||||
func runTunnel(conn *websocket.Conn, user *model.User) {
|
||||
defer conn.Close()
|
||||
|
||||
log.Printf("用户 %s 已连接", user.Username)
|
||||
|
||||
go func() {
|
||||
ticker := time.NewTicker(pingPeriod)
|
||||
defer ticker.Stop()
|
||||
for range ticker.C {
|
||||
conn.SetWriteDeadline(time.Now().Add(writeTimeout))
|
||||
if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
conn.SetPongHandler(func(string) error {
|
||||
conn.SetReadDeadline(time.Now().Add(readTimeout))
|
||||
return nil
|
||||
})
|
||||
|
||||
for {
|
||||
conn.SetReadDeadline(time.Now().Add(readTimeout))
|
||||
messageType, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
log.Printf("用户 %s 断开连接: %v", user.Username, err)
|
||||
return
|
||||
}
|
||||
|
||||
conn.SetWriteDeadline(time.Now().Add(writeTimeout))
|
||||
if err := conn.WriteMessage(messageType, data); err != nil {
|
||||
log.Printf("用户 %s 发送失败: %v", user.Username, err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user