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