feat: 实现第三层 TUN VPN 服务,支持后台 IP 分配

- 新增 TUN 设备层(water 库),分 linux/darwin 平台配置 IP/路由/MTU
- 实现 IP 分配管理:动态池自动分配 + 按用户静态预留,支持热更新
- 实现 PacketSwitch 共享 TUN 包转发:源 IP 防伪、按目的 IP 查表转发、allow-c2c
- 重写隧道:自研简化 WS 协议(文本帧 JSON 控制 init/ready,二进制帧=原始 IP 包)
- VpnService 单例管理 TUN 生命周期,子网变更踢线重建,预留增删热更新
- 新增 vpn_settings/vpn_reservations 表,AutoMigrate + 默认设置 seed
- 新增 Admin API:settings 读写、status、clients、reservations CRUD
- 前端新增 VpnView(/admin/vpn):状态面板/设置表单/在线客户端/静态预留
- main.go 启动时按 DB 设置初始化 VPN 服务
This commit is contained in:
2026-07-03 14:49:56 +08:00
parent 61189a53ec
commit 44b51b3b04
19 changed files with 1571 additions and 25 deletions
+22 -1
View File
@@ -40,10 +40,14 @@ func Init(cfg *config.DatabaseConfig) error {
return fmt.Errorf("数据库连接失败: %w", err)
}
if err := DB.AutoMigrate(&model.User{}, &model.Session{}); err != nil {
if err := DB.AutoMigrate(&model.User{}, &model.Session{}, &model.VpnSetting{}, &model.VpnReservation{}); err != nil {
return fmt.Errorf("数据库迁移失败: %w", err)
}
if err := seedDefaultVpnSettings(); err != nil {
return fmt.Errorf("初始化 VPN 设置失败: %w", err)
}
if err := seedDefaultAdmin(cfg); err != nil {
return fmt.Errorf("创建默认管理员失败: %w", err)
}
@@ -98,6 +102,23 @@ func seedDefaultAdmin(cfg *config.DatabaseConfig) error {
return nil
}
func seedDefaultVpnSettings() error {
var s model.VpnSetting
if err := DB.First(&s, model.VpnSettingSingletonID).Error; err == nil {
return nil
}
s = model.VpnSetting{
ID: model.VpnSettingSingletonID,
Enabled: false,
Subnet: "192.168.3.0/24",
MTU: 1420,
InterfaceName: "",
DoLocalIPConfig: true,
DoRemoteIPConfig: true,
}
return DB.Create(&s).Error
}
func generateRandomPassword(length int) (string, error) {
const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
b := make([]byte, length)
+312
View File
@@ -0,0 +1,312 @@
package handler
import (
"net"
"net/http"
"strconv"
"lmvpn/internal/db"
"lmvpn/internal/model"
"lmvpn/internal/vpn"
"github.com/apparentlymart/go-cidr/cidr"
"github.com/gin-gonic/gin"
)
type vpnSettingsResponse struct {
Enabled bool `json:"enabled"`
Subnet string `json:"subnet"`
MTU int `json:"mtu"`
InterfaceName string `json:"interface_name"`
AllowClientToClient bool `json:"allow_client_to_client"`
DoLocalIPConfig bool `json:"do_local_ip_config"`
DoRemoteIPConfig bool `json:"do_remote_ip_config"`
}
type updateVpnSettingsRequest struct {
Enabled *bool `json:"enabled"`
Subnet *string `json:"subnet"`
MTU *int `json:"mtu"`
InterfaceName *string `json:"interface_name"`
AllowClientToClient *bool `json:"allow_client_to_client"`
DoLocalIPConfig *bool `json:"do_local_ip_config"`
DoRemoteIPConfig *bool `json:"do_remote_ip_config"`
}
func loadVpnSettings() (model.VpnSetting, error) {
var s model.VpnSetting
err := db.DB.First(&s, model.VpnSettingSingletonID).Error
return s, err
}
func loadReservationsMap() (map[uint]string, error) {
var rows []model.VpnReservation
if err := db.DB.Find(&rows).Error; err != nil {
return nil, err
}
out := make(map[uint]string, len(rows))
for _, r := range rows {
out[r.UserID] = r.IPAddress
}
return out, nil
}
func ApplyVpnFromDB(svc *vpn.VpnService) error {
s, err := loadVpnSettings()
if err != nil {
return err
}
reservations, err := loadReservationsMap()
if err != nil {
return err
}
return svc.ApplySettings(s, reservations)
}
func validateSubnet(subnet string) error {
ip, ipNet, err := net.ParseCIDR(subnet)
if err != nil {
return err
}
if ip.To4() == nil {
return errIPv4Only
}
ones, _ := ipNet.Mask.Size()
if ones > 30 {
return errSubnetTooSmall
}
return nil
}
var (
errIPv4Only = errStr("仅支持 IPv4 子网")
errSubnetTooSmall = errStr("子网前缀长度不能大于 /30")
errIPNotInSubnet = errStr("IP 不在子网范围内")
errIPReserved = errStr("该 IP 已被预留")
errIPIsServer = errStr("该 IP 为服务器 IP,不可预留")
)
type errStr string
func (e errStr) Error() string { return string(e) }
func GetVpnSettings(c *gin.Context) {
s, err := loadVpnSettings()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "加载设置失败"})
return
}
c.JSON(http.StatusOK, vpnSettingsResponse{
Enabled: s.Enabled,
Subnet: s.Subnet,
MTU: s.MTU,
InterfaceName: s.InterfaceName,
AllowClientToClient: s.AllowClientToClient,
DoLocalIPConfig: s.DoLocalIPConfig,
DoRemoteIPConfig: s.DoRemoteIPConfig,
})
}
func UpdateVpnSettings(c *gin.Context) {
var req updateVpnSettingsRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "参数错误"})
return
}
s, err := loadVpnSettings()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "加载设置失败"})
return
}
if req.Subnet != nil && *req.Subnet != s.Subnet {
if err := validateSubnet(*req.Subnet); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
s.Subnet = *req.Subnet
}
if req.MTU != nil {
if *req.MTU < 500 || *req.MTU > 65535 {
c.JSON(http.StatusBadRequest, gin.H{"error": "MTU 范围 500-65535"})
return
}
s.MTU = *req.MTU
}
if req.InterfaceName != nil {
s.InterfaceName = *req.InterfaceName
}
if req.Enabled != nil {
s.Enabled = *req.Enabled
}
if req.AllowClientToClient != nil {
s.AllowClientToClient = *req.AllowClientToClient
}
if req.DoLocalIPConfig != nil {
s.DoLocalIPConfig = *req.DoLocalIPConfig
}
if req.DoRemoteIPConfig != nil {
s.DoRemoteIPConfig = *req.DoRemoteIPConfig
}
if err := db.DB.Save(&s).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "保存设置失败"})
return
}
if err := ApplyVpnFromDB(vpn.VPN); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "应用设置失败: " + err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"message": "设置已更新"})
}
type vpnStatusResponse struct {
Enabled bool `json:"enabled"`
Online int `json:"online"`
UsedIPs int `json:"used_ips"`
Capacity uint64 `json:"capacity"`
Clients []vpn.ClientInfo `json:"clients"`
}
func GetVpnStatus(c *gin.Context) {
s, err := loadVpnSettings()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "加载设置失败"})
return
}
used, cap := vpn.VPN.AllocStats()
clients := vpn.VPN.ClientList()
c.JSON(http.StatusOK, vpnStatusResponse{
Enabled: s.Enabled,
Online: len(clients),
UsedIPs: used,
Capacity: cap,
Clients: clients,
})
}
type reservationResponse struct {
ID uint `json:"id"`
UserID uint `json:"user_id"`
Username string `json:"username"`
IPAddress string `json:"ip_address"`
CreatedAt string `json:"created_at"`
}
func ListVpnReservations(c *gin.Context) {
var rows []model.VpnReservation
if err := db.DB.Order("id asc").Find(&rows).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "加载预留失败"})
return
}
userIDs := make([]uint, 0, len(rows))
for _, r := range rows {
userIDs = append(userIDs, r.UserID)
}
nameMap := make(map[uint]string)
if len(userIDs) > 0 {
var users []model.User
db.DB.Where("id IN ?", userIDs).Find(&users)
for _, u := range users {
nameMap[u.ID] = u.Username
}
}
out := make([]reservationResponse, len(rows))
for i, r := range rows {
out[i] = reservationResponse{
ID: r.ID,
UserID: r.UserID,
Username: nameMap[r.UserID],
IPAddress: r.IPAddress,
CreatedAt: r.CreatedAt.Format("2006-01-02 15:04:05"),
}
}
c.JSON(http.StatusOK, gin.H{"reservations": out})
}
type createReservationRequest struct {
UserID uint `json:"user_id" binding:"required"`
IPAddress string `json:"ip_address" binding:"required"`
}
func CreateVpnReservation(c *gin.Context) {
var req createReservationRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "参数错误"})
return
}
var user model.User
if err := db.DB.First(&user, req.UserID).Error; err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "用户不存在"})
return
}
s, err := loadVpnSettings()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "加载设置失败"})
return
}
_, ipNet, err := net.ParseCIDR(s.Subnet)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "子网配置错误"})
return
}
ip := net.ParseIP(req.IPAddress)
if ip == nil || !ipNet.Contains(ip) {
c.JSON(http.StatusBadRequest, gin.H{"error": errIPNotInSubnet.Error()})
return
}
serverIP, _ := cidr.Host(ipNet, 1)
if ip.Equal(serverIP) {
c.JSON(http.StatusBadRequest, gin.H{"error": errIPIsServer.Error()})
return
}
var count int64
db.DB.Model(&model.VpnReservation{}).Where("ip_address = ?", req.IPAddress).Count(&count)
if count > 0 {
c.JSON(http.StatusBadRequest, gin.H{"error": errIPReserved.Error()})
return
}
var existUser model.VpnReservation
if err := db.DB.Where("user_id = ?", req.UserID).First(&existUser).Error; err == nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "该用户已有预留 IP"})
return
}
r := model.VpnReservation{UserID: req.UserID, IPAddress: req.IPAddress}
if err := db.DB.Create(&r).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "创建预留失败"})
return
}
if vpn.VPN.Running() {
vpn.VPN.AddReservation(req.UserID, req.IPAddress)
}
c.JSON(http.StatusOK, gin.H{"message": "预留已创建"})
}
func DeleteVpnReservation(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "参数错误"})
return
}
var r model.VpnReservation
if err := db.DB.First(&r, id).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "预留不存在"})
return
}
if err := db.DB.Delete(&r).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "删除失败"})
return
}
if vpn.VPN.Running() {
vpn.VPN.RemoveReservation(r.UserID)
}
c.JSON(http.StatusOK, gin.H{"message": "删除成功"})
}
+32
View File
@@ -0,0 +1,32 @@
package model
import "time"
const VpnSettingSingletonID = 1
type VpnSetting struct {
ID uint `gorm:"primaryKey"`
Enabled bool `gorm:"default:false"`
Subnet string `gorm:"size:64;not null"`
MTU int `gorm:"default:1420"`
InterfaceName string `gorm:"size:16"`
AllowClientToClient bool `gorm:"default:false"`
DoLocalIPConfig bool `gorm:"default:true"`
DoRemoteIPConfig bool `gorm:"default:true"`
UpdatedAt time.Time
}
func (VpnSetting) TableName() string {
return "vpn_settings"
}
type VpnReservation struct {
ID uint `gorm:"primaryKey;autoIncrement"`
UserID uint `gorm:"uniqueIndex;not null"`
IPAddress string `gorm:"size:64;uniqueIndex;not null"`
CreatedAt time.Time `gorm:"autoCreateTime"`
}
func (VpnReservation) TableName() string {
return "vpn_reservations"
}
+7
View File
@@ -34,6 +34,13 @@ func Setup(r *gin.Engine) {
admin.PUT("/users/:id", handler.UpdateUser)
admin.DELETE("/users/:id", handler.DeleteUser)
admin.DELETE("/users/:id/sessions", handler.AdminRevokeUserSessions)
admin.GET("/vpn/settings", handler.GetVpnSettings)
admin.PUT("/vpn/settings", handler.UpdateVpnSettings)
admin.GET("/vpn/status", handler.GetVpnStatus)
admin.GET("/vpn/reservations", handler.ListVpnReservations)
admin.POST("/vpn/reservations", handler.CreateVpnReservation)
admin.DELETE("/vpn/reservations/:id", handler.DeleteVpnReservation)
}
distDir := http.Dir("./dist")
+131
View File
@@ -0,0 +1,131 @@
package vpn
import (
"errors"
"fmt"
"net"
"sync"
"github.com/apparentlymart/go-cidr/cidr"
)
type AllocationManager struct {
mu sync.Mutex
net *net.IPNet
serverIP net.IP
used map[string]bool
reservedByUser map[uint]string
reservedSet map[string]bool
}
func NewAllocationManager(ipNet *net.IPNet, serverIP net.IP, reservations map[uint]string) *AllocationManager {
m := &AllocationManager{
net: ipNet,
serverIP: serverIP,
used: make(map[string]bool),
reservedByUser: make(map[uint]string),
reservedSet: make(map[string]bool),
}
for uid, ip := range reservations {
m.reservedByUser[uid] = ip
m.reservedSet[ip] = true
}
return m
}
func (m *AllocationManager) ServerIP() net.IP { return m.serverIP }
func (m *AllocationManager) Subnet() *net.IPNet { return m.net }
func (m *AllocationManager) Allocate(userID uint) (net.IP, error) {
m.mu.Lock()
defer m.mu.Unlock()
if ipStr, ok := m.reservedByUser[userID]; ok {
if m.used[ipStr] {
return nil, fmt.Errorf("用户预留 IP %s 已被占用", ipStr)
}
m.used[ipStr] = true
return net.ParseIP(ipStr), nil
}
count := cidr.AddressCount(m.net)
maxIndex := int(count - 1)
for i := 2; i < maxIndex; i++ {
ip, err := cidr.Host(m.net, i)
if err != nil {
continue
}
ipStr := ip.String()
if m.used[ipStr] || m.reservedSet[ipStr] {
continue
}
m.used[ipStr] = true
return ip, nil
}
return nil, errors.New("可用 IP 地址已耗尽")
}
func (m *AllocationManager) Release(ip net.IP) {
if ip == nil {
return
}
m.mu.Lock()
defer m.mu.Unlock()
delete(m.used, ip.String())
}
func (m *AllocationManager) IsReserved(ipStr string) bool {
m.mu.Lock()
defer m.mu.Unlock()
return m.reservedSet[ipStr]
}
func (m *AllocationManager) ReservedByUser(userID uint) (string, bool) {
m.mu.Lock()
defer m.mu.Unlock()
ip, ok := m.reservedByUser[userID]
return ip, ok
}
func (m *AllocationManager) ReservedList() map[uint]string {
m.mu.Lock()
defer m.mu.Unlock()
out := make(map[uint]string, len(m.reservedByUser))
for k, v := range m.reservedByUser {
out[k] = v
}
return out
}
func (m *AllocationManager) UsedCount() int {
m.mu.Lock()
defer m.mu.Unlock()
return len(m.used)
}
func (m *AllocationManager) Capacity() uint64 {
count := cidr.AddressCount(m.net)
if count < 3 {
return 0
}
return count - 3
}
func (m *AllocationManager) AddReservation(userID uint, ipStr string) {
m.mu.Lock()
defer m.mu.Unlock()
if old, ok := m.reservedByUser[userID]; ok {
delete(m.reservedSet, old)
}
m.reservedByUser[userID] = ipStr
m.reservedSet[ipStr] = true
}
func (m *AllocationManager) RemoveReservation(userID uint) {
m.mu.Lock()
defer m.mu.Unlock()
if old, ok := m.reservedByUser[userID]; ok {
delete(m.reservedByUser, userID)
delete(m.reservedSet, old)
}
}
+14
View File
@@ -0,0 +1,14 @@
package vpn
type initMessage struct {
Type string `json:"type"`
IP string `json:"ip"`
Prefix int `json:"prefix"`
MTU int `json:"mtu"`
ServerIP string `json:"server_ip"`
}
type controlMessage struct {
Type string `json:"type"`
Message string `json:"message,omitempty"`
}
+269
View File
@@ -0,0 +1,269 @@
package vpn
import (
"errors"
"fmt"
"log"
"net"
"sync"
"lmvpn/internal/model"
"github.com/apparentlymart/go-cidr/cidr"
)
type VpnService struct {
mu sync.RWMutex
settings model.VpnSetting
net *net.IPNet
serverIP net.IP
prefix int
alloc *AllocationManager
switchx *PacketSwitch
tun *TUNInterface
tunDone chan struct{}
running bool
clients map[*tunnelConn]struct{}
}
func NewVpnService() *VpnService {
return &VpnService{
clients: make(map[*tunnelConn]struct{}),
}
}
func (s *VpnService) Running() bool {
s.mu.RLock()
defer s.mu.RUnlock()
return s.running
}
func (s *VpnService) Settings() model.VpnSetting {
s.mu.RLock()
defer s.mu.RUnlock()
return s.settings
}
func (s *VpnService) parseNet(subnet string) (*net.IPNet, net.IP, int, error) {
_, ipNet, err := net.ParseCIDR(subnet)
if err != nil {
return nil, nil, 0, fmt.Errorf("子网格式错误: %w", err)
}
ones, _ := ipNet.Mask.Size()
serverIP, err := cidr.Host(ipNet, 1)
if err != nil {
return nil, nil, 0, fmt.Errorf("计算服务器 IP 失败: %w", err)
}
return ipNet, serverIP, ones, nil
}
func (s *VpnService) ApplySettings(settings model.VpnSetting, reservations map[uint]string) error {
s.mu.Lock()
if s.running {
s.mu.Unlock()
_ = s.Stop()
s.mu.Lock()
}
s.settings = settings
s.mu.Unlock()
if !settings.Enabled {
return nil
}
ipNet, serverIP, prefix, err := s.parseNet(settings.Subnet)
if err != nil {
return err
}
tun, err := CreateTUN(settings.InterfaceName)
if err != nil {
return err
}
var peerIP net.IP = serverIP
if settings.DoLocalIPConfig {
if err := tun.Configure(serverIP, prefix, peerIP); err != nil {
_ = tun.Close()
return fmt.Errorf("配置 TUN 失败: %w", err)
}
}
if err := tun.SetMTU(settings.MTU); err != nil {
log.Printf("警告: 设置 MTU 失败: %v", err)
}
s.mu.Lock()
s.net = ipNet
s.serverIP = serverIP
s.prefix = prefix
s.alloc = NewAllocationManager(ipNet, serverIP, reservations)
s.switchx = NewPacketSwitch(settings.AllowClientToClient)
s.tun = tun
s.tunDone = make(chan struct{})
s.running = true
s.mu.Unlock()
go s.serveTUN()
log.Printf("VPN 服务已启动: tun=%s subnet=%s server=%s mtu=%d", tun.Name(), ipNet.String(), serverIP.String(), settings.MTU)
return nil
}
func (s *VpnService) serveTUN() {
s.mu.RLock()
tun := s.tun
switchx := s.switchx
done := s.tunDone
bufSize := s.settings.MTU + 64
s.mu.RUnlock()
packet := make([]byte, bufSize)
for {
n, err := tun.Iface.Read(packet)
if err != nil {
log.Printf("TUN 读取结束: %v", err)
close(done)
return
}
if n < 1 {
continue
}
targets := switchx.RouteFromTUN(packet[:n])
for _, t := range targets {
_ = t.WritePacket(packet[:n])
}
}
}
func (s *VpnService) Stop() error {
s.mu.Lock()
if !s.running {
s.mu.Unlock()
return nil
}
s.running = false
tun := s.tun
done := s.tunDone
clients := s.clients
s.clients = make(map[*tunnelConn]struct{})
s.mu.Unlock()
for c := range clients {
c.close()
}
if tun != nil {
_ = tun.Close()
if done != nil {
<-done
}
}
log.Printf("VPN 服务已停止")
return nil
}
func (s *VpnService) Allocate(user *model.User) (net.IP, error) {
s.mu.RLock()
alloc := s.alloc
s.mu.RUnlock()
if alloc == nil {
return nil, errors.New("VPN 服务未运行")
}
return alloc.Allocate(user.ID)
}
func (s *VpnService) WriteToTUN(packet []byte) error {
s.mu.RLock()
tun := s.tun
s.mu.RUnlock()
if tun == nil {
return errors.New("TUN 未就绪")
}
_, err := tun.Iface.Write(packet)
return err
}
func (s *VpnService) RouteFromClient(src SwitchConn, packet []byte) []SwitchConn {
s.mu.RLock()
switchx := s.switchx
s.mu.RUnlock()
if switchx == nil {
return nil
}
return switchx.RouteFromClient(src, packet)
}
func (s *VpnService) registerClient(c *tunnelConn) {
s.mu.Lock()
s.switchx.Register(c)
s.clients[c] = struct{}{}
s.mu.Unlock()
}
func (s *VpnService) unregisterClient(c *tunnelConn) {
s.mu.Lock()
if s.switchx != nil {
s.switchx.Unregister(c)
}
delete(s.clients, c)
if s.alloc != nil {
s.alloc.Release(c.assignedIP)
}
s.mu.Unlock()
}
func (s *VpnService) ServerIP() net.IP {
s.mu.RLock()
defer s.mu.RUnlock()
return s.serverIP
}
func (s *VpnService) Prefix() int {
s.mu.RLock()
defer s.mu.RUnlock()
return s.prefix
}
func (s *VpnService) AllocStats() (used int, capacity uint64) {
s.mu.RLock()
alloc := s.alloc
s.mu.RUnlock()
if alloc == nil {
return 0, 0
}
return alloc.UsedCount(), alloc.Capacity()
}
func (s *VpnService) AddReservation(userID uint, ipStr string) {
s.mu.RLock()
alloc := s.alloc
s.mu.RUnlock()
if alloc != nil {
alloc.AddReservation(userID, ipStr)
}
}
func (s *VpnService) RemoveReservation(userID uint) {
s.mu.RLock()
alloc := s.alloc
s.mu.RUnlock()
if alloc != nil {
alloc.RemoveReservation(userID)
}
}
func (s *VpnService) ClientList() []ClientInfo {
s.mu.RLock()
out := make([]ClientInfo, 0, len(s.clients))
for c := range s.clients {
out = append(out, c.info())
}
s.mu.RUnlock()
return out
}
type ClientInfo struct {
Username string `json:"username"`
IP string `json:"ip"`
ConnectedAt string `json:"connected_at"`
}
var VPN *VpnService
+141
View File
@@ -0,0 +1,141 @@
package vpn
import (
"net"
"sync"
"github.com/songgao/water/waterutil"
)
type SwitchConn interface {
WritePacket(data []byte) error
AssignedIP() net.IP
}
type ipKey [16]byte
func ipToKey(ip net.IP) ipKey {
var k ipKey
copy(k[:], ip.To16())
return k
}
type PacketSwitch struct {
allowClientToClient bool
mu sync.RWMutex
table map[ipKey]SwitchConn
}
func NewPacketSwitch(allowClientToClient bool) *PacketSwitch {
return &PacketSwitch{
allowClientToClient: allowClientToClient,
table: make(map[ipKey]SwitchConn),
}
}
func (s *PacketSwitch) SetAllowClientToClient(v bool) {
s.mu.Lock()
s.allowClientToClient = v
s.mu.Unlock()
}
func (s *PacketSwitch) Register(c SwitchConn) {
k := ipToKey(c.AssignedIP())
s.mu.Lock()
s.table[k] = c
s.mu.Unlock()
}
func (s *PacketSwitch) Unregister(c SwitchConn) {
k := ipToKey(c.AssignedIP())
s.mu.Lock()
if cur, ok := s.table[k]; ok && cur == c {
delete(s.table, k)
}
s.mu.Unlock()
}
func (s *PacketSwitch) findByIP(ip net.IP) SwitchConn {
s.mu.RLock()
c := s.table[ipToKey(ip)]
s.mu.RUnlock()
return c
}
func (s *PacketSwitch) allExcept(skip SwitchConn) []SwitchConn {
s.mu.RLock()
out := make([]SwitchConn, 0, len(s.table))
for _, c := range s.table {
if c == skip {
continue
}
out = append(out, c)
}
s.mu.RUnlock()
return out
}
func parseIPAddrs(packet []byte) (src, dest net.IP, ok bool) {
if len(packet) < 1 {
return nil, nil, false
}
switch {
case waterutil.IsIPv4(packet):
if len(packet) < 20 {
return nil, nil, false
}
return waterutil.IPv4Source(packet), waterutil.IPv4Destination(packet), true
case waterutil.IsIPv6(packet):
if len(packet) < 40 {
return nil, nil, false
}
src = make(net.IP, 16)
copy(src, packet[8:24])
dest = make(net.IP, 16)
copy(dest, packet[24:40])
return src, dest, true
}
return nil, nil, false
}
func (s *PacketSwitch) allowC2C() bool {
s.mu.RLock()
v := s.allowClientToClient
s.mu.RUnlock()
return v
}
func (s *PacketSwitch) RouteFromClient(src SwitchConn, packet []byte) []SwitchConn {
srcIP, dest, ok := parseIPAddrs(packet)
if !ok {
return nil
}
// anti-spoof: enforce assigned source IP
if srcIP != nil && !srcIP.Equal(src.AssignedIP()) {
return nil
}
if dest.IsGlobalUnicast() {
if c := s.findByIP(dest); c != nil && s.allowC2C() {
return []SwitchConn{c}
}
return nil
}
if s.allowC2C() {
return s.allExcept(src)
}
return nil
}
func (s *PacketSwitch) RouteFromTUN(packet []byte) []SwitchConn {
_, dest, ok := parseIPAddrs(packet)
if !ok {
return nil
}
if dest.IsGlobalUnicast() {
if c := s.findByIP(dest); c != nil {
return []SwitchConn{c}
}
return nil
}
return s.allExcept(nil)
}
+42
View File
@@ -0,0 +1,42 @@
package vpn
import (
"fmt"
"os"
"os/exec"
"strings"
"github.com/songgao/water"
)
type TUNInterface struct {
Iface *water.Interface
}
func CreateTUN(name string) (*TUNInterface, error) {
cfg := water.Config{DeviceType: water.TUN}
cfg.Name = name
ifce, err := water.New(cfg)
if err != nil {
return nil, fmt.Errorf("创建 TUN 设备失败: %w", err)
}
return &TUNInterface{Iface: ifce}, nil
}
func (t *TUNInterface) Name() string {
return t.Iface.Name()
}
func (t *TUNInterface) Close() error {
return t.Iface.Close()
}
func execCmd(name string, arg ...string) error {
cmd := exec.Command(name, arg...)
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
if err := cmd.Run(); err != nil {
return fmt.Errorf("command %s %s: %w", name, strings.Join(arg, " "), err)
}
return nil
}
+39
View File
@@ -0,0 +1,39 @@
//go:build darwin
package vpn
import (
"fmt"
"net"
)
func inetFamily(ip net.IP) string {
if ip.To4() == nil {
return "inet6"
}
return "inet"
}
func (t *TUNInterface) Configure(localIP net.IP, prefix int, peerIP net.IP) error {
if localIP == nil {
return execCmd("ifconfig", t.Name(), "up")
}
localCidr := fmt.Sprintf("%s/%d", localIP.String(), prefix)
inetType := inetFamily(localIP)
var err error
if t.Iface.IsTUN() && inetType == "inet" {
err = execCmd("ifconfig", t.Name(), inetType, localCidr, peerIP.String(), "up")
} else {
err = execCmd("ifconfig", t.Name(), inetType, localCidr, "up")
}
return err
}
func (t *TUNInterface) AddSubnetRoute(subnet *net.IPNet) error {
inetType := inetFamily(subnet.IP)
return execCmd("route", "add", fmt.Sprintf("-%s", inetType), "-net", subnet.String(), "-interface", t.Name())
}
func (t *TUNInterface) SetMTU(mtu int) error {
return execCmd("ifconfig", t.Name(), "mtu", fmt.Sprintf("%d", mtu))
}
+33
View File
@@ -0,0 +1,33 @@
//go:build linux
package vpn
import (
"fmt"
"net"
)
func (t *TUNInterface) Configure(localIP net.IP, prefix int, peerIP net.IP) error {
if err := execCmd("ip", "link", "set", "dev", t.Name(), "up"); err != nil {
return err
}
if localIP == nil {
return nil
}
localCidr := fmt.Sprintf("%s/%d", localIP.String(), prefix)
args := []string{"addr", "add", "dev", t.Name(), localCidr, "peer", peerIP.String()}
if err := execCmd("ip", args...); err != nil {
if err2 := execCmd("ip", "addr", "add", "dev", t.Name(), localCidr); err2 != nil {
return err
}
}
return nil
}
func (t *TUNInterface) AddSubnetRoute(subnet *net.IPNet) error {
return execCmd("ip", "route", "add", subnet.String(), "dev", t.Name())
}
func (t *TUNInterface) SetMTU(mtu int) error {
return execCmd("ip", "link", "set", "dev", t.Name(), "mtu", fmt.Sprintf("%d", mtu))
}
+20
View File
@@ -0,0 +1,20 @@
//go:build !linux && !darwin
package vpn
import (
"errors"
"net"
)
func (t *TUNInterface) Configure(localIP net.IP, prefix int, peerIP net.IP) error {
return errors.New("TUN 配置当前平台不支持")
}
func (t *TUNInterface) AddSubnetRoute(subnet *net.IPNet) error {
return errors.New("TUN 路由当前平台不支持")
}
func (t *TUNInterface) SetMTU(mtu int) error {
return errors.New("TUN MTU 当前平台不支持")
}
+137 -9
View File
@@ -1,8 +1,11 @@
package vpn
import (
"encoding/json"
"log"
"net"
"sync"
"sync/atomic"
"time"
"lmvpn/internal/model"
@@ -13,6 +16,7 @@ import (
const (
readTimeout = 60 * time.Second
writeTimeout = 10 * time.Second
readyTimeout = 30 * time.Second
pingPeriod = 30 * time.Second
maxMessageSize = 1 << 20
maxConnsPerUser = 3
@@ -23,13 +27,71 @@ var (
activeConnsMu sync.Mutex
)
type tunnelConn struct {
conn *websocket.Conn
user *model.User
svc *VpnService
assignedIP net.IP
connectedAt time.Time
writeMu sync.Mutex
ready atomic.Bool
rxBytes atomic.Int64
txBytes atomic.Int64
}
func (c *tunnelConn) AssignedIP() net.IP { return c.assignedIP }
func (c *tunnelConn) WritePacket(data []byte) error {
if !c.ready.Load() || len(data) == 0 {
return nil
}
c.writeMu.Lock()
defer c.writeMu.Unlock()
c.conn.SetWriteDeadline(time.Now().Add(writeTimeout))
if err := c.conn.WriteMessage(websocket.BinaryMessage, data); err != nil {
return err
}
c.txBytes.Add(int64(len(data)))
return nil
}
func (c *tunnelConn) writeControl(v interface{}) error {
data, err := json.Marshal(v)
if err != nil {
return err
}
c.writeMu.Lock()
defer c.writeMu.Unlock()
c.conn.SetWriteDeadline(time.Now().Add(writeTimeout))
return c.conn.WriteMessage(websocket.TextMessage, data)
}
func (c *tunnelConn) close() {
c.writeMu.Lock()
_ = c.conn.Close()
c.writeMu.Unlock()
}
func (c *tunnelConn) info() ClientInfo {
return ClientInfo{
Username: c.user.Username,
IP: c.assignedIP.String(),
ConnectedAt: c.connectedAt.Format("2006-01-02 15:04:05"),
}
}
func runTunnel(conn *websocket.Conn, user *model.User) {
defer conn.Close()
if VPN == nil || !VPN.Running() {
_ = sendJSON(conn, controlMessage{Type: "error", Message: "VPN 服务未启用"})
return
}
activeConnsMu.Lock()
if activeConns[user.ID] >= maxConnsPerUser {
activeConnsMu.Unlock()
sendJSON(conn, authResponse{Type: "auth_err", Message: "连接数已达上限"})
_ = sendJSON(conn, controlMessage{Type: "error", Message: "连接数已达上限"})
return
}
activeConns[user.ID]++
@@ -44,18 +106,52 @@ func runTunnel(conn *websocket.Conn, user *model.User) {
activeConnsMu.Unlock()
}()
log.Printf("用户 %s 已连接", user.Username)
ip, err := VPN.Allocate(user)
if err != nil {
_ = sendJSON(conn, controlMessage{Type: "error", Message: "分配 IP 失败: " + err.Error()})
return
}
tc := &tunnelConn{
conn: conn,
user: user,
svc: VPN,
assignedIP: ip,
connectedAt: time.Now(),
}
VPN.registerClient(tc)
defer VPN.unregisterClient(tc)
settings := VPN.Settings()
initMsg := initMessage{
Type: "init",
IP: ip.String(),
Prefix: VPN.Prefix(),
MTU: settings.MTU,
ServerIP: VPN.ServerIP().String(),
}
if err := tc.writeControl(initMsg); err != nil {
log.Printf("用户 %s 发送 init 失败: %v", user.Username, err)
return
}
log.Printf("用户 %s 已连接,分配 IP %s", user.Username, ip.String())
conn.SetReadLimit(maxMessageSize)
conn.SetReadDeadline(time.Now().Add(readyTimeout))
readyDeadline := time.Now().Add(readyTimeout)
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 {
tc.writeMu.Lock()
if err := conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(writeTimeout)); err != nil {
tc.writeMu.Unlock()
return
}
tc.writeMu.Unlock()
}
}()
@@ -65,17 +161,49 @@ func runTunnel(conn *websocket.Conn, user *model.User) {
})
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
if messageType == websocket.TextMessage {
var msg controlMessage
if err := json.Unmarshal(data, &msg); err != nil {
continue
}
if msg.Type == "ready" && !tc.ready.Load() {
tc.ready.Store(true)
conn.SetReadDeadline(time.Now().Add(readTimeout))
log.Printf("用户 %s 就绪 (IP %s)", user.Username, ip.String())
}
continue
}
if messageType != websocket.BinaryMessage {
continue
}
if !tc.ready.Load() {
if time.Now().After(readyDeadline) {
log.Printf("用户 %s 等待 ready 超时", user.Username)
return
}
conn.SetReadDeadline(readyDeadline)
continue
}
tc.rxBytes.Add(int64(len(data)))
targets := VPN.RouteFromClient(tc, data)
if len(targets) == 0 {
if err := VPN.WriteToTUN(data); err != nil {
log.Printf("用户 %s 写入 TUN 失败: %v", user.Username, err)
}
continue
}
for _, t := range targets {
_ = t.WritePacket(data)
}
}
}