mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-07 10:16:38 +08:00
feat: complete Go backend control-plane parity
Bridge Java-to-Go runtime behavior by enforcing forward ownership checks, wiring node command dispatch/diagnostics, and adding contract coverage so migrated APIs can run with production semantics.
This commit is contained in:
@@ -0,0 +1,685 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"go-backend/internal/ws"
|
||||||
|
)
|
||||||
|
|
||||||
|
var errForwardNotFound = errors.New("forward not found")
|
||||||
|
|
||||||
|
type forwardRecord struct {
|
||||||
|
ID int64
|
||||||
|
UserID int64
|
||||||
|
UserName string
|
||||||
|
Name string
|
||||||
|
TunnelID int64
|
||||||
|
RemoteAddr string
|
||||||
|
Strategy string
|
||||||
|
Status int
|
||||||
|
}
|
||||||
|
|
||||||
|
type tunnelRecord struct {
|
||||||
|
ID int64
|
||||||
|
Type int
|
||||||
|
Status int
|
||||||
|
}
|
||||||
|
|
||||||
|
type forwardPortRecord struct {
|
||||||
|
NodeID int64
|
||||||
|
Port int
|
||||||
|
}
|
||||||
|
|
||||||
|
type nodeRecord struct {
|
||||||
|
ID int64
|
||||||
|
Name string
|
||||||
|
ServerIP string
|
||||||
|
TCPListenAddr string
|
||||||
|
UDPListenAddr string
|
||||||
|
InterfaceName string
|
||||||
|
}
|
||||||
|
|
||||||
|
type chainNodeRecord struct {
|
||||||
|
ChainType int
|
||||||
|
Inx int64
|
||||||
|
NodeID int64
|
||||||
|
NodeName string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) resolveForwardAccess(r *http.Request, forwardID int64) (*forwardRecord, int64, int, error) {
|
||||||
|
userID, roleID, err := userRoleFromRequest(r)
|
||||||
|
if err != nil {
|
||||||
|
return nil, 0, 0, err
|
||||||
|
}
|
||||||
|
forward, err := h.ensureForwardAccessByActor(userID, roleID, forwardID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, userID, roleID, err
|
||||||
|
}
|
||||||
|
return forward, userID, roleID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) ensureForwardAccessByActor(actorUserID int64, actorRole int, forwardID int64) (*forwardRecord, error) {
|
||||||
|
forward, err := h.getForwardRecord(forwardID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if actorRole != 0 && forward.UserID != actorUserID {
|
||||||
|
return nil, errForwardNotFound
|
||||||
|
}
|
||||||
|
return forward, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) ensureTunnelPermission(userID int64, roleID int, tunnelID int64) error {
|
||||||
|
if roleID == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var count int
|
||||||
|
err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? AND status = 1`, userID, tunnelID).Scan(&count)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if count <= 0 {
|
||||||
|
return errors.New("你没有该隧道的权限")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) {
|
||||||
|
row := h.repo.DB().QueryRow(`
|
||||||
|
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
|
||||||
|
FROM forward WHERE id = ? LIMIT 1
|
||||||
|
`, forwardID)
|
||||||
|
var fr forwardRecord
|
||||||
|
err := row.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return nil, errForwardNotFound
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(fr.Strategy) == "" {
|
||||||
|
fr.Strategy = "fifo"
|
||||||
|
}
|
||||||
|
return &fr, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) {
|
||||||
|
row := h.repo.DB().QueryRow(`SELECT id, type, status FROM tunnel WHERE id = ? LIMIT 1`, tunnelID)
|
||||||
|
var tr tunnelRecord
|
||||||
|
err := row.Scan(&tr.ID, &tr.Type, &tr.Status)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return nil, errors.New("隧道不存在")
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &tr, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) listForwardsByTunnel(tunnelID int64) ([]forwardRecord, error) {
|
||||||
|
rows, err := h.repo.DB().Query(`
|
||||||
|
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
|
||||||
|
FROM forward
|
||||||
|
WHERE tunnel_id = ?
|
||||||
|
ORDER BY id ASC
|
||||||
|
`, tunnelID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
result := make([]forwardRecord, 0)
|
||||||
|
for rows.Next() {
|
||||||
|
var fr forwardRecord
|
||||||
|
if err := rows.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(fr.Strategy) == "" {
|
||||||
|
fr.Strategy = "fifo"
|
||||||
|
}
|
||||||
|
result = append(result, fr)
|
||||||
|
}
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error) {
|
||||||
|
rows, err := h.repo.DB().Query(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
result := make([]forwardPortRecord, 0)
|
||||||
|
for rows.Next() {
|
||||||
|
var item forwardPortRecord
|
||||||
|
if err := rows.Scan(&item.NodeID, &item.Port); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
result = append(result, item)
|
||||||
|
}
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
|
||||||
|
row := h.repo.DB().QueryRow(`
|
||||||
|
SELECT id, name, server_ip, tcp_listen_addr, udp_listen_addr, interface_name
|
||||||
|
FROM node
|
||||||
|
WHERE id = ?
|
||||||
|
LIMIT 1
|
||||||
|
`, nodeID)
|
||||||
|
var n nodeRecord
|
||||||
|
var tcpListen sql.NullString
|
||||||
|
var udpListen sql.NullString
|
||||||
|
var iface sql.NullString
|
||||||
|
err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &tcpListen, &udpListen, &iface)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return nil, errors.New("节点不存在")
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
n.TCPListenAddr = strings.TrimSpace(tcpListen.String)
|
||||||
|
n.UDPListenAddr = strings.TrimSpace(udpListen.String)
|
||||||
|
n.InterfaceName = strings.TrimSpace(iface.String)
|
||||||
|
if n.TCPListenAddr == "" {
|
||||||
|
n.TCPListenAddr = "[::]"
|
||||||
|
}
|
||||||
|
if n.UDPListenAddr == "" {
|
||||||
|
n.UDPListenAddr = "[::]"
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(n.Name) == "" {
|
||||||
|
n.Name = fmt.Sprintf("node_%d", n.ID)
|
||||||
|
}
|
||||||
|
return &n, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int, error) {
|
||||||
|
row := h.repo.DB().QueryRow(`
|
||||||
|
SELECT ut.id, sl.speed
|
||||||
|
FROM user_tunnel ut
|
||||||
|
LEFT JOIN speed_limit sl ON sl.id = ut.speed_id
|
||||||
|
WHERE ut.user_id = ? AND ut.tunnel_id = ?
|
||||||
|
LIMIT 1
|
||||||
|
`, userID, tunnelID)
|
||||||
|
var userTunnelID int64
|
||||||
|
var speed sql.NullInt64
|
||||||
|
err := row.Scan(&userTunnelID, &speed)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return 0, nil, nil
|
||||||
|
}
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
if !speed.Valid || speed.Int64 <= 0 {
|
||||||
|
return userTunnelID, nil, nil
|
||||||
|
}
|
||||||
|
v := int(speed.Int64)
|
||||||
|
return userTunnelID, &v, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) syncForwardServices(forward *forwardRecord, method string, allowFallbackAdd bool) error {
|
||||||
|
if h == nil || forward == nil {
|
||||||
|
return errors.New("invalid forward sync context")
|
||||||
|
}
|
||||||
|
|
||||||
|
tunnel, err := h.getTunnelRecord(forward.TunnelID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
ports, err := h.listForwardPorts(forward.ID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if len(ports) == 0 {
|
||||||
|
return errors.New("转发入口端口不存在")
|
||||||
|
}
|
||||||
|
|
||||||
|
userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
|
||||||
|
|
||||||
|
for _, fp := range ports {
|
||||||
|
node, err := h.getNodeRecord(fp.NodeID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiter)
|
||||||
|
_, err = h.sendNodeCommand(node.ID, method, services, true, false)
|
||||||
|
if err != nil && allowFallbackAdd && method == "UpdateService" {
|
||||||
|
_, err = h.sendNodeCommand(node.ID, "AddService", services, true, false)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("节点 %s 下发失败: %w", node.Name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) controlForwardServices(forward *forwardRecord, commandType string, tolerateNotFound bool) error {
|
||||||
|
if h == nil || forward == nil {
|
||||||
|
return errors.New("invalid forward control context")
|
||||||
|
}
|
||||||
|
ports, err := h.listForwardPorts(forward.ID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if len(ports) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
userTunnelID, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
base := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
|
||||||
|
payload := map[string]interface{}{
|
||||||
|
"services": []string{base + "_tcp", base + "_udp"},
|
||||||
|
}
|
||||||
|
seen := map[int64]struct{}{}
|
||||||
|
for _, fp := range ports {
|
||||||
|
if _, ok := seen[fp.NodeID]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[fp.NodeID] = struct{}{}
|
||||||
|
_, err := h.sendNodeCommand(fp.NodeID, commandType, payload, false, tolerateNotFound)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) applyNodeProtocolChange(nodeID int64, httpVal, tlsVal, socksVal int) error {
|
||||||
|
_, err := h.sendNodeCommand(nodeID, "SetProtocol", map[string]interface{}{
|
||||||
|
"http": httpVal,
|
||||||
|
"tls": tlsVal,
|
||||||
|
"socks": socksVal,
|
||||||
|
}, false, false)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) sendNodeCommand(nodeID int64, commandType string, data interface{}, tolerateExists bool, tolerateNotFound bool) (ws.CommandResult, error) {
|
||||||
|
result, err := h.wsServer.SendCommand(nodeID, commandType, data, 12*time.Second)
|
||||||
|
if err == nil {
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
msg := strings.ToLower(strings.TrimSpace(err.Error()))
|
||||||
|
if tolerateExists {
|
||||||
|
if strings.Contains(msg, "exists") || strings.Contains(msg, "already") || strings.Contains(msg, "已存在") {
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if tolerateNotFound {
|
||||||
|
if strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") {
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return result, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) diagnoseForwardRuntime(forward *forwardRecord) (map[string]interface{}, error) {
|
||||||
|
if forward == nil {
|
||||||
|
return nil, errForwardNotFound
|
||||||
|
}
|
||||||
|
targets := splitRemoteTargets(forward.RemoteAddr)
|
||||||
|
if len(targets) == 0 {
|
||||||
|
return nil, errors.New("目标地址不能为空")
|
||||||
|
}
|
||||||
|
ip, port, err := parseTargetAddress(targets[0])
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("目标地址格式错误")
|
||||||
|
}
|
||||||
|
|
||||||
|
ports, err := h.listForwardPorts(forward.ID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(ports) == 0 {
|
||||||
|
return nil, errors.New("转发入口端口不存在")
|
||||||
|
}
|
||||||
|
|
||||||
|
results := make([]map[string]interface{}, 0, len(ports))
|
||||||
|
seen := map[int64]struct{}{}
|
||||||
|
for _, fp := range ports {
|
||||||
|
if _, ok := seen[fp.NodeID]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[fp.NodeID] = struct{}{}
|
||||||
|
node, nodeErr := h.getNodeRecord(fp.NodeID)
|
||||||
|
nodeName := fmt.Sprintf("node_%d", fp.NodeID)
|
||||||
|
if nodeErr == nil {
|
||||||
|
nodeName = node.Name
|
||||||
|
}
|
||||||
|
|
||||||
|
resultItem := map[string]interface{}{
|
||||||
|
"nodeName": nodeName,
|
||||||
|
"nodeId": strconv.FormatInt(fp.NodeID, 10),
|
||||||
|
"targetIp": ip,
|
||||||
|
"targetPort": port,
|
||||||
|
"averageTime": 0,
|
||||||
|
"packetLoss": 100,
|
||||||
|
}
|
||||||
|
|
||||||
|
if nodeErr != nil {
|
||||||
|
resultItem["success"] = false
|
||||||
|
resultItem["description"] = "节点信息读取失败"
|
||||||
|
resultItem["errorMessage"] = nodeErr.Error()
|
||||||
|
results = append(results, resultItem)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
pingData, pingErr := h.tcpPingViaNode(fp.NodeID, ip, port)
|
||||||
|
if pingErr != nil {
|
||||||
|
resultItem["success"] = false
|
||||||
|
resultItem["description"] = "诊断失败"
|
||||||
|
resultItem["errorMessage"] = pingErr.Error()
|
||||||
|
} else {
|
||||||
|
resultItem["success"] = asBool(pingData["success"], false)
|
||||||
|
resultItem["description"] = "诊断完成"
|
||||||
|
resultItem["averageTime"] = asFloat(pingData["averageTime"], 0)
|
||||||
|
resultItem["packetLoss"] = asFloat(pingData["packetLoss"], 100)
|
||||||
|
if msg := asString(pingData["errorMessage"]); msg != "" {
|
||||||
|
resultItem["errorMessage"] = msg
|
||||||
|
}
|
||||||
|
}
|
||||||
|
results = append(results, resultItem)
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := map[string]interface{}{
|
||||||
|
"forwardName": forward.Name,
|
||||||
|
"timestamp": time.Now().UnixMilli(),
|
||||||
|
"results": results,
|
||||||
|
}
|
||||||
|
return payload, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{}, error) {
|
||||||
|
tunnel, err := h.getTunnelRecord(tunnelID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var tunnelName string
|
||||||
|
if err := h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName); err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return nil, errors.New("隧道不存在")
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
chainRows, err := h.listChainNodesForTunnel(tunnelID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(chainRows) == 0 {
|
||||||
|
return nil, errors.New("隧道配置不完整")
|
||||||
|
}
|
||||||
|
|
||||||
|
results := make([]map[string]interface{}, 0, len(chainRows))
|
||||||
|
seen := map[int64]struct{}{}
|
||||||
|
for _, row := range chainRows {
|
||||||
|
if _, ok := seen[row.NodeID]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[row.NodeID] = struct{}{}
|
||||||
|
|
||||||
|
desc := fmt.Sprintf("节点(%s)->外网", row.NodeName)
|
||||||
|
switch row.ChainType {
|
||||||
|
case 1:
|
||||||
|
desc = fmt.Sprintf("入口(%s)->外网", row.NodeName)
|
||||||
|
case 2:
|
||||||
|
desc = fmt.Sprintf("中继(%s)->外网", row.NodeName)
|
||||||
|
case 3:
|
||||||
|
desc = fmt.Sprintf("出口(%s)->外网", row.NodeName)
|
||||||
|
}
|
||||||
|
|
||||||
|
resultItem := map[string]interface{}{
|
||||||
|
"nodeName": row.NodeName,
|
||||||
|
"nodeId": strconv.FormatInt(row.NodeID, 10),
|
||||||
|
"targetIp": "www.google.com",
|
||||||
|
"targetPort": 443,
|
||||||
|
"averageTime": 0,
|
||||||
|
"packetLoss": 100,
|
||||||
|
"fromChainType": row.ChainType,
|
||||||
|
"description": desc,
|
||||||
|
}
|
||||||
|
if row.ChainType == 2 {
|
||||||
|
resultItem["fromInx"] = row.Inx
|
||||||
|
}
|
||||||
|
|
||||||
|
pingData, pingErr := h.tcpPingViaNode(row.NodeID, "www.google.com", 443)
|
||||||
|
if pingErr != nil {
|
||||||
|
resultItem["success"] = false
|
||||||
|
resultItem["description"] = desc + " 连通性检查失败"
|
||||||
|
resultItem["errorMessage"] = pingErr.Error()
|
||||||
|
} else {
|
||||||
|
resultItem["success"] = asBool(pingData["success"], false)
|
||||||
|
resultItem["description"] = desc + " 连通性检查完成"
|
||||||
|
resultItem["averageTime"] = asFloat(pingData["averageTime"], 0)
|
||||||
|
resultItem["packetLoss"] = asFloat(pingData["packetLoss"], 100)
|
||||||
|
if msg := asString(pingData["errorMessage"]); msg != "" {
|
||||||
|
resultItem["errorMessage"] = msg
|
||||||
|
}
|
||||||
|
}
|
||||||
|
results = append(results, resultItem)
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := map[string]interface{}{
|
||||||
|
"tunnelName": tunnelName,
|
||||||
|
"tunnelType": map[bool]string{true: "端口转发", false: "隧道转发"}[tunnel.Type == 1],
|
||||||
|
"timestamp": time.Now().UnixMilli(),
|
||||||
|
"results": results,
|
||||||
|
}
|
||||||
|
return payload, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) {
|
||||||
|
rows, err := h.repo.DB().Query(`
|
||||||
|
SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, n.name
|
||||||
|
FROM chain_tunnel ct
|
||||||
|
LEFT JOIN node n ON n.id = ct.node_id
|
||||||
|
WHERE ct.tunnel_id = ?
|
||||||
|
ORDER BY ct.chain_type ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
|
||||||
|
`, tunnelID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
result := make([]chainNodeRecord, 0)
|
||||||
|
for rows.Next() {
|
||||||
|
var item chainNodeRecord
|
||||||
|
var name sql.NullString
|
||||||
|
if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &name); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(name.String) == "" {
|
||||||
|
item.NodeName = fmt.Sprintf("node_%d", item.NodeID)
|
||||||
|
} else {
|
||||||
|
item.NodeName = name.String
|
||||||
|
}
|
||||||
|
result = append(result, item)
|
||||||
|
}
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int) (map[string]interface{}, error) {
|
||||||
|
res, err := h.sendNodeCommand(nodeID, "TcpPing", map[string]interface{}{
|
||||||
|
"ip": ip,
|
||||||
|
"port": port,
|
||||||
|
"count": 4,
|
||||||
|
"timeout": 5000,
|
||||||
|
}, false, false)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if res.Data == nil {
|
||||||
|
return nil, errors.New("节点未返回诊断数据")
|
||||||
|
}
|
||||||
|
return res.Data, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func splitRemoteTargets(remoteAddr string) []string {
|
||||||
|
parts := strings.Split(remoteAddr, ",")
|
||||||
|
out := make([]string, 0, len(parts))
|
||||||
|
for _, part := range parts {
|
||||||
|
part = strings.TrimSpace(part)
|
||||||
|
if part == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
out = append(out, processServerAddress(part))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseTargetAddress(addr string) (string, int, error) {
|
||||||
|
addr = strings.TrimSpace(addr)
|
||||||
|
if addr == "" {
|
||||||
|
return "", 0, errors.New("empty address")
|
||||||
|
}
|
||||||
|
host, portStr, err := net.SplitHostPort(addr)
|
||||||
|
if err != nil {
|
||||||
|
idx := strings.LastIndex(addr, ":")
|
||||||
|
if idx <= 0 || idx >= len(addr)-1 {
|
||||||
|
return "", 0, err
|
||||||
|
}
|
||||||
|
host = strings.TrimSpace(addr[:idx])
|
||||||
|
portStr = strings.TrimSpace(addr[idx+1:])
|
||||||
|
}
|
||||||
|
port, err := strconv.Atoi(strings.TrimSpace(portStr))
|
||||||
|
if err != nil || port <= 0 || port > 65535 {
|
||||||
|
return "", 0, errors.New("invalid port")
|
||||||
|
}
|
||||||
|
host = strings.Trim(strings.TrimSpace(host), "[]")
|
||||||
|
if host == "" {
|
||||||
|
return "", 0, errors.New("invalid host")
|
||||||
|
}
|
||||||
|
return host, port, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildForwardServiceBase(forwardID, userID, userTunnelID int64) string {
|
||||||
|
return fmt.Sprintf("%d_%d_%d", forwardID, userID, userTunnelID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiter *int) []map[string]interface{} {
|
||||||
|
protocols := []string{"tcp", "udp"}
|
||||||
|
services := make([]map[string]interface{}, 0, 2)
|
||||||
|
targets := splitRemoteTargets(forward.RemoteAddr)
|
||||||
|
strategy := strings.TrimSpace(forward.Strategy)
|
||||||
|
if strategy == "" {
|
||||||
|
strategy = "fifo"
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, protocol := range protocols {
|
||||||
|
listenerAddr := node.TCPListenAddr
|
||||||
|
if protocol == "udp" {
|
||||||
|
listenerAddr = node.UDPListenAddr
|
||||||
|
}
|
||||||
|
service := map[string]interface{}{
|
||||||
|
"name": fmt.Sprintf("%s_%s", baseName, protocol),
|
||||||
|
"addr": fmt.Sprintf("%s:%d", listenerAddr, port),
|
||||||
|
"handler": map[string]interface{}{
|
||||||
|
"type": protocol,
|
||||||
|
},
|
||||||
|
"listener": map[string]interface{}{
|
||||||
|
"type": protocol,
|
||||||
|
},
|
||||||
|
"forwarder": map[string]interface{}{
|
||||||
|
"nodes": buildForwarderNodes(targets),
|
||||||
|
"selector": map[string]interface{}{
|
||||||
|
"strategy": strategy,
|
||||||
|
"maxFails": 1,
|
||||||
|
"failTimeout": "600s",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if protocol == "udp" {
|
||||||
|
service["listener"].(map[string]interface{})["metadata"] = map[string]interface{}{"keepAlive": true}
|
||||||
|
}
|
||||||
|
if tunnel != nil && tunnel.Type == 2 {
|
||||||
|
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
|
||||||
|
}
|
||||||
|
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
|
||||||
|
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
|
||||||
|
}
|
||||||
|
if limiter != nil && *limiter > 0 {
|
||||||
|
service["limiter"] = strconv.Itoa(*limiter)
|
||||||
|
}
|
||||||
|
services = append(services, service)
|
||||||
|
}
|
||||||
|
|
||||||
|
return services
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildForwarderNodes(targets []string) []map[string]interface{} {
|
||||||
|
nodes := make([]map[string]interface{}, 0, len(targets))
|
||||||
|
for i, addr := range targets {
|
||||||
|
nodes = append(nodes, map[string]interface{}{
|
||||||
|
"name": fmt.Sprintf("node_%d", i+1),
|
||||||
|
"addr": addr,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return nodes
|
||||||
|
}
|
||||||
|
|
||||||
|
func processServerAddress(serverAddr string) string {
|
||||||
|
serverAddr = strings.TrimSpace(serverAddr)
|
||||||
|
if serverAddr == "" {
|
||||||
|
return serverAddr
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(serverAddr, "[") {
|
||||||
|
return serverAddr
|
||||||
|
}
|
||||||
|
idx := strings.LastIndex(serverAddr, ":")
|
||||||
|
if idx < 0 {
|
||||||
|
if looksLikeIPv6(serverAddr) {
|
||||||
|
return "[" + serverAddr + "]"
|
||||||
|
}
|
||||||
|
return serverAddr
|
||||||
|
}
|
||||||
|
host := strings.TrimSpace(serverAddr[:idx])
|
||||||
|
port := strings.TrimSpace(serverAddr[idx+1:])
|
||||||
|
if host == "" || port == "" {
|
||||||
|
return serverAddr
|
||||||
|
}
|
||||||
|
if looksLikeIPv6(host) {
|
||||||
|
return "[" + host + "]:" + port
|
||||||
|
}
|
||||||
|
return serverAddr
|
||||||
|
}
|
||||||
|
|
||||||
|
func looksLikeIPv6(address string) bool {
|
||||||
|
return strings.Count(address, ":") >= 2
|
||||||
|
}
|
||||||
|
|
||||||
|
func asBool(v interface{}, def bool) bool {
|
||||||
|
s := strings.TrimSpace(strings.ToLower(asString(v)))
|
||||||
|
if s == "" {
|
||||||
|
return def
|
||||||
|
}
|
||||||
|
switch s {
|
||||||
|
case "1", "t", "true", "yes", "y":
|
||||||
|
return true
|
||||||
|
case "0", "f", "false", "no", "n":
|
||||||
|
return false
|
||||||
|
default:
|
||||||
|
return def
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -3,6 +3,7 @@ package handler
|
|||||||
import (
|
import (
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sort"
|
"sort"
|
||||||
@@ -63,26 +64,80 @@ func (h *Handler) WebSocketHandler() http.Handler {
|
|||||||
func (h *Handler) Register(mux *http.ServeMux) {
|
func (h *Handler) Register(mux *http.ServeMux) {
|
||||||
mux.HandleFunc("/api/v1/user/login", h.login)
|
mux.HandleFunc("/api/v1/user/login", h.login)
|
||||||
mux.HandleFunc("/api/v1/user/list", h.userList)
|
mux.HandleFunc("/api/v1/user/list", h.userList)
|
||||||
|
mux.HandleFunc("/api/v1/user/create", h.userCreate)
|
||||||
|
mux.HandleFunc("/api/v1/user/update", h.userUpdate)
|
||||||
|
mux.HandleFunc("/api/v1/user/delete", h.userDelete)
|
||||||
|
mux.HandleFunc("/api/v1/user/reset", h.userResetFlow)
|
||||||
mux.HandleFunc("/api/v1/config/get", h.getConfigByName)
|
mux.HandleFunc("/api/v1/config/get", h.getConfigByName)
|
||||||
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
||||||
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
|
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
|
||||||
mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig)
|
mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig)
|
||||||
mux.HandleFunc("/api/v1/captcha/check", h.checkCaptcha)
|
mux.HandleFunc("/api/v1/captcha/check", h.checkCaptcha)
|
||||||
|
mux.HandleFunc("/api/v1/captcha/generate", h.captchaGenerate)
|
||||||
|
mux.HandleFunc("/api/v1/captcha/verify", h.captchaVerify)
|
||||||
mux.HandleFunc("/api/v1/user/package", h.userPackage)
|
mux.HandleFunc("/api/v1/user/package", h.userPackage)
|
||||||
mux.HandleFunc("/api/v1/user/updatePassword", h.updatePassword)
|
mux.HandleFunc("/api/v1/user/updatePassword", h.updatePassword)
|
||||||
mux.HandleFunc("/api/v1/node/list", h.nodeList)
|
mux.HandleFunc("/api/v1/node/list", h.nodeList)
|
||||||
|
mux.HandleFunc("/api/v1/node/create", h.nodeCreate)
|
||||||
|
mux.HandleFunc("/api/v1/node/update", h.nodeUpdate)
|
||||||
|
mux.HandleFunc("/api/v1/node/delete", h.nodeDelete)
|
||||||
|
mux.HandleFunc("/api/v1/node/install", h.nodeInstall)
|
||||||
|
mux.HandleFunc("/api/v1/node/update-order", h.nodeUpdateOrder)
|
||||||
|
mux.HandleFunc("/api/v1/node/batch-delete", h.nodeBatchDelete)
|
||||||
|
mux.HandleFunc("/api/v1/node/check-status", h.nodeCheckStatus)
|
||||||
mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
|
mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/update", h.tunnelUpdate)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/delete", h.tunnelDelete)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/diagnose", h.tunnelDiagnose)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/update-order", h.tunnelUpdateOrder)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/batch-delete", h.tunnelBatchDelete)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/batch-redeploy", h.tunnelBatchRedeploy)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/user/assign", h.userTunnelAssign)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/user/batch-assign", h.userTunnelBatchAssign)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/user/remove", h.userTunnelRemove)
|
||||||
|
mux.HandleFunc("/api/v1/tunnel/user/update", h.userTunnelUpdate)
|
||||||
mux.HandleFunc("/api/v1/forward/list", h.forwardList)
|
mux.HandleFunc("/api/v1/forward/list", h.forwardList)
|
||||||
|
mux.HandleFunc("/api/v1/forward/create", h.forwardCreate)
|
||||||
|
mux.HandleFunc("/api/v1/forward/update", h.forwardUpdate)
|
||||||
|
mux.HandleFunc("/api/v1/forward/delete", h.forwardDelete)
|
||||||
|
mux.HandleFunc("/api/v1/forward/force-delete", h.forwardForceDelete)
|
||||||
|
mux.HandleFunc("/api/v1/forward/pause", h.forwardPause)
|
||||||
|
mux.HandleFunc("/api/v1/forward/resume", h.forwardResume)
|
||||||
|
mux.HandleFunc("/api/v1/forward/diagnose", h.forwardDiagnose)
|
||||||
|
mux.HandleFunc("/api/v1/forward/update-order", h.forwardUpdateOrder)
|
||||||
|
mux.HandleFunc("/api/v1/forward/batch-delete", h.forwardBatchDelete)
|
||||||
|
mux.HandleFunc("/api/v1/forward/batch-pause", h.forwardBatchPause)
|
||||||
|
mux.HandleFunc("/api/v1/forward/batch-resume", h.forwardBatchResume)
|
||||||
|
mux.HandleFunc("/api/v1/forward/batch-redeploy", h.forwardBatchRedeploy)
|
||||||
|
mux.HandleFunc("/api/v1/forward/batch-change-tunnel", h.forwardBatchChangeTunnel)
|
||||||
mux.HandleFunc("/api/v1/speed-limit/list", h.speedLimitList)
|
mux.HandleFunc("/api/v1/speed-limit/list", h.speedLimitList)
|
||||||
|
mux.HandleFunc("/api/v1/speed-limit/create", h.speedLimitCreate)
|
||||||
|
mux.HandleFunc("/api/v1/speed-limit/update", h.speedLimitUpdate)
|
||||||
|
mux.HandleFunc("/api/v1/speed-limit/delete", h.speedLimitDelete)
|
||||||
|
mux.HandleFunc("/api/v1/speed-limit/tunnels", h.tunnelList)
|
||||||
mux.HandleFunc("/api/v1/tunnel/user/tunnel", h.userTunnelVisibleList)
|
mux.HandleFunc("/api/v1/tunnel/user/tunnel", h.userTunnelVisibleList)
|
||||||
mux.HandleFunc("/api/v1/tunnel/user/list", h.userTunnelList)
|
mux.HandleFunc("/api/v1/tunnel/user/list", h.userTunnelList)
|
||||||
mux.HandleFunc("/api/v1/group/tunnel/list", h.tunnelGroupList)
|
mux.HandleFunc("/api/v1/group/tunnel/list", h.tunnelGroupList)
|
||||||
|
mux.HandleFunc("/api/v1/group/tunnel/create", h.groupTunnelCreate)
|
||||||
|
mux.HandleFunc("/api/v1/group/tunnel/update", h.groupTunnelUpdate)
|
||||||
|
mux.HandleFunc("/api/v1/group/tunnel/delete", h.groupTunnelDelete)
|
||||||
|
mux.HandleFunc("/api/v1/group/tunnel/assign", h.groupTunnelAssign)
|
||||||
mux.HandleFunc("/api/v1/group/user/list", h.userGroupList)
|
mux.HandleFunc("/api/v1/group/user/list", h.userGroupList)
|
||||||
|
mux.HandleFunc("/api/v1/group/user/create", h.groupUserCreate)
|
||||||
|
mux.HandleFunc("/api/v1/group/user/update", h.groupUserUpdate)
|
||||||
|
mux.HandleFunc("/api/v1/group/user/delete", h.groupUserDelete)
|
||||||
|
mux.HandleFunc("/api/v1/group/user/assign", h.groupUserAssign)
|
||||||
mux.HandleFunc("/api/v1/group/permission/list", h.groupPermissionList)
|
mux.HandleFunc("/api/v1/group/permission/list", h.groupPermissionList)
|
||||||
|
mux.HandleFunc("/api/v1/group/permission/assign", h.groupPermissionAssign)
|
||||||
|
mux.HandleFunc("/api/v1/group/permission/remove", h.groupPermissionRemove)
|
||||||
|
mux.HandleFunc("/api/v1/open_api/sub_store", h.openAPISubStore)
|
||||||
|
|
||||||
mux.HandleFunc("/flow/test", h.flowTest)
|
mux.HandleFunc("/flow/test", h.flowTest)
|
||||||
mux.HandleFunc("/flow/config", h.flowConfig)
|
mux.HandleFunc("/flow/config", h.flowConfig)
|
||||||
mux.HandleFunc("/flow/upload", h.flowUpload)
|
mux.HandleFunc("/flow/upload", h.flowUpload)
|
||||||
|
mux.HandleFunc("/error", h.errorPage)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
||||||
@@ -240,11 +295,26 @@ func (h *Handler) forwardList(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
userID, roleID, err := userRoleFromRequest(r)
|
||||||
|
if err != nil {
|
||||||
|
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
items, err := h.repo.ListForwards()
|
items, err := h.repo.ListForwards()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if roleID != 0 {
|
||||||
|
filtered := make([]map[string]interface{}, 0, len(items))
|
||||||
|
for _, item := range items {
|
||||||
|
if asInt64(item["userId"], 0) == userID {
|
||||||
|
filtered = append(filtered, item)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
items = filtered
|
||||||
|
}
|
||||||
response.WriteJSON(w, response.OK(items))
|
response.WriteJSON(w, response.OK(items))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -262,6 +332,92 @@ func (h *Handler) speedLimitList(w http.ResponseWriter, r *http.Request) {
|
|||||||
response.WriteJSON(w, response.OK(items))
|
response.WriteJSON(w, response.OK(items))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodGet {
|
||||||
|
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if h == nil || h.repo == nil || h.repo.DB() == nil {
|
||||||
|
response.WriteJSON(w, response.Err(-2, "database unavailable"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
username := strings.TrimSpace(r.URL.Query().Get("user"))
|
||||||
|
password := strings.TrimSpace(r.URL.Query().Get("pwd"))
|
||||||
|
tunnel := strings.TrimSpace(r.URL.Query().Get("tunnel"))
|
||||||
|
if tunnel == "" {
|
||||||
|
tunnel = "-1"
|
||||||
|
}
|
||||||
|
|
||||||
|
if username == "" {
|
||||||
|
response.WriteJSON(w, response.ErrDefault("用户不能为空"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if password == "" {
|
||||||
|
response.WriteJSON(w, response.ErrDefault("密码不能为空"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
user, err := h.repo.GetUserByUsername(username)
|
||||||
|
if err != nil {
|
||||||
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if user == nil || user.Pwd != security.MD5(password) {
|
||||||
|
response.WriteJSON(w, response.ErrDefault("鉴权失败"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
const giga = int64(1024 * 1024 * 1024)
|
||||||
|
headerValue := ""
|
||||||
|
|
||||||
|
if tunnel == "-1" {
|
||||||
|
headerValue = buildSubscriptionHeader(user.OutFlow, user.InFlow, user.Flow*giga, user.ExpTime/1000)
|
||||||
|
} else {
|
||||||
|
tunnelID, parseErr := strconv.ParseInt(tunnel, 10, 64)
|
||||||
|
if parseErr != nil || tunnelID <= 0 {
|
||||||
|
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var userID int64
|
||||||
|
var inFlow int64
|
||||||
|
var outFlow int64
|
||||||
|
var flow int64
|
||||||
|
var expTime int64
|
||||||
|
err = h.repo.DB().QueryRow(`SELECT user_id, in_flow, out_flow, flow, exp_time FROM user_tunnel WHERE id = ? LIMIT 1`, tunnelID).
|
||||||
|
Scan(&userID, &inFlow, &outFlow, &flow, &expTime)
|
||||||
|
if err != nil {
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if userID != user.ID {
|
||||||
|
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
headerValue = buildSubscriptionHeader(outFlow, inFlow, flow*giga, expTime/1000)
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("subscription-userinfo", headerValue)
|
||||||
|
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||||
|
_, _ = w.Write([]byte(headerValue))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) errorPage(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "text/html; charset=UTF-8")
|
||||||
|
w.WriteHeader(http.StatusNotFound)
|
||||||
|
_, _ = w.Write([]byte("<!DOCTYPE html><html lang='zh-CN'><head><meta charset='UTF-8'><meta name='viewport' content='width=device-width, initial-scale=1.0'><title>错误 404</title></head><body><div style='min-height:100vh;display:flex;align-items:center;justify-content:center;flex-direction:column;font-family:-apple-system,BlinkMacSystemFont,Segoe UI,Arial,sans-serif;'><div style='font-size:6rem;color:#333;font-weight:300;'>404</div><div style='font-size:1.2rem;color:#666;'>你推开了后端的大门,却发现里面只有寂寞。</div></div></body></html>"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildSubscriptionHeader(upload, download, total, expire int64) string {
|
||||||
|
return fmt.Sprintf("upload=%d; download=%d; total=%d; expire=%d", download, upload, total, expire)
|
||||||
|
}
|
||||||
|
|
||||||
func (h *Handler) userTunnelVisibleList(w http.ResponseWriter, r *http.Request) {
|
func (h *Handler) userTunnelVisibleList(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.Method != http.MethodPost {
|
if r.Method != http.MethodPost {
|
||||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||||
@@ -738,6 +894,18 @@ func userIDFromRequest(r *http.Request) (int64, error) {
|
|||||||
return parseUserID(claims.Sub)
|
return parseUserID(claims.Sub)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func userRoleFromRequest(r *http.Request) (int64, int, error) {
|
||||||
|
claims, ok := r.Context().Value(middleware.ClaimsContextKey).(auth.Claims)
|
||||||
|
if !ok {
|
||||||
|
return 0, 0, strconv.ErrSyntax
|
||||||
|
}
|
||||||
|
userID, err := parseUserID(claims.Sub)
|
||||||
|
if err != nil {
|
||||||
|
return 0, 0, err
|
||||||
|
}
|
||||||
|
return userID, claims.RoleID, nil
|
||||||
|
}
|
||||||
|
|
||||||
func nullableNullInt64(v sql.NullInt64) interface{} {
|
func nullableNullInt64(v sql.NullInt64) interface{} {
|
||||||
if v.Valid {
|
if v.Valid {
|
||||||
return v.Int64
|
return v.Int64
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -25,6 +25,13 @@ type Repository struct {
|
|||||||
db *sql.DB
|
db *sql.DB
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *Repository) DB() *sql.DB {
|
||||||
|
if r == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return r.db
|
||||||
|
}
|
||||||
|
|
||||||
type User struct {
|
type User struct {
|
||||||
ID int64
|
ID int64
|
||||||
User string
|
User string
|
||||||
|
|||||||
@@ -2,11 +2,14 @@ package ws
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
|
|
||||||
@@ -38,15 +41,36 @@ type nodeSession struct {
|
|||||||
conn *connWrap
|
conn *connWrap
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type commandResponse struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Success bool `json:"success"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
Data json.RawMessage `json:"data,omitempty"`
|
||||||
|
RequestID string `json:"requestId,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type pendingRequest struct {
|
||||||
|
nodeID int64
|
||||||
|
ch chan CommandResult
|
||||||
|
}
|
||||||
|
|
||||||
|
type CommandResult struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Success bool `json:"success"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
Data map[string]interface{} `json:"data,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
type Server struct {
|
type Server struct {
|
||||||
repo *sqlite.Repository
|
repo *sqlite.Repository
|
||||||
jwtSecret string
|
jwtSecret string
|
||||||
upgrader websocket.Upgrader
|
upgrader websocket.Upgrader
|
||||||
|
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
admins map[*connWrap]struct{}
|
admins map[*connWrap]struct{}
|
||||||
nodes map[int64]*nodeSession
|
nodes map[int64]*nodeSession
|
||||||
byConn map[*websocket.Conn]*nodeSession
|
byConn map[*websocket.Conn]*nodeSession
|
||||||
|
pending map[string]pendingRequest
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewServer(repo *sqlite.Repository, jwtSecret string) *Server {
|
func NewServer(repo *sqlite.Repository, jwtSecret string) *Server {
|
||||||
@@ -56,9 +80,10 @@ func NewServer(repo *sqlite.Repository, jwtSecret string) *Server {
|
|||||||
upgrader: websocket.Upgrader{
|
upgrader: websocket.Upgrader{
|
||||||
CheckOrigin: func(r *http.Request) bool { return true },
|
CheckOrigin: func(r *http.Request) bool { return true },
|
||||||
},
|
},
|
||||||
admins: make(map[*connWrap]struct{}),
|
admins: make(map[*connWrap]struct{}),
|
||||||
nodes: make(map[int64]*nodeSession),
|
nodes: make(map[int64]*nodeSession),
|
||||||
byConn: make(map[*websocket.Conn]*nodeSession),
|
byConn: make(map[*websocket.Conn]*nodeSession),
|
||||||
|
pending: make(map[string]pendingRequest),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -150,6 +175,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
|||||||
delete(s.byConn, conn)
|
delete(s.byConn, conn)
|
||||||
s.mu.Unlock()
|
s.mu.Unlock()
|
||||||
if needOfflineBroadcast {
|
if needOfflineBroadcast {
|
||||||
|
s.failPendingForNode(nodeID, "节点连接已断开")
|
||||||
_ = s.repo.UpdateNodeStatus(nodeID, 0)
|
_ = s.repo.UpdateNodeStatus(nodeID, 0)
|
||||||
s.broadcastStatus(nodeID, 0)
|
s.broadcastStatus(nodeID, 0)
|
||||||
}
|
}
|
||||||
@@ -163,10 +189,186 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
|||||||
}
|
}
|
||||||
|
|
||||||
msg := decryptIfNeeded(payload, secret)
|
msg := decryptIfNeeded(payload, secret)
|
||||||
|
s.tryResolvePending(nodeID, msg)
|
||||||
s.broadcastInfo(nodeID, msg)
|
s.broadcastInfo(nodeID, msg)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, timeout time.Duration) (CommandResult, error) {
|
||||||
|
if s == nil {
|
||||||
|
return CommandResult{}, errors.New("server not initialized")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(cmdType) == "" {
|
||||||
|
return CommandResult{}, errors.New("command type is empty")
|
||||||
|
}
|
||||||
|
if timeout <= 0 {
|
||||||
|
timeout = 10 * time.Second
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.RLock()
|
||||||
|
ns, ok := s.nodes[nodeID]
|
||||||
|
s.mu.RUnlock()
|
||||||
|
if !ok || ns == nil || ns.conn == nil || ns.conn.conn == nil {
|
||||||
|
return CommandResult{}, errors.New("节点不在线")
|
||||||
|
}
|
||||||
|
|
||||||
|
requestID := fmt.Sprintf("%d_%d", nodeID, time.Now().UnixNano())
|
||||||
|
ch := make(chan CommandResult, 1)
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
s.pending[requestID] = pendingRequest{nodeID: nodeID, ch: ch}
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
cleanup := func() {
|
||||||
|
s.mu.Lock()
|
||||||
|
if p, exists := s.pending[requestID]; exists {
|
||||||
|
delete(s.pending, requestID)
|
||||||
|
close(p.ch)
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
cmdPayload := map[string]interface{}{
|
||||||
|
"type": cmdType,
|
||||||
|
"data": data,
|
||||||
|
"requestId": requestID,
|
||||||
|
}
|
||||||
|
rawCmd, err := json.Marshal(cmdPayload)
|
||||||
|
if err != nil {
|
||||||
|
cleanup()
|
||||||
|
return CommandResult{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
messageData := rawCmd
|
||||||
|
if strings.TrimSpace(ns.secret) != "" {
|
||||||
|
crypto, err := security.NewAESCrypto(ns.secret)
|
||||||
|
if err != nil {
|
||||||
|
cleanup()
|
||||||
|
return CommandResult{}, err
|
||||||
|
}
|
||||||
|
encrypted, err := crypto.Encrypt(rawCmd)
|
||||||
|
if err != nil {
|
||||||
|
cleanup()
|
||||||
|
return CommandResult{}, err
|
||||||
|
}
|
||||||
|
wrapper := map[string]interface{}{
|
||||||
|
"encrypted": true,
|
||||||
|
"data": encrypted,
|
||||||
|
"timestamp": time.Now().UnixMilli(),
|
||||||
|
}
|
||||||
|
messageData, err = json.Marshal(wrapper)
|
||||||
|
if err != nil {
|
||||||
|
cleanup()
|
||||||
|
return CommandResult{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ns.conn.mu.Lock()
|
||||||
|
err = ns.conn.conn.WriteMessage(websocket.TextMessage, messageData)
|
||||||
|
ns.conn.mu.Unlock()
|
||||||
|
if err != nil {
|
||||||
|
cleanup()
|
||||||
|
return CommandResult{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case result, ok := <-ch:
|
||||||
|
if !ok {
|
||||||
|
return CommandResult{}, errors.New("命令通道已关闭")
|
||||||
|
}
|
||||||
|
if !result.Success {
|
||||||
|
if strings.TrimSpace(result.Message) == "" {
|
||||||
|
result.Message = "命令执行失败"
|
||||||
|
}
|
||||||
|
return result, errors.New(result.Message)
|
||||||
|
}
|
||||||
|
return result, nil
|
||||||
|
case <-time.After(timeout):
|
||||||
|
cleanup()
|
||||||
|
return CommandResult{}, errors.New("等待节点响应超时")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) tryResolvePending(nodeID int64, message string) {
|
||||||
|
if s == nil || strings.TrimSpace(message) == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var resp commandResponse
|
||||||
|
if err := json.Unmarshal([]byte(message), &resp); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(resp.RequestID) == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
p, ok := s.pending[resp.RequestID]
|
||||||
|
if ok {
|
||||||
|
delete(s.pending, resp.RequestID)
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if p.nodeID != nodeID {
|
||||||
|
select {
|
||||||
|
case p.ch <- CommandResult{Type: resp.Type, Success: false, Message: "节点响应与请求不匹配"}:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
close(p.ch)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
result := CommandResult{
|
||||||
|
Type: resp.Type,
|
||||||
|
Success: resp.Success,
|
||||||
|
Message: resp.Message,
|
||||||
|
}
|
||||||
|
if len(resp.Data) > 0 {
|
||||||
|
var data map[string]interface{}
|
||||||
|
if err := json.Unmarshal(resp.Data, &data); err == nil {
|
||||||
|
result.Data = data
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case p.ch <- result:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
close(p.ch)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) failPendingForNode(nodeID int64, message string) {
|
||||||
|
if s == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
type pair struct {
|
||||||
|
id string
|
||||||
|
pr pendingRequest
|
||||||
|
}
|
||||||
|
items := make([]pair, 0)
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
for id, pr := range s.pending {
|
||||||
|
if pr.nodeID != nodeID {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
items = append(items, pair{id: id, pr: pr})
|
||||||
|
delete(s.pending, id)
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
for _, item := range items {
|
||||||
|
select {
|
||||||
|
case item.pr.ch <- CommandResult{Success: false, Message: message}:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
close(item.pr.ch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Server) broadcastStatus(nodeID int64, status int) {
|
func (s *Server) broadcastStatus(nodeID int64, status int) {
|
||||||
payload := map[string]interface{}{
|
payload := map[string]interface{}{
|
||||||
"id": strconv.FormatInt(nodeID, 10),
|
"id": strconv.FormatInt(nodeID, 10),
|
||||||
|
|||||||
@@ -0,0 +1,128 @@
|
|||||||
|
package contract_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"go-backend/internal/auth"
|
||||||
|
"go-backend/internal/http/response"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestForwardOwnershipAndScopeContracts(t *testing.T) {
|
||||||
|
secret := "contract-jwt-secret"
|
||||||
|
router, repo := setupContractRouter(t, secret)
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
|
||||||
|
if _, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||||
|
VALUES(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||||
|
`, now, now); err != nil {
|
||||||
|
t.Fatalf("insert user: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
res, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
`, "contract-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("insert tunnel: %v", err)
|
||||||
|
}
|
||||||
|
tunnelID, err := res.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get tunnel id: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
resAdmin, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
||||||
|
`, 1, "admin_user", "admin-forward", tunnelID, "1.1.1.1:443", "fifo", now, now, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("insert admin forward: %v", err)
|
||||||
|
}
|
||||||
|
adminForwardID, err := resAdmin.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get admin forward id: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
resUser, err := repo.DB().Exec(`
|
||||||
|
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
||||||
|
`, 2, "normal_user", "user-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 1)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("insert user forward: %v", err)
|
||||||
|
}
|
||||||
|
userForwardID, err := resUser.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get user forward id: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate user token: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("non-owner cannot delete another user's forward", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/delete", bytes.NewBufferString(`{"id":`+jsonNumber(adminForwardID)+`}`))
|
||||||
|
req.Header.Set("Authorization", userToken)
|
||||||
|
res := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(res, req)
|
||||||
|
|
||||||
|
assertCodeMsg(t, res, -1, "转发不存在")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("non-admin forward list is scoped to owner", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/list", bytes.NewBufferString(`{}`))
|
||||||
|
req.Header.Set("Authorization", userToken)
|
||||||
|
res := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(res, req)
|
||||||
|
|
||||||
|
var out response.R
|
||||||
|
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||||
|
t.Fatalf("decode response: %v", err)
|
||||||
|
}
|
||||||
|
if out.Code != 0 {
|
||||||
|
t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
arr, ok := out.Data.([]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected array data, got %T", out.Data)
|
||||||
|
}
|
||||||
|
if len(arr) != 1 {
|
||||||
|
t.Fatalf("expected 1 forward, got %d", len(arr))
|
||||||
|
}
|
||||||
|
item, ok := arr[0].(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected object item, got %T", arr[0])
|
||||||
|
}
|
||||||
|
if got := int64(item["id"].(float64)); got != userForwardID {
|
||||||
|
t.Fatalf("expected forward id %d, got %d", userForwardID, got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("diagnose no longer returns hardcoded success", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/diagnose", bytes.NewBufferString(`{"forwardId":`+jsonNumber(userForwardID)+`}`))
|
||||||
|
req.Header.Set("Authorization", userToken)
|
||||||
|
res := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(res, req)
|
||||||
|
|
||||||
|
var out response.R
|
||||||
|
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||||
|
t.Fatalf("decode response: %v", err)
|
||||||
|
}
|
||||||
|
if out.Code == 0 {
|
||||||
|
t.Fatalf("expected non-zero code for missing runtime path, got success")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func jsonNumber(v int64) string {
|
||||||
|
return strconv.FormatInt(v, 10)
|
||||||
|
}
|
||||||
@@ -0,0 +1,154 @@
|
|||||||
|
package contract_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"go-backend/internal/auth"
|
||||||
|
httpserver "go-backend/internal/http"
|
||||||
|
"go-backend/internal/http/handler"
|
||||||
|
"go-backend/internal/http/response"
|
||||||
|
"go-backend/internal/store/sqlite"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestOpenAPISubStoreContracts(t *testing.T) {
|
||||||
|
router, repo := setupContractRouter(t, "contract-jwt-secret")
|
||||||
|
|
||||||
|
const tunnelFlowGB = int64(500)
|
||||||
|
const tunnelInFlow = int64(123)
|
||||||
|
const tunnelOutFlow = int64(456)
|
||||||
|
const tunnelExpTimeMs = int64(2727251700000)
|
||||||
|
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
res, err := repo.DB().Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||||
|
"contract-tunnel", 1.0, 1, "tls", 1, now, now, 1, nil, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("insert tunnel: %v", err)
|
||||||
|
}
|
||||||
|
tunnelID, err := res.LastInsertId()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("last insert id: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := repo.DB().Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, ?, ?, ?, ?, ?)`,
|
||||||
|
1, tunnelID, 99999, tunnelFlowGB, tunnelInFlow, tunnelOutFlow, 1, tunnelExpTimeMs, 1); err != nil {
|
||||||
|
t.Fatalf("insert user_tunnel: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("default user subscription payload", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user", nil)
|
||||||
|
resp := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(resp, req)
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read body: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
expected := "upload=0; download=0; total=107373108658176; expire=2727251700"
|
||||||
|
if string(body) != expected {
|
||||||
|
t.Fatalf("expected body %q, got %q", expected, string(body))
|
||||||
|
}
|
||||||
|
if got := resp.Header().Get("subscription-userinfo"); got != expected {
|
||||||
|
t.Fatalf("expected subscription-userinfo %q, got %q", expected, got)
|
||||||
|
}
|
||||||
|
if !strings.Contains(resp.Header().Get("Content-Type"), "text/plain") {
|
||||||
|
t.Fatalf("expected text/plain content type, got %q", resp.Header().Get("Content-Type"))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("tunnel scoped subscription payload", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user&tunnel="+strconv.FormatInt(tunnelID, 10), nil)
|
||||||
|
resp := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(resp, req)
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read body: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
expected := "upload=123; download=456; total=536870912000; expire=2727251700"
|
||||||
|
if string(body) != expected {
|
||||||
|
t.Fatalf("expected body %q, got %q", expected, string(body))
|
||||||
|
}
|
||||||
|
if got := resp.Header().Get("subscription-userinfo"); got != expected {
|
||||||
|
t.Fatalf("expected subscription-userinfo %q, got %q", expected, got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("invalid credentials returns contract error", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=wrong", nil)
|
||||||
|
resp := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(resp, req)
|
||||||
|
|
||||||
|
assertCodeMsg(t, resp, -1, "鉴权失败")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("missing tunnel returns contract error", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user&tunnel=999999", nil)
|
||||||
|
resp := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(resp, req)
|
||||||
|
|
||||||
|
assertCodeMsg(t, resp, -1, "隧道不存在")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSpeedLimitTunnelsRouteAlias(t *testing.T) {
|
||||||
|
secret := "contract-jwt-secret"
|
||||||
|
router, _ := setupContractRouter(t, secret)
|
||||||
|
|
||||||
|
t.Run("missing token blocked", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil)
|
||||||
|
resp := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(resp, req)
|
||||||
|
|
||||||
|
assertCodeMsg(t, resp, 401, "未登录或token已过期")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("admin token receives success envelope", func(t *testing.T) {
|
||||||
|
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate token: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil)
|
||||||
|
req.Header.Set("Authorization", token)
|
||||||
|
resp := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(resp, req)
|
||||||
|
|
||||||
|
var out response.R
|
||||||
|
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||||
|
t.Fatalf("decode response: %v", err)
|
||||||
|
}
|
||||||
|
if out.Code != 0 {
|
||||||
|
t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) {
|
||||||
|
t.Helper()
|
||||||
|
dbPath := filepath.Join(t.TempDir(), "contract.db")
|
||||||
|
repo, err := sqlite.Open(dbPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open sqlite: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_ = repo.Close()
|
||||||
|
})
|
||||||
|
|
||||||
|
h := handler.New(repo, jwtSecret)
|
||||||
|
return httpserver.NewRouter(h, jwtSecret), repo
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user