Files
lmvpn_server/internal/vpn/tunnel.go
T
kevin 8398067acb feat: 连接列表显示客户端真实 IP,支持反代和 CDN 场景
新增可配置的 real_ip_headers(默认 CF-Connecting-IP > X-Real-IP >
X-Forwarded-For 降级取值),trusted_proxies 用于 gin 代理信任链。
WebSocket 连接建立时捕获真实 IP 并在 /profile 和 /admin 连接列表展示,
登录会话 IP 同步改用真实 IP。
2026-07-10 16:58:14 +08:00

283 lines
6.5 KiB
Go

package vpn
import (
"encoding/json"
"log"
"net"
"sync"
"sync/atomic"
"time"
"lmvpn/internal/db"
"lmvpn/internal/model"
"github.com/gorilla/websocket"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
readTimeout = 60 * time.Second
writeTimeout = 10 * time.Second
readyTimeout = 30 * time.Second
pingPeriod = 30 * time.Second
maxMessageSize = 1 << 20
)
var (
activeConns = make(map[uint]int)
activeConnsMu sync.Mutex
)
type tunnelConn struct {
conn *websocket.Conn
user *model.User
svc *VpnService
assignedIP net.IP
assignedIP6 net.IP
realIP string
connectedAt time.Time
writeMu sync.Mutex
ready atomic.Bool
rxBytes atomic.Int64
txBytes atomic.Int64
flushedRx atomic.Int64
flushedTx atomic.Int64
}
func (c *tunnelConn) AssignedIP() net.IP { return c.assignedIP }
func (c *tunnelConn) AssignedIP6() net.IP { return c.assignedIP6 }
func (c *tunnelConn) flushDelta() (rx, tx int64) {
curRx := c.rxBytes.Load()
curTx := c.txBytes.Load()
rx = curRx - c.flushedRx.Swap(curRx)
tx = curTx - c.flushedTx.Swap(curTx)
return
}
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 {
ci := ClientInfo{
UserID: c.user.ID,
Username: c.user.Username,
IP: c.assignedIP.String(),
RealIP: c.realIP,
ConnectedAt: c.connectedAt.Format("2006-01-02 15:04:05"),
RxBytes: c.rxBytes.Load(),
TxBytes: c.txBytes.Load(),
}
if c.assignedIP6 != nil {
ci.IP6 = c.assignedIP6.String()
}
return ci
}
func runTunnel(conn *websocket.Conn, user *model.User, realIP string) {
defer conn.Close()
if VPN == nil || !VPN.Running() {
_ = sendJSON(conn, controlMessage{Type: "error", Message: "VPN 服务未启用"})
return
}
activeConnsMu.Lock()
maxConns := VPN.Settings().MaxConnsPerUser
if maxConns <= 0 {
maxConns = 30
}
if activeConns[user.ID] >= maxConns {
activeConnsMu.Unlock()
_ = sendJSON(conn, controlMessage{Type: "error", Message: "连接数已达上限"})
return
}
activeConns[user.ID]++
activeConnsMu.Unlock()
defer func() {
activeConnsMu.Lock()
activeConns[user.ID]--
if activeConns[user.ID] <= 0 {
delete(activeConns, user.ID)
}
activeConnsMu.Unlock()
}()
ip4, ip6, 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: ip4,
assignedIP6: ip6,
realIP: realIP,
connectedAt: time.Now(),
}
VPN.registerClient(tc)
defer func() {
rx, tx := tc.flushDelta()
recordTraffic(tc.user.ID, rx, tx)
VPN.unregisterClient(tc)
}()
settings := VPN.Settings()
initMsg := initMessage{
Type: "init",
IP: ip4.String(),
Prefix: VPN.Prefix(),
MTU: settings.MTU,
ServerIP: VPN.ServerIP().String(),
}
if ip6 != nil {
initMsg.IP6 = ip6.String()
initMsg.Prefix6 = VPN.Prefix6()
initMsg.ServerIP6 = VPN.ServerIP6().String()
}
if err := tc.writeControl(initMsg); err != nil {
log.Printf("用户 %s 发送 init 失败: %v", user.Username, err)
return
}
log.Printf("用户 %s 已连接,分配 IP %s", user.Username, ip4.String())
if ip6 != nil {
log.Printf(" IPv6: %s", ip6.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 {
tc.writeMu.Lock()
if err := conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(writeTimeout)); err != nil {
tc.writeMu.Unlock()
return
}
tc.writeMu.Unlock()
}
}()
conn.SetPongHandler(func(string) error {
conn.SetReadDeadline(time.Now().Add(readTimeout))
return nil
})
for {
messageType, data, err := conn.ReadMessage()
if 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, ip4.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 isDrop(targets) {
continue
}
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)
}
}
}
func recordTraffic(userID uint, rx, tx int64) {
if rx == 0 && tx == 0 {
return
}
today := time.Now().Format("2006-01-02")
userStat := model.UserTrafficStat{UserID: userID, Date: today, RxBytes: rx, TxBytes: tx}
if err := db.DB.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "user_id"}, {Name: "date"}},
DoUpdates: clause.Assignments(map[string]interface{}{
"rx_bytes": gorm.Expr("rx_bytes + ?", rx),
"tx_bytes": gorm.Expr("tx_bytes + ?", tx),
}),
}).Create(&userStat).Error; err != nil {
log.Printf("记录用户流量失败: %v", err)
}
stat := model.TrafficStat{Date: today, RxBytes: rx, TxBytes: tx}
if err := db.DB.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "date"}},
DoUpdates: clause.Assignments(map[string]interface{}{
"rx_bytes": gorm.Expr("rx_bytes + ?", rx),
"tx_bytes": gorm.Expr("tx_bytes + ?", tx),
}),
}).Create(&stat).Error; err != nil {
log.Printf("记录流量失败: %v", err)
}
}