This commit is contained in:
qaq
2025-06-12 15:32:13 +08:00
commit 959ff9e2f2
146 changed files with 35933 additions and 0 deletions
@@ -0,0 +1,21 @@
package com.admin;
import com.baomidou.mybatisplus.annotation.DbType;
import com.baomidou.mybatisplus.extension.plugins.MybatisPlusInterceptor;
import com.baomidou.mybatisplus.extension.plugins.inner.PaginationInnerInterceptor;
import org.mybatis.spring.annotation.MapperScan;
import org.springframework.boot.SpringApplication;
import org.springframework.boot.autoconfigure.SpringBootApplication;
import org.springframework.context.annotation.Bean;
import org.springframework.scheduling.annotation.EnableAsync;
@SpringBootApplication
@EnableAsync
public class AdminApplication {
public static void main(String[] args) {
SpringApplication.run(AdminApplication.class, args);
}
}
@@ -0,0 +1,148 @@
package com.admin;
import com.baomidou.mybatisplus.core.exceptions.MybatisPlusException;
import com.baomidou.mybatisplus.core.toolkit.StringPool;
import com.baomidou.mybatisplus.core.toolkit.StringUtils;
import com.baomidou.mybatisplus.generator.AutoGenerator;
import com.baomidou.mybatisplus.generator.InjectionConfig;
import com.baomidou.mybatisplus.generator.config.*;
import com.baomidou.mybatisplus.generator.config.po.TableInfo;
import com.baomidou.mybatisplus.generator.config.rules.NamingStrategy;
import com.baomidou.mybatisplus.generator.engine.FreemarkerTemplateEngine;
import java.util.ArrayList;
import java.util.List;
import java.util.Scanner;
// 演示例子,执行 main 方法控制台输入模块表名回车自动生成对应项目目录中
public class CodeGenerator {
/**
* <p>
* 读取控制台内容
* </p>
*/
public static String scanner(String tip) {
Scanner scanner = new Scanner(System.in);
StringBuilder help = new StringBuilder();
help.append("请输入" + tip + ":");
System.out.println(help.toString());
if (scanner.hasNext()) {
String ipt = scanner.next();
if (StringUtils.isNotBlank(ipt)) {
return ipt;
}
}
throw new MybatisPlusException("请输入正确的" + tip + "!");
}
public static void main(String[] args) {
// 代码生成器
AutoGenerator mpg = new AutoGenerator();
// 全局配置
GlobalConfig gc = new GlobalConfig();
String projectPath = System.getProperty("user.dir");
gc.setOutputDir(projectPath + "/src/main/java");
gc.setAuthor("QAQ");
gc.setOpen(false);
// gc.setSwagger2(true); 实体属性 Swagger2 注解
gc.setServiceName("%sService");
mpg.setGlobalConfig(gc);
// 数据源配置 - 使用环境变量
DataSourceConfig dsc = new DataSourceConfig();
String dbHost = System.getenv("DB_HOST");
String dbName = System.getenv("DB_NAME");
String dbUser = System.getenv("DB_USER");
String dbPassword = System.getenv("DB_PASSWORD");
if (dbHost == null || dbName == null || dbUser == null || dbPassword == null) {
throw new MybatisPlusException("请设置数据库环境变量: DB_HOST, DB_NAME, DB_USER, DB_PASSWORD");
}
dsc.setUrl("jdbc:mysql://" + dbHost + "/" + dbName + "?useUnicode=true&useSSL=false&characterEncoding=utf8&serverTimezone=Asia/Shanghai");
dsc.setDriverName("com.mysql.cj.jdbc.Driver");
dsc.setUsername(dbUser);
dsc.setPassword(dbPassword);
mpg.setDataSource(dsc);
// 包配置
PackageConfig pc = new PackageConfig();
// pc.setModuleName(scanner("模块名"));
pc.setParent("com.admin");
mpg.setPackageInfo(pc);
// 自定义配置
InjectionConfig cfg = new InjectionConfig() {
@Override
public void initMap() {
// to do nothing
}
};
// 如果模板引擎是 freemarker
String templatePath = "/templates/mapper.xml.ftl";
// 如果模板引擎是 velocity
// String templatePath = "/templates/mapper.xml.vm";
// 自定义输出配置
List<FileOutConfig> focList = new ArrayList<>();
// 自定义配置会被优先输出
focList.add(new FileOutConfig(templatePath) {
@Override
public String outputFile(TableInfo tableInfo) {
// 自定义输出文件名 , 如果你 Entity 设置了前后缀、此处注意 xml 的名称会跟着发生变化!!
return projectPath + "/src/main/resources/mapper/" + pc.getModuleName()
+ "/" + tableInfo.getEntityName() + "Mapper" + StringPool.DOT_XML;
}
});
/*
cfg.setFileCreate(new IFileCreate() {
@Override
public boolean isCreate(ConfigBuilder configBuilder, FileType fileType, String filePath) {
// 判断自定义文件夹是否需要创建
checkDir("调用默认方法创建的目录,自定义目录用");
if (fileType == FileType.MAPPER) {
// 已经生成 mapper 文件判断存在,不想重新生成返回 false
return !new File(filePath).exists();
}
// 允许生成模板文件
return true;
}
});
*/
cfg.setFileOutConfigList(focList);
mpg.setCfg(cfg);
// 配置模板
TemplateConfig templateConfig = new TemplateConfig();
// 配置自定义输出模板
//指定自定义模板路径,注意不要带上.ftl/.vm, 会根据使用的模板引擎自动识别
// templateConfig.setEntity("templates/entity2.java");
// templateConfig.setService();
// templateConfig.setController();
templateConfig.setXml(null);
mpg.setTemplate(templateConfig);
// 策略配置
StrategyConfig strategy = new StrategyConfig();
strategy.setNaming(NamingStrategy.underline_to_camel);
strategy.setColumnNaming(NamingStrategy.underline_to_camel);
strategy.setSuperEntityClass("com.admin.entity.BaseEntity");
strategy.setEntityLombokModel(true);
strategy.setRestControllerStyle(true);
// 公共父类
strategy.setSuperControllerClass("com.admin.controller.BaseController");
strategy.setSuperEntityColumns("id", "created_time", "updated_time", "status");
strategy.setInclude(scanner("表名,多个英文逗号分割").split(","));
strategy.setControllerMappingHyphenStyle(true);
// strategy.setTablePrefix("sys_");//动态调整
mpg.setStrategy(strategy);
mpg.setTemplateEngine(new FreemarkerTemplateEngine());
mpg.execute();
}
}
@@ -0,0 +1,15 @@
package com.admin.common.annotation;
import java.lang.annotation.ElementType;
import java.lang.annotation.Retention;
import java.lang.annotation.RetentionPolicy;
import java.lang.annotation.Target;
/**
* 权限控制注解
* 用于标记需要管理员权限的方法(role_id = 0)
*/
@Target(ElementType.METHOD)
@Retention(RetentionPolicy.RUNTIME)
public @interface RequireRole {
}
@@ -0,0 +1,9 @@
package com.admin.common.aop;
import java.lang.annotation.*;
@Target({ElementType.METHOD})
@Retention(RetentionPolicy.RUNTIME)
@Documented
public @interface LogAnnotation {}
@@ -0,0 +1,193 @@
package com.admin.common.aop;
import cn.hutool.core.util.ArrayUtil;
import com.admin.common.utils.JwtUtil;
import com.alibaba.fastjson.JSON;
import com.admin.common.utils.HttpContextUtils;
import com.admin.common.utils.IpUtils;
import lombok.extern.slf4j.Slf4j;
import org.aspectj.lang.JoinPoint;
import org.aspectj.lang.annotation.*;
import org.aspectj.lang.reflect.CodeSignature;
import org.aspectj.lang.reflect.MethodSignature;
import org.springframework.stereotype.Component;
import javax.servlet.http.HttpServletRequest;
import java.lang.reflect.Method;
import java.util.Arrays;
import java.util.HashMap;
import java.util.Map;
@Component
@Aspect
@Slf4j
public class LogAspect {
@Pointcut("@annotation(com.admin.common.aop.LogAnnotation)")
public void pt() {
}
/**
* 返回后通知(@AfterReturning):在某连接点(joinpoint)
* 正常完成后执行的通知:例如,一个方法没有抛出任何异常,正常返回
* 方法执行完毕之后
* 注意在这里不能使用ProceedingJoinPoint
* 不然会报错ProceedingJoinPoint is only supported for around advice
* crmAspect()指向需要控制的方法
* returning 注解返回值
*
* @param joinPoint
* @param returnValue 返回值
* @throws Exception
*/
@AfterReturning(value = "pt()", returning = "returnValue")
public void log(JoinPoint joinPoint, Object returnValue) throws Throwable {
// 获取请求信息
HttpServletRequest request = HttpContextUtils.getHttpServletRequest();
// 获取请求方法类型(POST/GET等)
String requestMethod = request.getMethod();
// 获取用户ID
String authorization = request.getHeader("Authorization") + "";
Object user_id = "未登录"; // 请求用户的id
if (!authorization.equals("null")) {
user_id = JwtUtil.getUserIdFromToken(authorization);
}
// 获取请求IP
String ipAddr = IpUtils.getIpAddr(request);
// 获取方法签名信息
MethodSignature signature = (MethodSignature) joinPoint.getSignature();
Method method = signature.getMethod();
// 获取控制器方法名
String className = joinPoint.getTarget().getClass().getName();
String methodName = signature.getName();
String controllerMethod = className + "." + methodName;
// 获取请求参数
String requestParams = getRequestParams(joinPoint);
// 获取返回参数
String responseParams = returnValue != null ? JSON.toJSONString(returnValue) : "无返回值";
// 合并为一条完整的日志信息
String logMessage = String.format(
"【请求日志】用户ID:[%s], IP地址:[%s], 请求方式:[%s], 控制器方法:[%s], 请求参数:[%s], 返回参数:[%s]", user_id, ipAddr, requestMethod, controllerMethod, requestParams, responseParams
);
// 打印单条完整日志
log.info(logMessage);
}
/**
* 抛出异常后通知(@AfterThrowing):方法抛出异常退出时执行的通知
* 注意在这里不能使用ProceedingJoinPoint
* 不然会报错ProceedingJoinPoint is only supported for around advice
* throwing注解为错误信息
*
* @param joinPoint
* @param ex
*/
@AfterThrowing(value = "pt()", throwing = "ex")
public void recordLog(JoinPoint joinPoint, Exception ex) {
try {
// 获取请求信息
HttpServletRequest request = HttpContextUtils.getHttpServletRequest();
// 获取请求方法类型(POST/GET等)
String requestMethod = request.getMethod();
// 获取用户ID
String authorization = request.getHeader("Authorization") + "";
Object user_id = "未登录"; // 请求用户的id
if (!authorization.equals("null")) {
user_id = JwtUtil.getUserIdFromToken(authorization);
}
// 获取请求IP
String ipAddr = IpUtils.getIpAddr(request);
// 获取方法签名信息
MethodSignature signature = (MethodSignature) joinPoint.getSignature();
Method method = signature.getMethod();
// 获取控制器方法名
String className = joinPoint.getTarget().getClass().getName();
String methodName = signature.getName();
String controllerMethod = className + "." + methodName;
// 获取请求参数
String requestParams = getRequestParams(joinPoint);
// 获取异常信息
String exceptionMsg = ex != null ? ex.getMessage() : "未知异常";
// 合并为一条完整的异常日志信息
String errorMessage = String.format(
"【异常日志】用户ID:[%s], IP地址:[%s], 请求方式:[%s], 控制器方法:[%s], 请求参数:[%s], 异常信息:[%s]", user_id, ipAddr, requestMethod, controllerMethod, requestParams, exceptionMsg
);
// 打印单条完整异常日志
log.error(errorMessage, ex);
} catch (Exception e) {
log.error("记录异常日志时出错: {}", e.getMessage());
}
}
/**
* 获取请求参数
*/
private String getRequestParams(JoinPoint joinPoint) {
try {
Object[] args = joinPoint.getArgs();
if (args.length == 0) {
return "无参数";
} else if (args[0] != null && args[0].toString().contains("SecurityContextHolderAwareRequestWrapper")) {
return JSON.toJSONString(Arrays.toString(ArrayUtil.remove(args, 0)));
} else {
// 检查是否只有一个参数且已经是JSON字符串格式
if (args.length == 1 && args[0] != null) {
// 如果参数本身就是字符串且是JSON格式,直接返回
if (args[0] instanceof String && ((String) args[0]).startsWith("{") && ((String) args[0]).endsWith("}")) {
return (String) args[0];
}
// 如果参数是普通对象,直接序列化
try {
return JSON.toJSONString(args[0]);
} catch (Exception e) {
// 如果序列化失败,再尝试使用参数名映射
Map<String, Object> map = new HashMap<>();
String[] names = ((CodeSignature) joinPoint.getSignature()).getParameterNames();
if (names != null) {
map.put(names[0], args[0]);
return JSON.toJSONString(map);
}
return JSON.toJSONString(args[0]);
}
} else {
// 多个参数时,使用参数名映射
Map<String, Object> map = new HashMap<>();
String[] names = ((CodeSignature) joinPoint.getSignature()).getParameterNames();
if (names != null) {
for (int i = 0; i < names.length; i++) {
map.put(names[i], args[i]);
}
}
return JSON.toJSONString(map);
}
}
} catch (Exception e) {
return "获取参数失败: " + e.getMessage();
}
}
}
@@ -0,0 +1,49 @@
package com.admin.common.aop;
import com.admin.common.annotation.RequireRole;
import com.admin.common.lang.R;
import com.admin.common.utils.JwtUtil;
import org.aspectj.lang.ProceedingJoinPoint;
import org.aspectj.lang.annotation.Around;
import org.aspectj.lang.annotation.Aspect;
import org.springframework.stereotype.Component;
import org.springframework.web.context.request.RequestContextHolder;
import org.springframework.web.context.request.ServletRequestAttributes;
import javax.servlet.http.HttpServletRequest;
/**
* 权限控制切面
* 处理 @RequireRole 注解,检查管理员权限(role_id = 0)
* 注意:JWT拦截器已经验证了token的有效性,这里只需要检查权限
*/
@Aspect
@Component
public class RoleAspect {
@Around("@annotation(requireRole)")
public Object checkRole(ProceedingJoinPoint joinPoint, RequireRole requireRole) throws Throwable {
// 获取当前请求
ServletRequestAttributes attributes = (ServletRequestAttributes) RequestContextHolder.getRequestAttributes();
if (attributes == null) {
return R.err(500, "无法获取请求信息");
}
HttpServletRequest request = attributes.getRequest();
String token = request.getHeader("Authorization");
// JWT拦截器已经验证过token存在且有效,这里直接获取role_id
Integer roleId = JwtUtil.getRoleIdFromToken(token);
if (roleId == null) {
return R.err(401, "无法获取用户权限信息");
}
// 检查是否为管理员(role_id = 0)
if (roleId != 0) {
return R.err(403, "权限不足,仅管理员可操作");
}
// 权限检查通过,执行原方法
return joinPoint.proceed();
}
}
@@ -0,0 +1,18 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
@Data
public class ChangePasswordDto {
@NotBlank(message = "当前密码不能为空")
private String currentPassword;
@NotBlank(message = "新密码不能为空")
private String newPassword;
@NotBlank(message = "确认密码不能为空")
private String confirmPassword;
}
@@ -0,0 +1,20 @@
package com.admin.common.dto;
import lombok.Data;
@Data
public class FlowDto {
// [{n=41_tcp, t=cc, u=73225, d=35043}, {n=41_tcp, t=conn, u=35043, d=73225}]
// 转发id_类型
private String n;
// 是请求还是接收
private String t;
// 上传流量
private Long u;
// 下载流量
private Long d;
}
@@ -0,0 +1,18 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
@Data
public class ForwardDto {
@NotBlank(message = "转发名称不能为空")
private String name;
@NotNull(message = "隧道ID不能为空")
private Integer tunnelId;
@NotBlank(message = "远程地址不能为空")
private String remoteAddr;
}
@@ -0,0 +1,24 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
@Data
public class ForwardUpdateDto {
@NotNull(message = "ID不能为空")
private Long id;
@NotNull(message = "用户ID不能为空")
private Integer userId;
@NotBlank(message = "转发名称不能为空")
private String name;
@NotNull(message = "隧道ID不能为空")
private Integer tunnelId;
@NotBlank(message = "远程地址不能为空")
private String remoteAddr;
}
@@ -0,0 +1,111 @@
package com.admin.common.dto;
import lombok.Data;
/**
* <p>
* 转发信息及关联隧道信息DTO
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Data
public class ForwardWithTunnelDto {
/**
* 转发记录ID
*/
private Long id;
/**
* 转发名称
*/
private String name;
/**
* 入口端口
*/
private Integer inPort;
/**
* 远程地址
*/
private String remoteAddr;
/**
* 转发状态
*/
private Integer status;
/**
* 创建时间
*/
private Long createdTime;
/**
* 更新时间
*/
private Long updatedTime;
// 以下为隧道相关字段
/**
* 隧道名称
*/
private String tunnelName;
/**
* 入口IP
*/
private String inIp;
private String userName;
/**
* 用户ID
*/
private Integer userId;
/**
* 隧道ID
*/
private Integer tunnelId;
/**
* 入站流量(字节)
*/
private Long inFlow;
/**
* 出站流量(字节)
*/
private Long outFlow;
// /**
// * 入口端口开始
// */
// private Integer inPortSta;
//
// /**
// * 入口端口结束
// */
// private Integer inPortEnd;
//
// /**
// * 出口IP
// */
// private String outIp;
//
// /**
// * 出口端口开始
// */
// private Integer outIpSta;
//
// /**
// * 出口端口结束
// */
// private Integer outIpEnd;
}
@@ -0,0 +1,10 @@
package com.admin.common.dto;
import lombok.Data;
@Data
public class GostDto {
private Integer code;
private String msg;
}
@@ -0,0 +1,17 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
@Data
public class LoginDto {
@NotBlank(message = "用户名不能为空")
private String username;
@NotBlank(message = "密码不能为空")
private String password;
}
@@ -0,0 +1,20 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Max;
import javax.validation.constraints.Min;
@Data
public class NodeDto {
@NotBlank(message = "节点名称不能为空")
private String name;
@NotNull(message = "控制端口不能为空")
@Min(value = 1, message = "端口号必须在1-65535之间")
@Max(value = 65535, message = "端口号必须在1-65535之间")
private Integer port;
}
@@ -0,0 +1,23 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Max;
import javax.validation.constraints.Min;
@Data
public class NodeUpdateDto {
@NotNull(message = "节点ID不能为空")
private Long id;
@NotBlank(message = "节点名称不能为空")
private String name;
@NotNull(message = "控制端口不能为空")
@Min(value = 1, message = "端口号必须在1-65535之间")
@Max(value = 65535, message = "端口号必须在1-65535之间")
private Integer port;
}
@@ -0,0 +1,22 @@
package com.admin.common.dto;
import lombok.Data;
@Data
public class PageDto {
/**
* 当前页码,默认为1
*/
private Long current = 1L;
/**
* 每页显示条数,默认为10
*/
private Long size = 10L;
/**
* 搜索关键字(可选)
*/
private String keyword;
}
@@ -0,0 +1,23 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Min;
@Data
public class SpeedLimitDto {
@NotBlank(message = "限速规则名称不能为空")
private String name;
@NotNull(message = "速度限制不能为空")
@Min(value = 1, message = "速度限制必须大于0")
private Integer speed;
@NotNull(message = "隧道ID不能为空")
private Long tunnelId;
@NotBlank(message = "隧道名称不能为空")
private String tunnelName;
}
@@ -0,0 +1,26 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Min;
@Data
public class SpeedLimitUpdateDto {
@NotNull(message = "ID不能为空")
private Long id;
@NotBlank(message = "限速规则名称不能为空")
private String name;
@NotNull(message = "速度限制不能为空")
@Min(value = 1, message = "速度限制必须大于0")
private Integer speed;
@NotNull(message = "隧道ID不能为空")
private Long tunnelId;
@NotBlank(message = "隧道名称不能为空")
private String tunnelName;
}
@@ -0,0 +1,57 @@
package com.admin.common.dto;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
/**
* 系统信息DTO
* 对应Go客户端上报的系统信息结构
*/
@Data
public class SystemInfoDto {
/**
* 主机IP地址
*/
@JsonProperty("host_ip")
private String hostIp;
/**
* 开机时间(秒)
*/
@JsonProperty("uptime")
private Long uptime;
/**
* 接收字节数
*/
@JsonProperty("bytes_received")
private Long bytesReceived;
/**
* 发送字节数
*/
@JsonProperty("bytes_transmitted")
private Long bytesTransmitted;
/**
* CPU使用率(百分比)
*/
@JsonProperty("cpu_usage")
private Double cpuUsage;
/**
* 内存使用率(百分比)
*/
@JsonProperty("memory_usage")
private Double memoryUsage;
/**
* 上报时间戳
*/
private Long timestamp;
public SystemInfoDto() {
this.timestamp = System.currentTimeMillis();
}
}
@@ -0,0 +1,42 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Min;
import javax.validation.constraints.Max;
@Data
public class TunnelDto {
@NotBlank(message = "隧道名称不能为空")
private String name;
@NotNull(message = "入口节点不能为空")
private Long inNodeId;
@NotNull(message = "入口端口开始不能为空")
@Min(value = 1, message = "入口端口开始必须大于0")
@Max(value = 65535, message = "入口端口开始不能超过65535")
private Integer inPortSta;
@NotNull(message = "入口端口结束不能为空")
@Min(value = 1, message = "入口端口结束必须大于0")
@Max(value = 65535, message = "入口端口结束不能超过65535")
private Integer inPortEnd;
// 出口节点ID,当type=1时可以为空,会自动设置为入口节点ID
private Long outNodeId;
// 出口端口开始,当type=1时可以为空,会自动设置为入口端口
private Integer outIpSta;
// 出口端口结束,当type=1时可以为空,会自动设置为入口端口
private Integer outIpEnd;
@NotNull(message = "隧道类型不能为空")
private Integer type;
@NotNull(message = "流量计算类型不能为空")
private Integer flow;
}
@@ -0,0 +1,13 @@
package com.admin.common.dto;
import lombok.Data;
@Data
public class TunnelListDto {
private Integer id;
private String name;
}
@@ -0,0 +1,36 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Min;
@Data
public class UserDto {
@NotBlank(message = "姓名不能为空")
private String name;
@NotBlank(message = "用户名不能为空")
private String user;
@NotBlank(message = "密码不能为空")
private String pwd;
@NotNull(message = "流量不能为空")
@Min(value = 0, message = "流量不能小于0")
private Long flow;
@NotNull(message = "转发数量不能为空")
@Min(value = 0, message = "转发数量不能小于0")
private Integer num;
@NotNull(message = "过期时间不能为空")
private Long expTime;
@NotNull(message = "流量重置时间不能为空")
private Long flowResetTime;
private Integer status;
}
@@ -0,0 +1,85 @@
package com.admin.common.dto;
import lombok.Data;
import java.util.List;
/**
* 用户套餐信息DTO
*/
@Data
public class UserPackageDto {
/**
* 用户基本信息
*/
private UserInfoDto userInfo;
/**
* 用户隧道权限列表
*/
private List<UserTunnelDetailDto> tunnelPermissions;
/**
* 用户转发列表
*/
private List<UserForwardDetailDto> forwards;
/**
* 用户基本信息
*/
@Data
public static class UserInfoDto {
private Long id;
private String name;
private String user;
private Integer status;
private Long flow; // 总流量配额(GB)
private Long inFlow; // 已用入站流量(字节)
private Long outFlow; // 已用出站流量(字节)
private Integer num; // 转发数量配额
private Long expTime; // 过期时间
private Long flowResetTime; // 流量重置时间
private Long createdTime;
private Long updatedTime;
}
/**
* 用户隧道权限详情
*/
@Data
public static class UserTunnelDetailDto {
private Integer id;
private Integer userId;
private Integer tunnelId;
private String tunnelName;
private Integer tunnelFlow; // 隧道流量计算类型(1-单向,2-双向)
private Long flow; // 隧道流量配额(GB)
private Long inFlow; // 隧道已用入站流量(字节)
private Long outFlow; // 隧道已用出站流量(字节)
private Integer num; // 隧道转发数量配额
private Long flowResetTime; // 流量重置时间
private Long expTime; // 隧道权限过期时间
private Integer speedId;
private String speedLimitName;
private Integer speed;
}
/**
* 用户转发详情
*/
@Data
public static class UserForwardDetailDto {
private Long id;
private String name;
private Integer tunnelId;
private String tunnelName;
private String inIp;
private Integer inPort;
private String remoteAddr;
private Long inFlow; // 转发入站流量(字节)
private Long outFlow; // 转发出站流量(字节)
private Integer status;
private Long createdTime;
}
}
@@ -0,0 +1,40 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Min;
@Data
public class UserTunnelDto {
@NotNull(message = "用户ID不能为空")
private Integer userId;
@NotNull(message = "隧道ID不能为空")
private Integer tunnelId;
@NotNull(message = "流量限制不能为空")
@Min(value = 0, message = "流量限制不能小于0")
private Long flow;
@NotNull(message = "转发数量不能为空")
@Min(value = 0, message = "转发数量不能小于0")
private Integer num;
/**
* 流量重置时间(时间戳)
*/
@NotNull(message = "流量重置时间不能为空")
private Long flowResetTime;
/**
* 到期时间(时间戳)
*/
@NotNull(message = "到期时间不能为空")
private Long expTime;
/**
* 限速规则ID(可选,null表示不限速)
*/
private Integer speedId;
}
@@ -0,0 +1,13 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotNull;
@Data
public class UserTunnelQueryDto {
@NotNull
private Integer userId;
}
@@ -0,0 +1,37 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Min;
@Data
public class UserTunnelUpdateDto {
@NotNull(message = "用户隧道权限ID不能为空")
private Integer id;
@NotNull(message = "流量限制不能为空")
@Min(value = 0, message = "流量限制不能小于0")
private Long flow;
@NotNull(message = "转发数量不能为空")
@Min(value = 0, message = "转发数量不能小于0")
private Integer num;
/**
* 流量重置时间(时间戳)
*/
@NotNull(message = "流量重置时间不能为空")
private Long flowResetTime;
/**
* 到期时间(时间戳)
*/
@NotNull(message = "到期时间不能为空")
private Long expTime;
/**
* 限速规则ID(可选,null表示不限速)
*/
private Integer speedId;
}
@@ -0,0 +1,120 @@
package com.admin.common.dto;
import lombok.Data;
/**
* <p>
* 用户隧道权限及隧道详细信息DTO
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Data
public class UserTunnelWithDetailDto {
/**
* 用户隧道权限记录ID
*/
private Integer id;
/**
* 用户ID
*/
private Integer userId;
/**
* 隧道ID
*/
private Integer tunnelId;
/**
* 流量限制
*/
private Integer flow;
/**
* 转发数量限制
*/
private Integer num;
/**
* 流量重置时间(时间戳)
*/
private Long flowResetTime;
/**
* 到期时间(时间戳)
*/
private Long expTime;
/**
* 限速规则ID
*/
private Integer speedId;
/**
* 限速规则名称
*/
private String speedLimitName;
/**
* 限速值
*/
private Integer speed;
/**
* 隧道名称
*/
private String tunnelName;
/**
* 隧道流量计算类型(1-单向,2-双向)
*/
private Integer tunnelFlow;
/**
* 入站流量(字节)
*/
private Long inFlow;
/**
* 出站流量(字节)
*/
private Long outFlow;
// /**
// * 入口IP
// */
// private String inIp;
//
// /**
// * 入口端口开始
// */
// private Integer inPortSta;
//
// /**
// * 入口端口结束
// */
// private Integer inPortEnd;
//
// /**
// * 出口IP
// */
// private String outIp;
//
// /**
// * 出口端口开始
// */
// private Integer outIpSta;
//
// /**
// * 出口端口结束
// */
// private Integer outIpEnd;
//
// /**
// * 隧道类型(1-端口转发,2-隧道转发)
// */
// private Integer type;
}
@@ -0,0 +1,38 @@
package com.admin.common.dto;
import lombok.Data;
import javax.validation.constraints.NotBlank;
import javax.validation.constraints.NotNull;
import javax.validation.constraints.Min;
@Data
public class UserUpdateDto {
@NotNull(message = "用户ID不能为空")
private Long id;
@NotBlank(message = "姓名不能为空")
private String name;
@NotBlank(message = "用户名不能为空")
private String user;
private String pwd; // 更新时密码可选
@NotNull(message = "流量不能为空")
@Min(value = 0, message = "流量不能小于0")
private Long flow;
@NotNull(message = "转发数量不能为空")
@Min(value = 0, message = "转发数量不能小于0")
private Integer num;
@NotNull(message = "过期时间不能为空")
private Long expTime;
@NotNull(message = "流量重置时间不能为空")
private Long flowResetTime;
private Integer status;
}
@@ -0,0 +1,40 @@
package com.admin.common.exception;
import com.admin.common.lang.R;
import lombok.extern.slf4j.Slf4j;
import org.springframework.validation.BindingResult;
import org.springframework.validation.ObjectError;
import org.springframework.web.bind.MethodArgumentNotValidException;
import org.springframework.web.bind.annotation.ExceptionHandler;
import org.springframework.web.bind.annotation.RestControllerAdvice;
@Slf4j
@RestControllerAdvice
public class GlobalExceptionHandler {
//
// 实体校验异常捕获
//@ResponseStatus(HttpStatus.BAD_REQUEST)
@ExceptionHandler(value = MethodArgumentNotValidException.class)
public R MethodArgumentNotValidException(MethodArgumentNotValidException e) {
BindingResult result = e.getBindingResult();
ObjectError objectError = result.getAllErrors().stream().findFirst().get();
log.error("实体校验异常:----------------{}", objectError.getDefaultMessage());
return R.err(500, objectError.getDefaultMessage());
}
// 未授权异常捕获
@ExceptionHandler(value = UnauthorizedException.class)
public R handleUnauthorizedException(UnauthorizedException e) {
log.error("未授权异常:----------------{}", e.getMessage());
return R.err(401, e.getMessage());
}
@ExceptionHandler(value = Exception.class)
public R Exception(Exception e){
log.error("异常:----------------{}", e.getMessage());
return R.err(-2, "异常错误");
}
}
@@ -0,0 +1,18 @@
package com.admin.common.exception;
import org.springframework.http.client.ClientHttpResponse;
import org.springframework.web.client.ResponseErrorHandler;
import java.io.IOException;
public class HttpErrorHandler implements ResponseErrorHandler {
@Override
public boolean hasError(ClientHttpResponse response) throws IOException {
return false;
}
@Override
public void handleError(ClientHttpResponse clientHttpResponse) throws IOException {
}
}
@@ -0,0 +1,15 @@
package com.admin.common.exception;
/**
* 未授权异常类
*/
public class UnauthorizedException extends RuntimeException {
public UnauthorizedException(String message) {
super(message);
}
public UnauthorizedException(String message, Throwable cause) {
super(message, cause);
}
}
@@ -0,0 +1,33 @@
package com.admin.common.interceptor;
import com.admin.common.exception.UnauthorizedException;
import com.admin.common.utils.JwtUtil;
import org.springframework.util.StringUtils;
import org.springframework.web.servlet.HandlerInterceptor;
import javax.servlet.http.HttpServletRequest;
import javax.servlet.http.HttpServletResponse;
/**
* JWT拦截器,验证用户是否登录
*/
public class JwtInterceptor implements HandlerInterceptor {
@Override
public boolean preHandle(HttpServletRequest request, HttpServletResponse response, Object handler) {
String token = request.getHeader("Authorization");
if (!StringUtils.hasText(token)) {
throw new UnauthorizedException("未登录或token已过期");
}
if (!JwtUtil.validateToken(token)) {
throw new UnauthorizedException("无效的token或token已过期");
}
return true;
}
}
@@ -0,0 +1,49 @@
package com.admin.common.lang;
import lombok.Data;
@Data
public class R {
private int code = 0;
private String msg = "操作成功";
private long ts = System.currentTimeMillis();
private Object data;
public static R ok(Object data){
R m = new R();
m.setData(data);
return m;
}
public static R ok(){
return new R();
}
public static R err(int code, String msg){
R m = new R();
m.setCode(code);
m.setMsg(msg);
return m;
}
public static R err(String msg){
R m = new R();
m.setCode(-1);
m.setMsg(msg);
return m;
}
public static R err(){
R m = new R();
m.setCode(-1);
m.setMsg("请求失败");
return m;
}
}
@@ -0,0 +1,559 @@
package com.admin.common.task;
import com.admin.mapper.UserMapper;
import com.admin.mapper.ForwardMapper;
import com.admin.mapper.UserTunnelMapper;
import com.admin.mapper.TunnelMapper;
import com.admin.mapper.NodeMapper;
import com.admin.service.UserService;
import com.admin.service.UserTunnelService;
import com.admin.entity.Forward;
import com.admin.entity.Tunnel;
import com.admin.entity.Node;
import com.admin.entity.UserTunnel;
import com.admin.common.dto.GostDto;
import com.admin.common.utils.GostUtil;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import lombok.extern.slf4j.Slf4j;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.CommandLineRunner;
import org.springframework.stereotype.Component;
import javax.annotation.Resource;
import java.util.List;
import java.util.concurrent.DelayQueue;
import java.util.concurrent.Executors;
@Component
@Slf4j
public class DelayQueueManager implements CommandLineRunner {
private final DelayQueue<DelayTask> delayQueue = new DelayQueue<>();
@Autowired
UserMapper userMapper;
@Resource
ForwardMapper forwardMapper;
@Resource
UserTunnelMapper userTunnelMapper;
@Resource
TunnelMapper tunnelMapper;
@Resource
NodeMapper nodeMapper;
/**
* 加入到延时队列中
*
* @param task
*/
public void put(DelayTask task) {
log.info("加入延时任务:{}", task);
delayQueue.put(task);
}
/**
* 取消延时任务
*
* @param task
* @return
*/
public boolean remove(DelayTask task) {
log.info("取消延时任务:{}", task);
return delayQueue.remove(task);
}
/**
* 取消延时任务
*
* @param taskid
* @return
*/
public boolean remove(String taskid) {
return remove(new DelayTask(new TaskBase(taskid), 0));
}
/**
* 取消延时任务
*
* @param taskid
* @return
*/
public boolean remove_a(String taskid) {
return remove(new DelayTask(new TaskBase(taskid), 0));
}
@Override
public void run(String... args) throws Exception {
log.info("初始化延时队列");
Executors.newSingleThreadExecutor().execute(new Thread(this::excuteThread));
// 初始化用户账号到期延时任务
initUserExpirationTasks();
// 初始化用户隧道到期延时任务
initUserTunnelExpirationTasks();
}
/**
* 延时任务执行线程
*/
private void excuteThread() {
while (true) {
try {
DelayTask task = delayQueue.take();
processTask(task);
} catch (InterruptedException e) {
break;
}
}
}
/**
* 内部执行延时任务
*
* @param task
*/
private void processTask(DelayTask task) {
log.info("执行延时任务:{}", task.getData().toString());
TaskBase data = task.getData();
switch (data.getType()){
case "1": // 账号到期延迟任务
handleUserExpiration(data.getData());
break;
case "2": // 隧道到期延迟任务
handleUserTunnelExpiration(data.getData());
break;
default:
log.error("未知延时任务类型:{}", data.getType());
break;
}
}
/**
* 处理用户账号到期
*
* @param userId 用户ID
*/
private void handleUserExpiration(String userId) {
try {
log.info("处理用户账号到期,用户ID:{}", userId);
Long userIdLong = Long.parseLong(userId);
// 获取用户信息
com.admin.entity.User user = userMapper.selectById(userIdLong);
if (user == null) {
log.warn("用户不存在,用户ID:{}", userId);
return;
}
// 检查用户是否确实已过期
if (user.getExpTime() != null && user.getExpTime() > System.currentTimeMillis()) {
log.info("用户未过期,无需处理,用户ID:{},过期时间:{}", userId, user.getExpTime());
return;
}
// 禁用用户账号
user.setStatus(0); // 设置为禁用状态
user.setUpdatedTime(System.currentTimeMillis());
int i = userMapper.updateById(user);
if (i != 0) {
log.info("用户账号已禁用,用户ID:{}", userId);
// 清理用户相关的活跃连接和服务
cleanupUserServices(userIdLong);
} else {
log.error("禁用用户账号失败,用户ID:{}", userId);
}
} catch (Exception e) {
log.error("处理用户账号到期异常,用户ID:{},错误:{}", userId, e.getMessage(), e);
}
}
/**
* 清理用户相关服务
*
* @param userId 用户ID
*/
private void cleanupUserServices(Long userId) {
try {
log.info("暂停用户相关转发服务,用户ID:{}", userId);
// 获取用户的所有转发
QueryWrapper<Forward> forwardQuery = new QueryWrapper<>();
forwardQuery.eq("user_id", userId);
List<Forward> userForwards = forwardMapper.selectList(forwardQuery);
log.info("找到用户转发数量:{},用户ID:{}", userForwards.size(), userId);
for (Forward forward : userForwards) {
try {
// 暂停转发服务
pauseForwardService(forward, userId);
} catch (Exception e) {
log.error("暂停转发服务失败,转发ID:{},用户ID:{},错误:{}", forward.getId(), userId, e.getMessage());
}
}
} catch (Exception e) {
log.error("清理用户服务失败,用户ID:{},错误:{}", userId, e.getMessage(), e);
}
}
/**
* 暂停转发服务
*
* @param forward 转发对象
* @param userId 用户ID
*/
private void pauseForwardService(Forward forward, Long userId) {
try {
Tunnel tunnel = tunnelMapper.selectById(forward.getTunnelId());
if (tunnel == null) {
log.warn("隧道不存在,跳过暂停,转发ID:{},隧道ID:{}", forward.getId(), forward.getTunnelId());
return;
}
Node inNode = nodeMapper.selectById(tunnel.getInNodeId());
if (inNode == null) {
log.warn("入口节点不存在,跳过暂停,转发ID:{},节点ID:{}", forward.getId(), tunnel.getInNodeId());
return;
}
// 获取用户隧道关系
UserTunnel userTunnel = getUserTunnelRelation(userId, tunnel.getId());
if (userTunnel == null) {
log.warn("用户隧道关系不存在,跳过暂停,用户ID:{},隧道ID:{}", userId, tunnel.getId());
return;
}
String serviceName = buildServiceName(forward.getId(), userId, userTunnel.getId());
String nodeAddress = buildNodeAddress(inNode);
// 暂停主服务
GostDto result = GostUtil.PauseService(nodeAddress, serviceName, inNode.getSecret());
// 隧道转发需要同时暂停远端服务
if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD
Node outNode = nodeMapper.selectById(tunnel.getOutNodeId());
if (outNode != null) {
String outNodeAddress = buildNodeAddress(outNode);
GostDto remoteResult = GostUtil.PauseRemoteService(outNodeAddress, serviceName, outNode.getSecret());
if (!"OK".equals(remoteResult.getMsg())) {
log.warn("暂停远端服务失败,转发ID:{},用户ID:{},服务名:{},结果:{}",
forward.getId(), userId, serviceName, remoteResult.getMsg());
}
}
}
if ( "OK".equals(result.getMsg())) {
forward.setStatus(0);
forwardMapper.updateById(forward);
log.info("成功暂停转发服务,转发ID:{},用户ID:{},服务名:{}", forward.getId(), userId, serviceName);
} else {
log.warn("暂停转发服务失败,转发ID:{},用户ID:{},服务名:{},结果:{}",
forward.getId(), userId, serviceName, result.getMsg());
}
} catch (Exception e) {
log.error("暂停转发服务异常,转发ID:{},用户ID:{},错误:{}", forward.getId(), userId, e.getMessage(), e);
}
}
/**
* 获取用户隧道关系
*
* @param userId 用户ID
* @param tunnelId 隧道ID
* @return 用户隧道关系对象
*/
private UserTunnel getUserTunnelRelation(Long userId, Long tunnelId) {
try {
QueryWrapper<UserTunnel> query = new QueryWrapper<>();
query.eq("user_id", userId).eq("tunnel_id", tunnelId);
return userTunnelMapper.selectOne(query);
} catch (Exception e) {
log.error("获取用户隧道关系失败,用户ID:{},隧道ID:{},错误:{}", userId, tunnelId, e.getMessage());
return null;
}
}
/**
* 构建服务名称
*
* @param forwardId 转发ID
* @param userId 用户ID
* @param userTunnelId 用户隧道ID
* @return 服务名称
*/
private String buildServiceName(Long forwardId, Long userId, Integer userTunnelId) {
return forwardId + "_" + userId + "_" + userTunnelId;
}
/**
* 构建节点地址
*
* @param node 节点对象
* @return 节点地址字符串
*/
private String buildNodeAddress(Node node) {
return node.getIp() + ":" + node.getPort();
}
/**
* 初始化用户账号到期延时任务
* 查询所有非管理员的正常用户,为有到期时间且未过期的用户创建延时任务
*/
private void initUserExpirationTasks() {
try {
log.info("开始初始化用户账号到期延时任务");
QueryWrapper<com.admin.entity.User> userQuery = new QueryWrapper<>();
userQuery.ne("role_id", 0) // 排除管理员用户
.eq("status", 1) // 只查询启用状态的用户
.isNotNull("exp_time") // 只查询有到期时间的用户
.orderBy(true, true, "exp_time"); // 按到期时间排序
List<com.admin.entity.User> users = userMapper.selectList(userQuery);
for (com.admin.entity.User user : users) {
scheduleUserExpirationTask(user);
}
} catch (Exception e) {
log.error("初始化用户账号到期延时任务失败:{}", e.getMessage(), e);
}
}
/**
* 安排用户到期延时任务
*
* @param user 用户对象
*/
private void scheduleUserExpirationTask(com.admin.entity.User user) {
try {
if (user.getExpTime() != null && user.getExpTime() > System.currentTimeMillis()) {
// 创建延时任务
TaskBase taskBase = new TaskBase(user.getId().toString());
taskBase.setType("1"); // 账号到期延迟任务
long delayTime = user.getExpTime() - System.currentTimeMillis();
DelayTask delayTask = new DelayTask(taskBase, delayTime);
put(delayTask);
log.debug("已添加用户到期延时任务,用户ID:{},到期时间:{},剩余时间:{}ms",
user.getId(), user.getExpTime(), delayTime);
}
} catch (Exception e) {
log.error("添加用户到期延时任务失败,用户ID:{},错误:{}", user.getId(), e.getMessage(), e);
}
}
/**
* 处理用户隧道到期
*
* @param userTunnelId 用户隧道ID
*/
private void handleUserTunnelExpiration(String userTunnelId) {
try {
log.info("处理用户隧道到期,用户隧道ID:{}", userTunnelId);
Integer userTunnelIdInt = Integer.parseInt(userTunnelId);
// 获取用户隧道信息
UserTunnel userTunnel = userTunnelMapper.selectById(userTunnelIdInt);
if (userTunnel == null) {
log.warn("用户隧道不存在,用户隧道ID:{}", userTunnelId);
return;
}
// 检查用户隧道是否确实已过期
if (userTunnel.getExpTime() != null && userTunnel.getExpTime() > System.currentTimeMillis()) {
log.info("用户隧道未过期,无需处理,用户隧道ID:{},过期时间:{}", userTunnelId, userTunnel.getExpTime());
return;
}
log.info("用户隧道已过期,开始处理,用户隧道ID:{},用户ID:{},隧道ID:{}",
userTunnelId, userTunnel.getUserId(), userTunnel.getTunnelId());
// 暂停该用户在该隧道上的所有转发服务
cleanupUserTunnelServices(userTunnel);
// 禁用过期的用户隧道权限(设置status为0)
userTunnel.setStatus(0);
int updateResult = userTunnelMapper.updateById(userTunnel);
if (updateResult > 0) {
log.info("已禁用过期的用户隧道权限,用户隧道ID:{}", userTunnelId);
} else {
log.error("禁用过期用户隧道权限失败,用户隧道ID:{}", userTunnelId);
}
} catch (Exception e) {
log.error("处理用户隧道到期异常,用户隧道ID:{},错误:{}", userTunnelId, e.getMessage(), e);
}
}
/**
* 清理用户隧道相关服务
*
* @param userTunnel 用户隧道对象
*/
private void cleanupUserTunnelServices(UserTunnel userTunnel) {
try {
log.info("暂停用户隧道相关转发服务,用户ID:{},隧道ID:{}", userTunnel.getUserId(), userTunnel.getTunnelId());
// 获取该用户在该隧道上的所有转发
QueryWrapper<Forward> forwardQuery = new QueryWrapper<>();
forwardQuery.eq("user_id", userTunnel.getUserId())
.eq("tunnel_id", userTunnel.getTunnelId());
List<Forward> userTunnelForwards = forwardMapper.selectList(forwardQuery);
log.info("找到用户隧道转发数量:{},用户ID:{},隧道ID:{}",
userTunnelForwards.size(), userTunnel.getUserId(), userTunnel.getTunnelId());
for (Forward forward : userTunnelForwards) {
try {
// 暂停转发服务
pauseUserTunnelForwardService(forward, userTunnel);
} catch (Exception e) {
log.error("暂停用户隧道转发服务失败,转发ID:{},用户ID:{},隧道ID:{},错误:{}",
forward.getId(), userTunnel.getUserId(), userTunnel.getTunnelId(), e.getMessage());
}
}
} catch (Exception e) {
log.error("清理用户隧道服务失败,用户ID:{},隧道ID:{},错误:{}",
userTunnel.getUserId(), userTunnel.getTunnelId(), e.getMessage(), e);
}
}
/**
* 暂停用户隧道转发服务
*
* @param forward 转发对象
* @param userTunnel 用户隧道对象
*/
private void pauseUserTunnelForwardService(Forward forward, UserTunnel userTunnel) {
try {
Tunnel tunnel = tunnelMapper.selectById(forward.getTunnelId());
if (tunnel == null) {
log.warn("隧道不存在,跳过暂停,转发ID:{},隧道ID:{}", forward.getId(), forward.getTunnelId());
return;
}
Node inNode = nodeMapper.selectById(tunnel.getInNodeId());
if (inNode == null) {
log.warn("入口节点不存在,跳过暂停,转发ID:{},节点ID:{}", forward.getId(), tunnel.getInNodeId());
return;
}
String serviceName = buildServiceName(forward.getId(), Long.valueOf(userTunnel.getUserId()), userTunnel.getId());
String nodeAddress = buildNodeAddress(inNode);
// 暂停服务
GostDto result = GostUtil.PauseService(nodeAddress, serviceName, inNode.getSecret());
// 隧道转发需要同时暂停远端服务
if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD
Node outNode = nodeMapper.selectById(tunnel.getOutNodeId());
if (outNode != null) {
String outNodeAddress = buildNodeAddress(outNode);
GostDto remoteResult = GostUtil.PauseRemoteService(outNodeAddress, serviceName, outNode.getSecret());
if (!"OK".equals(remoteResult.getMsg())) {
log.warn("暂停远端服务失败,转发ID:{},用户ID:{},隧道ID:{},服务名:{},结果:{}",
forward.getId(), userTunnel.getUserId(), userTunnel.getTunnelId(), serviceName, remoteResult.getMsg());
}
}
}
if ("OK".equals(result.getMsg())) {
forward.setStatus(0);
forwardMapper.updateById(forward);
log.info("成功暂停用户隧道转发服务,转发ID:{},用户ID:{},隧道ID:{},服务名:{}",
forward.getId(), userTunnel.getUserId(), userTunnel.getTunnelId(), serviceName);
} else {
log.warn("暂停用户隧道转发服务失败,转发ID:{},用户ID:{},隧道ID:{},服务名:{},结果:{}",
forward.getId(), userTunnel.getUserId(), userTunnel.getTunnelId(), serviceName, result.getMsg());
}
} catch (Exception e) {
log.error("暂停用户隧道转发服务异常,转发ID:{},用户ID:{},隧道ID:{},错误:{}",
forward.getId(), userTunnel.getUserId(), userTunnel.getTunnelId(), e.getMessage(), e);
}
}
/**
* 初始化用户隧道到期延时任务
* 查询所有有到期时间且未过期的用户隧道权限,为其创建延时任务
*/
private void initUserTunnelExpirationTasks() {
try {
log.info("开始初始化用户隧道到期延时任务");
QueryWrapper<UserTunnel> userTunnelQuery = new QueryWrapper<>();
userTunnelQuery.eq("status", 1); // 按到期时间排序
List<UserTunnel> userTunnels = userTunnelMapper.selectList(userTunnelQuery);
int taskCount = 0;
for (UserTunnel userTunnel : userTunnels) {
scheduleUserTunnelExpirationTask(userTunnel);
}
log.info("完成初始化用户隧道到期延时任务,总计查询:{},添加任务:{}", userTunnels.size(), taskCount);
} catch (Exception e) {
log.error("初始化用户隧道到期延时任务失败:{}", e.getMessage(), e);
}
}
/**
* 安排用户隧道到期延时任务
*
* @param userTunnel 用户隧道对象
*/
private void scheduleUserTunnelExpirationTask(UserTunnel userTunnel) {
// 创建延时任务
TaskBase taskBase = new TaskBase(userTunnel.getId().toString());
taskBase.setType("2"); // 隧道到期延迟任务
long delayTime = userTunnel.getExpTime() - System.currentTimeMillis();
DelayTask delayTask = new DelayTask(taskBase, delayTime);
put(delayTask);
log.debug("已添加用户隧道到期延时任务,用户隧道ID:{},用户ID:{},隧道ID:{},到期时间:{},剩余时间:{}ms",
userTunnel.getId(), userTunnel.getUserId(), userTunnel.getTunnelId(),
userTunnel.getExpTime(), delayTime);
}
/**
* 添加用户隧道到期延时任务(公共方法,供其他服务调用)
*
* @param userTunnel 用户隧道对象
*/
public void addUserTunnelExpirationTask(UserTunnel userTunnel) {
scheduleUserTunnelExpirationTask(userTunnel);
}
/**
* 移除用户隧道到期延时任务(公共方法,供其他服务调用)
*
* @param userTunnelId 用户隧道ID
*/
public void removeUserTunnelExpirationTask(Integer userTunnelId) {
String taskId = userTunnelId.toString();
boolean removed = remove(taskId);
if (removed) {
log.info("已移除用户隧道到期延时任务,用户隧道ID:{}", userTunnelId);
} else {
log.debug("未找到需要移除的用户隧道到期延时任务,用户隧道ID:{}", userTunnelId);
}
}
}
@@ -0,0 +1,57 @@
package com.admin.common.task;
import java.util.concurrent.Delayed;
import java.util.concurrent.TimeUnit;
/**
* 延时任务
*/
public class DelayTask implements Delayed {
//任务参数
final private TaskBase data;
//任务的延时时间,单位毫秒
final private long expire;
/**
* 构造延时任务
*
* @param data 业务数据
* @param expire 任务延时时间(ms)
*/
public DelayTask(TaskBase data, long expire) {
super();
this.data = data;
this.expire = expire + System.currentTimeMillis();
}
public TaskBase getData() {
return data;
}
public long getExpire() {
return expire;
}
@Override
public boolean equals(Object obj) {
if (obj instanceof DelayTask) {
return this.data.getData().equals(((DelayTask) obj).getData().getData());
}
return false;
}
@Override
public String toString() {
return "{" + "data:" + data.toString() + "," + "延时时间:"+expire+"}";
}
@Override
public long getDelay(TimeUnit unit) {
return unit.convert(this.expire - System.currentTimeMillis(), unit);
}
@Override
public int compareTo(Delayed o) {
long delta = getDelay(TimeUnit.NANOSECONDS) - o.getDelay(TimeUnit.NANOSECONDS);
return (int) delta;
}
}
@@ -0,0 +1,164 @@
package com.admin.common.task;
import com.admin.entity.User;
import com.admin.entity.UserTunnel;
import com.admin.service.UserService;
import com.admin.service.UserTunnelService;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper;
import lombok.extern.slf4j.Slf4j;
import org.springframework.context.annotation.Configuration;
import org.springframework.scheduling.annotation.EnableScheduling;
import org.springframework.scheduling.annotation.Scheduled;
import javax.annotation.Resource;
import java.time.LocalDate;
import java.util.List;
@Slf4j
@Configuration
@EnableScheduling
public class ResetFlowAsync {
@Resource
UserService userService;
@Resource
UserTunnelService userTunnelService;
/**
* 每天0点执行流量重置任务
* 查询出用户和隧道的重置流量日期是今天的数据,将上下流量重置为0
* 考虑当月是29、30天,但是选择是31的这种边界情况
*
* 并发安全说明:
* - 使用setSql()进行原子SQL更新,只更新流量字段(in_flow, out_flow)
* - 不会影响DelayQueueManager的到期任务对status等其他字段的更新
* - 避免了并发修改导致的数据覆盖问题
*/
@Scheduled(cron = "0 0 0 * * ?")
public void reset_flow(){
log.info("开始执行流量重置任务");
try {
// 获取当前日期信息
LocalDate today = LocalDate.now();
int currentDay = today.getDayOfMonth(); // 当前是几号
int lastDayOfMonth = today.lengthOfMonth(); // 当月最后一天
log.info("当前日期: {}, 当月第{}天, 当月最后一天: {}", today, currentDay, lastDayOfMonth);
// 重置用户流量
resetUserFlow(currentDay, lastDayOfMonth);
// 重置用户隧道流量
resetUserTunnelFlow(currentDay, lastDayOfMonth);
log.info("流量重置任务执行完成");
} catch (Exception e) {
log.error("流量重置任务执行失败", e);
}
}
/**
* 重置用户流量
* @param currentDay 当前日期(几号)
* @param lastDayOfMonth 当月最后一天
*/
private void resetUserFlow(int currentDay, int lastDayOfMonth) {
try {
// flowResetTime字段存储的是1-31的数字,表示每月第几号重置
// 构建查询条件:重置日期等于今天,或者重置日期大于当月最大天数且今天是月末
QueryWrapper<User> queryWrapper = new QueryWrapper<>();
if (currentDay == lastDayOfMonth) {
// 如果今天是月末,查询重置日期等于今天或者大于当月最大天数的记录
// 例如:当月30天,但用户设置31号重置,则在30号执行重置
queryWrapper.and(wrapper -> wrapper.eq("flow_reset_time", currentDay)
.or().gt("flow_reset_time", lastDayOfMonth));
} else {
// 否则只查询重置日期等于今天的记录
queryWrapper.eq("flow_reset_time", currentDay);
}
// 查询需要重置的用户
List<User> usersToReset = userService.list(queryWrapper);
if (usersToReset.isEmpty()) {
log.info("没有需要重置流量的用户");
return;
}
log.info("找到{}个需要重置流量的用户", usersToReset.size());
// 批量重置用户流量 - 使用SQL原子操作避免与到期任务的并发冲突
for (User user : usersToReset) {
UpdateWrapper<User> updateWrapper = new UpdateWrapper<>();
updateWrapper.eq("id", user.getId())
.setSql("in_flow = 0, out_flow = 0"); // 使用SQL原子操作,只更新流量字段
boolean success = userService.update(null, updateWrapper);
if (success) {
log.info("用户[ID: {}, 用户名: {}]流量重置成功,重置日期: 每月{}号",
user.getId(), user.getUser(), user.getFlowResetTime());
} else {
log.error("用户[ID: {}, 用户名: {}]流量重置失败", user.getId(), user.getUser());
}
}
} catch (Exception e) {
log.error("重置用户流量失败", e);
}
}
/**
* 重置用户隧道流量
* @param currentDay 当前日期(几号)
* @param lastDayOfMonth 当月最后一天
*/
private void resetUserTunnelFlow(int currentDay, int lastDayOfMonth) {
try {
// flowResetTime字段存储的是1-31的数字,表示每月第几号重置
// 构建查询条件:重置日期等于今天,或者重置日期大于当月最大天数且今天是月末
QueryWrapper<UserTunnel> queryWrapper = new QueryWrapper<>();
if (currentDay == lastDayOfMonth) {
// 如果今天是月末,查询重置日期等于今天或者大于当月最大天数的记录
// 例如:当月30天,但用户设置31号重置,则在30号执行重置
queryWrapper.and(wrapper -> wrapper.eq("flow_reset_time", currentDay)
.or().gt("flow_reset_time", lastDayOfMonth));
} else {
// 否则只查询重置日期等于今天的记录
queryWrapper.eq("flow_reset_time", currentDay);
}
// 查询需要重置的用户隧道
List<UserTunnel> userTunnelsToReset = userTunnelService.list(queryWrapper);
if (userTunnelsToReset.isEmpty()) {
log.info("没有需要重置流量的用户隧道");
return;
}
log.info("找到{}个需要重置流量的用户隧道", userTunnelsToReset.size());
// 批量重置用户隧道流量 - 使用SQL原子操作避免与到期任务的并发冲突
for (UserTunnel userTunnel : userTunnelsToReset) {
UpdateWrapper<UserTunnel> updateWrapper = new UpdateWrapper<>();
updateWrapper.eq("id", userTunnel.getId())
.setSql("in_flow = 0, out_flow = 0"); // 使用SQL原子操作,只更新流量字段
boolean success = userTunnelService.update(null, updateWrapper);
if (success) {
log.info("用户隧道[ID: {}, 用户ID: {}, 隧道ID: {}]流量重置成功,重置日期: 每月{}号",
userTunnel.getId(), userTunnel.getUserId(), userTunnel.getTunnelId(), userTunnel.getFlowResetTime());
} else {
log.error("用户隧道[ID: {}, 用户ID: {}, 隧道ID: {}]流量重置失败",
userTunnel.getId(), userTunnel.getUserId(), userTunnel.getTunnelId());
}
}
} catch (Exception e) {
log.error("重置用户隧道流量失败", e);
}
}
}
@@ -0,0 +1,14 @@
package com.admin.common.task;
import lombok.Data;
@Data
public class TaskBase {
private String data;
private String type;
public TaskBase(String data) {
this.data = data;
}
}
@@ -0,0 +1,378 @@
package com.admin.common.utils;
import com.admin.common.dto.GostDto;
import com.alibaba.fastjson.JSONArray;
import com.alibaba.fastjson.JSONObject;
public class GostUtil {
private static final String API_BASE_URL = "/api/config/";
private static final String LIMITERS_ENDPOINT = "limiters";
private static final String SERVICES_ENDPOINT = "services";
private static final String CHAINS_ENDPOINT = "chains";
/**
* 添加限流器配置
* @param addr 服务器地址
* @param name 限流器名称
* @param speed 限速值(MB)
* @param secret 认证密钥
* @return 请求结果
*/
public static GostDto AddLimiters(String addr, Long name, String speed, String secret) {
JSONObject data = createLimiterData(name, speed);
String url = buildUrl(addr, LIMITERS_ENDPOINT);
return HttpUtils.post(url, data, secret);
}
/**
* 更新限流器配置
* @param addr 服务器地址
* @param name 限流器名称
* @param speed 限速值(MB)
* @param secret 认证密钥
* @return 请求结果
*/
public static GostDto UpdateLimiters(String addr, Long name, String speed, String secret) {
JSONObject data = createLimiterData(name, speed);
String url = buildUrl(addr, LIMITERS_ENDPOINT + "/" + name);
return HttpUtils.put(url, data, secret);
}
/**
* 删除限流器配置
* @param addr 服务器地址
* @param name 限流器名称
* @param secret 认证密钥
* @return 请求结果
*/
public static GostDto DeleteLimiters(String addr, Long name, String secret) {
String url = buildUrl(addr, LIMITERS_ENDPOINT + "/" + name);
return HttpUtils.delete(url, secret);
}
/**
* 创建限流器数据
*/
private static JSONObject createLimiterData(Long name, String speed) {
JSONObject data = new JSONObject();
data.put("name", name.toString());
JSONArray limits = new JSONArray();
limits.add("$ " + speed + "MB " + speed + "MB");
data.put("limits", limits);
return data;
}
/**
* 添加服务配置(支持端口转发和隧道转发)
* @param addr 服务器地址
* @param name 服务名称
* @param in_port 监听端口
* @param limiter 限流器ID
* @param remoteAddr 远程地址(端口转发时使用)
* @param secret 认证密钥
* @param fow_type 转发类型:1=端口转发,2=隧道转发
* @return 请求结果
*/
public static GostDto AddService(String addr, String name, Integer in_port, Integer limiter, String remoteAddr, String secret, Integer fow_type) {
JSONArray services = new JSONArray();
String[] protocols = {"tcp", "udp"};
for (String protocol : protocols) {
JSONObject service = createServiceConfig(name, in_port, limiter, remoteAddr, protocol, fow_type);
services.add(service);
}
String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch");
return HttpUtils.post(url, services, secret);
}
/**
* 更新服务配置(批量更新TCP和UDP服务)
* @param addr 服务器地址
* @param name 服务名称
* @param in_port 监听端口
* @param limiter 限流器ID
* @param remoteAddr 远程地址(端口转发时使用)
* @param secret 认证密钥
* @param fow_type 转发类型:1=端口转发,2=隧道转发
* @return 请求结果
*/
public static GostDto UpdateService(String addr, String name, Integer in_port, Integer limiter, String remoteAddr, String secret, Integer fow_type) {
JSONArray services = new JSONArray();
String[] protocols = {"tcp", "udp"};
for (String protocol : protocols) {
JSONObject service = createServiceConfig(name, in_port, limiter, remoteAddr, protocol, fow_type);
services.add(service);
}
String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch");
return HttpUtils.put(url, services, secret);
}
/**
* 删除服务配置(批量删除TCP和UDP服务)
* @param addr 服务器地址
* @param name 服务名称
* @param secret 认证密钥
* @return 请求结果
*/
public static GostDto DeleteService(String addr, String name, String secret) {
JSONObject data = new JSONObject();
JSONArray services = new JSONArray();
services.add(name + "_tcp");
services.add(name + "_udp");
data.put("services", services);
String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch");
return HttpUtils.delete(url, data, secret);
}
public static GostDto PauseService(String addr, String name, String secret) {
JSONObject data = new JSONObject();
JSONArray services = new JSONArray();
services.add(name + "_tcp");
services.add(name + "_udp");
data.put("services", services);
String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch/pause");
return HttpUtils.post(url, data, secret);
}
public static GostDto ResumeService(String addr, String name, String secret) {
JSONObject data = new JSONObject();
JSONArray services = new JSONArray();
services.add(name + "_tcp");
services.add(name + "_udp");
data.put("services", services);
String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch/resume");
return HttpUtils.post(url, data, secret);
}
public static GostDto PauseRemoteService(String addr, String name, String secret) {
JSONObject data = new JSONObject();
JSONArray services = new JSONArray();
services.add(name + "_tls");
data.put("services", services);
String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch/pause");
return HttpUtils.post(url, data, secret);
}
public static GostDto ResumeRemoteService(String addr, String name, String secret) {
JSONObject data = new JSONObject();
JSONArray services = new JSONArray();
services.add(name + "_tls");
data.put("services", services);
String url = buildUrl(addr, SERVICES_ENDPOINT + "/batch/resume");
return HttpUtils.post(url, data, secret);
}
public static GostDto AddChains(String addr, String name, String remoteAddr, String secret) {
JSONObject dialer = new JSONObject();
dialer.put("type", "tls");
JSONObject connector = new JSONObject();
connector.put("type", "relay");
JSONObject node = new JSONObject();
node.put("name", "node-" + name);
node.put("addr", remoteAddr);
node.put("connector", connector);
node.put("dialer", dialer);
JSONArray nodes = new JSONArray();
nodes.add(node);
JSONObject hop = new JSONObject();
hop.put("name", "hop-" + name);
hop.put("nodes", nodes);
JSONArray hops = new JSONArray();
hops.add(hop);
JSONObject data = new JSONObject();
data.put("name", name + "_chains");
data.put("hops", hops);
String url = buildUrl(addr, CHAINS_ENDPOINT);
return HttpUtils.post(url, data, secret);
}
public static GostDto UpdateChains(String addr, String name, String remoteAddr, String secret) {
JSONObject dialer = new JSONObject();
dialer.put("type", "tls");
JSONObject connector = new JSONObject();
connector.put("type", "relay");
JSONObject node = new JSONObject();
node.put("name", "node-" + name);
node.put("addr", remoteAddr);
node.put("connector", connector);
node.put("dialer", dialer);
JSONArray nodes = new JSONArray();
nodes.add(node);
JSONObject hop = new JSONObject();
hop.put("name", "hop-" + name);
hop.put("nodes", nodes);
JSONArray hops = new JSONArray();
hops.add(hop);
JSONObject data = new JSONObject();
data.put("name", name + "_chains");
data.put("hops", hops);
String url = buildUrl(addr, CHAINS_ENDPOINT + "/" + name + "_chains");
return HttpUtils.put(url, data, secret);
}
public static GostDto DeleteChains(String addr, String name, String secret) {
String url = buildUrl(addr, CHAINS_ENDPOINT + "/" + name + "_chains");
return HttpUtils.delete(url, secret);
}
public static GostDto AddRemoteService(String addr, String name, Integer out_port, String remoteAddr, String secret) {
JSONObject data = new JSONObject();
data.put("name", name + "_tls");
data.put("addr", ":"+out_port);
JSONObject handler = new JSONObject();
handler.put("type", "relay");
data.put("handler", handler);
JSONObject listener = new JSONObject();
listener.put("type", "tls");
data.put("listener", listener);
JSONObject forwarder = new JSONObject();
JSONArray nodes = new JSONArray();
JSONObject node = new JSONObject();
node.put("name", name + "_node");
node.put("addr", remoteAddr);
nodes.add(node);
forwarder.put("nodes", nodes);
data.put("forwarder", forwarder);
String url = buildUrl(addr, SERVICES_ENDPOINT);
return HttpUtils.post(url, data, secret);
}
public static GostDto UpdateRemoteService(String addr, String name, Integer out_port, String remoteAddr, String secret) {
JSONObject data = new JSONObject();
data.put("name", name + "_tls");
data.put("addr", ":"+out_port);
JSONObject handler = new JSONObject();
handler.put("type", "relay");
data.put("handler", handler);
JSONObject listener = new JSONObject();
listener.put("type", "tls");
data.put("listener", listener);
JSONObject forwarder = new JSONObject();
JSONArray nodes = new JSONArray();
JSONObject node = new JSONObject();
node.put("name", name + "_node");
node.put("addr", remoteAddr);
nodes.add(node);
forwarder.put("nodes", nodes);
data.put("forwarder", forwarder);
String url = buildUrl(addr, SERVICES_ENDPOINT + "/" + name + "_tls");
return HttpUtils.put(url, data, secret);
}
public static GostDto DeleteRemoteService(String addr, String name, String secret) {
String url = buildUrl(addr, SERVICES_ENDPOINT + "/" + name + "_tls");
return HttpUtils.delete(url, secret);
}
/**
* 创建单个服务配置
*/
private static JSONObject createServiceConfig(String name, Integer in_port, Integer limiter, String remoteAddr, String protocol, Integer fow_type) {
JSONObject service = new JSONObject();
service.put("name", name + "_" + protocol);
service.put("addr", ":" +in_port);
// 添加限流器配置
if (limiter != null) {
service.put("limiter", limiter.toString());
}
// 配置处理器
JSONObject handler = createHandler(protocol, name, fow_type);
service.put("handler", handler);
// 配置监听器
JSONObject listener = createListener(protocol);
service.put("listener", listener);
// 端口转发需要配置转发器
if (isPortForwarding(fow_type)) {
JSONObject forwarder = createForwarder(protocol, remoteAddr);
service.put("forwarder", forwarder);
}
return service;
}
/**
* 创建处理器配置
*/
private static JSONObject createHandler(String protocol, String name, Integer fow_type) {
JSONObject handler = new JSONObject();
handler.put("type", protocol);
// 隧道转发需要添加链配置
if (isTunnelForwarding(fow_type)) {
handler.put("chain", name + "_chains");
}
return handler;
}
/**
* 创建监听器配置
*/
private static JSONObject createListener(String protocol) {
JSONObject listener = new JSONObject();
listener.put("type", protocol);
return listener;
}
/**
* 创建转发器配置
*/
private static JSONObject createForwarder(String protocol, String remoteAddr) {
JSONObject forwarder = new JSONObject();
JSONArray nodes = new JSONArray();
JSONObject node = new JSONObject();
node.put("name", protocol);
node.put("addr", remoteAddr);
nodes.add(node);
forwarder.put("nodes", nodes);
return forwarder;
}
/**
* 判断是否为端口转发
*/
private static boolean isPortForwarding(Integer fow_type) {
return fow_type != null && fow_type == 1;
}
/**
* 判断是否为隧道转发
*/
private static boolean isTunnelForwarding(Integer fow_type) {
return fow_type != null && fow_type != 1;
}
/**
* 构建API URL
*/
private static String buildUrl(String addr, String endpoint) {
return "http://" + addr + API_BASE_URL + endpoint;
}
}
@@ -0,0 +1,14 @@
package com.admin.common.utils;
import org.springframework.web.context.request.RequestContextHolder;
import org.springframework.web.context.request.ServletRequestAttributes;
import javax.servlet.http.HttpServletRequest;
public class HttpContextUtils {
public static HttpServletRequest getHttpServletRequest(){
return ((ServletRequestAttributes) RequestContextHolder.getRequestAttributes()).getRequest();
}
}
@@ -0,0 +1,178 @@
package com.admin.common.utils;
import com.admin.common.dto.GostDto;
import com.admin.config.RestTemplateConfig;
import com.alibaba.fastjson.JSONObject;
import lombok.SneakyThrows;
import org.apache.http.HttpResponse;
import org.apache.http.NameValuePair;
import org.apache.http.client.config.RequestConfig;
import org.apache.http.client.entity.UrlEncodedFormEntity;
import org.apache.http.client.methods.CloseableHttpResponse;
import org.apache.http.client.methods.HttpGet;
import org.apache.http.client.methods.HttpPost;
import org.apache.http.client.utils.URIBuilder;
import org.apache.http.entity.ContentType;
import org.apache.http.entity.StringEntity;
import org.apache.http.impl.client.CloseableHttpClient;
import org.apache.http.impl.client.HttpClients;
import org.apache.http.message.BasicNameValuePair;
import org.apache.http.util.EntityUtils;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.http.*;
import org.springframework.http.client.ClientHttpResponse;
import org.springframework.web.client.ResponseErrorHandler;
import org.springframework.web.client.RestTemplate;
import java.io.IOException;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import java.util.*;
/**
* HTTP请求工具类
* 支持GET和POST请求,支持表单和JSON格式的请求体
*/
public class HttpUtils {
private static final Logger logger = LoggerFactory.getLogger(HttpUtils.class);
/**
* 自定义错误处理器,不抛出异常
*/
private static class NoOpResponseErrorHandler implements ResponseErrorHandler {
@Override
public boolean hasError(ClientHttpResponse response) throws IOException {
return false;
}
@Override
public void handleError(ClientHttpResponse response) throws IOException {
}
}
@SneakyThrows
public static GostDto post(String url, Object requestBody, String secret){
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON));
String auth = secret + ":" + secret;
String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8));
headers.set("Authorization", "Basic " + encodedAuth);
HttpEntity<Object> entity = new HttpEntity<>(requestBody, headers);
RestTemplate restTemplate = new RestTemplate(RestTemplateConfig.generateHttpRequestFactory());
restTemplate.setErrorHandler(new NoOpResponseErrorHandler());
try {
ResponseEntity<GostDto> response = restTemplate.postForEntity(url, entity, GostDto.class);
GostDto body = response.getBody();
if (body.getMsg() != null && body.getMsg().contains("exists")) {
body.setMsg("OK");
}
return body;
} catch (Exception e) {
GostDto gostDto = new GostDto();
gostDto.setCode(500);
gostDto.setMsg("请求失败");
return gostDto;
}
}
@SneakyThrows
public static GostDto put(String url, Object requestBody, String secret){
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON));
String auth = secret + ":" + secret;
String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8));
headers.set("Authorization", "Basic " + encodedAuth);
HttpEntity<Object> entity = new HttpEntity<>(requestBody, headers);
RestTemplate restTemplate = new RestTemplate(RestTemplateConfig.generateHttpRequestFactory());
restTemplate.setErrorHandler(new NoOpResponseErrorHandler());
try {
ResponseEntity<GostDto> response = restTemplate.exchange(
url,
HttpMethod.PUT,
entity,
GostDto.class
);
GostDto body = response.getBody();
return body;
} catch (Exception e) {
GostDto gostDto = new GostDto();
gostDto.setCode(500);
gostDto.setMsg("请求失败");
return gostDto;
}
}
@SneakyThrows
public static GostDto delete(String url, String secret) {
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON));
// Basic Auth
String auth = secret + ":" + secret;
String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8));
headers.set("Authorization", "Basic " + encodedAuth);
HttpEntity<Void> entity = new HttpEntity<>(headers);
RestTemplate restTemplate = new RestTemplate(RestTemplateConfig.generateHttpRequestFactory());
restTemplate.setErrorHandler(new NoOpResponseErrorHandler());
try {
ResponseEntity<GostDto> response = restTemplate.exchange(
url,
HttpMethod.DELETE,
entity,
GostDto.class
);
GostDto body = response.getBody();
if (body != null && body.getMsg() != null && body.getMsg().contains("not found")) {
body.setMsg("OK");
}
return body;
} catch (Exception e) {
GostDto gostDto = new GostDto();
gostDto.setCode(500);
gostDto.setMsg("请求失败");
return gostDto;
}
}
@SneakyThrows
public static GostDto delete(String url, JSONObject data, String secret) {
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.setAccept(Collections.singletonList(MediaType.APPLICATION_JSON));
// Basic Auth
String auth = secret + ":" + secret;
String encodedAuth = Base64.getEncoder().encodeToString(auth.getBytes(StandardCharsets.UTF_8));
headers.set("Authorization", "Basic " + encodedAuth);
HttpEntity<JSONObject> entity = new HttpEntity<>(data, headers);
RestTemplate restTemplate = new RestTemplate(RestTemplateConfig.generateHttpRequestFactory());
restTemplate.setErrorHandler(new NoOpResponseErrorHandler());
try {
ResponseEntity<GostDto> response = restTemplate.exchange(
url,
HttpMethod.DELETE,
entity,
GostDto.class
);
GostDto body = response.getBody();
if (body != null && body.getMsg() != null && body.getMsg().contains("not found")) {
body.setMsg("OK");
}
return body;
} catch (Exception e) {
GostDto gostDto = new GostDto();
gostDto.setCode(500);
gostDto.setMsg("请求失败");
return gostDto;
}
}
}
@@ -0,0 +1,47 @@
package com.admin.common.utils;
import javax.servlet.http.HttpServletRequest;
import java.net.InetAddress;
import java.net.UnknownHostException;
public class IpUtils {
public static String getIpAddr(HttpServletRequest request) {
String ipAddress = null;
try {
ipAddress = request.getHeader("x-forwarded-for");
if (ipAddress == null || ipAddress.length() == 0 || "unknown".equalsIgnoreCase(ipAddress)) {
ipAddress = request.getHeader("Proxy-Client-IP");
}
if (ipAddress == null || ipAddress.length() == 0 || "unknown".equalsIgnoreCase(ipAddress)) {
ipAddress = request.getHeader("WL-Proxy-Client-IP");
}
if (ipAddress == null || ipAddress.length() == 0 || "unknown".equalsIgnoreCase(ipAddress)) {
ipAddress = request.getRemoteAddr();
if (ipAddress.equals("127.0.0.1")) {
// 根据网卡取本机配置的IP
InetAddress inet = null;
try {
inet = InetAddress.getLocalHost();
} catch (UnknownHostException e) {
e.printStackTrace();
}
ipAddress = inet.getHostAddress();
}
}
// 对于通过多个代理的情况,第一个IP为客户端真实IP,多个IP按照','分割
if (ipAddress != null && ipAddress.length() > 15) {
// "***.***.***.***".length()
// = 15
if (ipAddress.indexOf(",") > 0) {
ipAddress = ipAddress.substring(0, ipAddress.indexOf(","));
}
}
} catch (Exception e) {
ipAddress="";
}
return ipAddress;
}
}
@@ -0,0 +1,194 @@
package com.admin.common.utils;
import com.admin.entity.User;
import com.alibaba.fastjson2.JSON;
import lombok.SneakyThrows;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Component;
import javax.annotation.PostConstruct;
import javax.crypto.Mac;
import javax.crypto.spec.SecretKeySpec;
import java.nio.charset.StandardCharsets;
import java.util.Base64;
import java.util.Date;
import java.util.HashMap;
import java.util.Map;
/**
* JWT工具类,不使用第三方库实现
*/
@Component
public class JwtUtil {
@Value("${jwt-secret}")
private String secretKey;
private static String SECRET_KEY;
// token有效期,7天
private static final long EXPIRE_TIME = 7 * 24 * 60 * 60 * 1000;
// 算法
private static final String ALGORITHM = "HmacSHA256";
@PostConstruct
public void init() {
SECRET_KEY = this.secretKey;
}
/**
* 生成JWT Token
*
* @param user 用户信息
* @return 生成的JWT Token
*/
public static String generateToken(User user) {
try {
long nowMillis = System.currentTimeMillis();
Date now = new Date(nowMillis);
Date expireDate = new Date(nowMillis + EXPIRE_TIME);
// Header
Map<String, Object> header = new HashMap<>();
header.put("alg", ALGORITHM);
header.put("typ", "JWT");
String headerJson = JSON.toJSONString(header);
String encodedHeader = Base64.getUrlEncoder().withoutPadding()
.encodeToString(headerJson.getBytes(StandardCharsets.UTF_8));
// Payload
Map<String, Object> payload = new HashMap<>();
payload.put("sub", user.getId().toString());
payload.put("iat", now.getTime() / 1000); // 发布时间
payload.put("exp", expireDate.getTime() / 1000); // 过期时间
payload.put("user", user.getUser());
payload.put("name", user.getName());
payload.put("role_id", user.getRoleId());
String payloadJson = JSON.toJSONString(payload);
String encodedPayload = Base64.getUrlEncoder().withoutPadding()
.encodeToString(payloadJson.getBytes(StandardCharsets.UTF_8));
// Signature
String signature = calculateSignature(encodedHeader, encodedPayload);
// Token
return encodedHeader + "." + encodedPayload + "." + signature;
} catch (Exception e) {
throw new RuntimeException("JWT token generation failed", e);
}
}
/**
* 验证JWT Token
*
* @param token JWT Token
* @return 验证是否通过
*/
public static boolean validateToken(String token) {
try {
if (token == null || token.isEmpty()) {
return false;
}
String[] parts = token.split("\\.");
if (parts.length != 3) {
return false;
}
String encodedHeader = parts[0];
String encodedPayload = parts[1];
String signature = parts[2];
// 验证签名
String expectedSignature = calculateSignature(encodedHeader, encodedPayload);
if (!expectedSignature.equals(signature)) {
return false;
}
// 验证过期时间
String decodedPayload = new String(Base64.getUrlDecoder().decode(encodedPayload), StandardCharsets.UTF_8);
Map<String, Object> payload = JSON.parseObject(decodedPayload, Map.class);
long exp = Long.parseLong(payload.get("exp").toString());
long now = System.currentTimeMillis() / 1000;
return exp > now;
} catch (Exception e) {
return false;
}
}
/**
* 从JWT Token中获取用户ID
*
* @param token JWT Token
* @return 用户ID
*/
public static Long getUserIdFromToken(String token) {
String[] parts = token.split("\\.");
String encodedPayload = parts[1];
String decodedPayload = new String(Base64.getUrlDecoder().decode(encodedPayload), StandardCharsets.UTF_8);
Map<String, Object> payload = JSON.parseObject(decodedPayload, Map.class);
return Long.parseLong(payload.get("sub").toString());
}
public static Integer getUserIdFromToken() {
String token = HttpContextUtils.getHttpServletRequest().getHeader("Authorization");
String[] parts = token.split("\\.");
String encodedPayload = parts[1];
String decodedPayload = new String(Base64.getUrlDecoder().decode(encodedPayload), StandardCharsets.UTF_8);
Map<String, Object> payload = JSON.parseObject(decodedPayload, Map.class);
return Integer.parseInt(payload.get("sub").toString());
}
public static String getNameFromToken() {
String token = HttpContextUtils.getHttpServletRequest().getHeader("Authorization");
String[] parts = token.split("\\.");
String encodedPayload = parts[1];
String decodedPayload = new String(Base64.getUrlDecoder().decode(encodedPayload), StandardCharsets.UTF_8);
Map<String, Object> payload = JSON.parseObject(decodedPayload, Map.class);
return payload.get("name").toString();
}
/**
* 从JWT Token中获取用户角色ID
*
* @param token JWT Token
* @return 角色ID
*/
public static Integer getRoleIdFromToken(String token) {
String[] parts = token.split("\\.");
String encodedPayload = parts[1];
String decodedPayload = new String(Base64.getUrlDecoder().decode(encodedPayload), StandardCharsets.UTF_8);
Map<String, Object> payload = JSON.parseObject(decodedPayload, Map.class);
return Integer.parseInt(payload.get("role_id").toString());
}
@SneakyThrows
public static Integer getRoleIdFromToken() {
String token = HttpContextUtils.getHttpServletRequest().getHeader("Authorization");
if (token == null || token.isEmpty()) throw new Exception();
String[] parts = token.split("\\.");
String encodedPayload = parts[1];
String decodedPayload = new String(Base64.getUrlDecoder().decode(encodedPayload), StandardCharsets.UTF_8);
Map<String, Object> payload = JSON.parseObject(decodedPayload, Map.class);
return Integer.parseInt(payload.get("role_id").toString());
}
/**
* 计算签名
*
* @param encodedHeader 编码后的头部
* @param encodedPayload 编码后的负载
* @return 签名
* @throws Exception 签名计算异常
*/
private static String calculateSignature(String encodedHeader, String encodedPayload) throws Exception {
String content = encodedHeader + "." + encodedPayload;
Mac hmac = Mac.getInstance(ALGORITHM);
SecretKeySpec secretKeySpec = new SecretKeySpec(SECRET_KEY.getBytes(StandardCharsets.UTF_8), ALGORITHM);
hmac.init(secretKeySpec);
byte[] signatureBytes = hmac.doFinal(content.getBytes(StandardCharsets.UTF_8));
return Base64.getUrlEncoder().withoutPadding().encodeToString(signatureBytes);
}
}
@@ -0,0 +1,172 @@
package com.admin.common.utils;
import java.security.MessageDigest;
import java.security.NoSuchAlgorithmException;
import java.security.SecureRandom;
import java.util.Base64;
/**
* MD5工具类
*/
public class Md5Util {
private static final String MD5_ALGORITHM = "MD5";
private static final String DEFAULT_SALT = "admin_salt_2024";
/**
* 基础MD5加密
*
* @param input 待加密字符串
* @return MD5加密后的字符串(32位小写)
*/
public static String md5(String input) {
if (input == null || input.isEmpty()) {
return null;
}
try {
MessageDigest md = MessageDigest.getInstance(MD5_ALGORITHM);
byte[] digest = md.digest(input.getBytes());
return bytesToHex(digest);
} catch (NoSuchAlgorithmException e) {
throw new RuntimeException("MD5算法不可用", e);
}
}
/**
* MD5加密(使用默认盐值)
*
* @param input 待加密字符串
* @return MD5加密后的字符串
*/
public static String md5WithSalt(String input) {
return md5WithSalt(input, DEFAULT_SALT);
}
/**
* MD5加密(使用自定义盐值)
*
* @param input 待加密字符串
* @param salt 盐值
* @return MD5加密后的字符串
*/
public static String md5WithSalt(String input, String salt) {
if (input == null || input.isEmpty()) {
return null;
}
if (salt == null) {
salt = DEFAULT_SALT;
}
return md5(input + salt);
}
/**
* 生成随机盐值
*
* @param length 盐值长度
* @return 随机盐值
*/
public static String generateSalt(int length) {
SecureRandom random = new SecureRandom();
byte[] salt = new byte[length];
random.nextBytes(salt);
return Base64.getEncoder().encodeToString(salt);
}
/**
* 生成默认长度(16字节)的随机盐值
*
* @return 随机盐值
*/
public static String generateSalt() {
return generateSalt(16);
}
/**
* 验证密码
*
* @param password 原始密码
* @param hashedPassword 已加密的密码
* @return 是否匹配
*/
public static boolean verify(String password, String hashedPassword) {
if (password == null || hashedPassword == null) {
return false;
}
String encrypted = md5WithSalt(password);
return encrypted.equals(hashedPassword);
}
/**
* 验证密码(使用自定义盐值)
*
* @param password 原始密码
* @param salt 盐值
* @param hashedPassword 已加密的密码
* @return 是否匹配
*/
public static boolean verify(String password, String salt, String hashedPassword) {
if (password == null || hashedPassword == null) {
return false;
}
String encrypted = md5WithSalt(password, salt);
return encrypted.equals(hashedPassword);
}
/**
* 多次MD5加密
*
* @param input 待加密字符串
* @param times 加密次数
* @return 加密后的字符串
*/
public static String md5Multiple(String input, int times) {
if (input == null || input.isEmpty() || times <= 0) {
return input;
}
String result = input;
for (int i = 0; i < times; i++) {
result = md5(result);
}
return result;
}
/**
* 字节数组转十六进制字符串
*
* @param bytes 字节数组
* @return 十六进制字符串
*/
private static String bytesToHex(byte[] bytes) {
StringBuilder result = new StringBuilder();
for (byte b : bytes) {
result.append(String.format("%02x", b));
}
return result.toString();
}
/**
* 获取文件的MD5值
*
* @param bytes 文件字节数组
* @return MD5值
*/
public static String getFileMd5(byte[] bytes) {
if (bytes == null || bytes.length == 0) {
return null;
}
try {
MessageDigest md = MessageDigest.getInstance(MD5_ALGORITHM);
byte[] digest = md.digest(bytes);
return bytesToHex(digest);
} catch (NoSuchAlgorithmException e) {
throw new RuntimeException("MD5算法不可用", e);
}
}
}
@@ -0,0 +1,616 @@
package com.admin.common.utils;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.data.redis.core.RedisTemplate;
import org.springframework.data.redis.core.ZSetOperations;
import org.springframework.stereotype.Component;
import org.springframework.util.CollectionUtils;
import java.util.Collection;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.TimeUnit;
@Component
public class RedisUtil {
@Autowired
private RedisTemplate redisTemplate;
/**
* 指定缓存失效时间
*
* @param key 键
* @param time 时间(秒)
* @return
*/
public boolean expire(String key, long time) {
try {
if (time > 0) {
redisTemplate.expire(key, time, TimeUnit.SECONDS);
}
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 根据key 获取过期时间
*
* @param key 键 不能为null
* @return 时间(秒) 返回0代表为永久有效
*/
public long getExpire(String key) {
return redisTemplate.getExpire(key, TimeUnit.SECONDS);
}
/**
* 判断key是否存在
*
* @param key 键
* @return true 存在 false不存在
*/
public boolean hasKey(String key) {
try {
return redisTemplate.hasKey(key);
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 删除缓存
*
* @param key 可以传一个值 或多个
*/
@SuppressWarnings("unchecked")
public void del(String... key) {
if (key != null && key.length > 0) {
if (key.length == 1) {
redisTemplate.delete(key[0]);
} else {
redisTemplate.delete(CollectionUtils.arrayToList(key));
}
}
}
//============================String=============================
/**
* 普通缓存获取
*
* @param key 键
* @return 值
*/
public Object get(String key) {
return key == null ? null : redisTemplate.opsForValue().get(key);
}
/**
* 普通缓存放入
*
* @param key 键
* @param value 值
* @return true成功 false失败
*/
public boolean set(String key, Object value) {
try {
redisTemplate.opsForValue().set(key, value);
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 普通缓存放入并设置时间
*
* @param key 键
* @param value 值
* @param time 时间(秒) time要大于0 如果time小于等于0 将设置无限期
* @return true成功 false 失败
*/
public boolean set(String key, Object value, long time) {
try {
if (time > 0) {
redisTemplate.opsForValue().set(key, value, time, TimeUnit.SECONDS);
} else {
set(key, value);
}
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 递增
*
* @param key 键
* @param delta 要增加几(大于0)
* @return
*/
public long incr(String key, long delta) {
if (delta < 0) {
throw new RuntimeException("递增因子必须大于0");
}
return redisTemplate.opsForValue().increment(key, delta);
}
/**
* 递减
*
* @param key 键
* @param delta 要减少几(小于0)
* @return
*/
public long decr(String key, long delta) {
if (delta < 0) {
throw new RuntimeException("递减因子必须大于0");
}
return redisTemplate.opsForValue().increment(key, -delta);
}
//================================Map=================================
/**
* HashGet
*
* @param key 键 不能为null
* @param item 项 不能为null
* @return 值
*/
public Object hget(String key, String item) {
return redisTemplate.opsForHash().get(key, item);
}
/**
* 获取hashKey对应的所有键值
*
* @param key 键
* @return 对应的多个键值
*/
public Map<Object, Object> hmget(String key) {
return redisTemplate.opsForHash().entries(key);
}
/**
* HashSet
*
* @param key 键
* @param map 对应多个键值
* @return true 成功 false 失败
*/
public boolean hmset(String key, Map<String, Object> map) {
try {
redisTemplate.opsForHash().putAll(key, map);
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* HashSet 并设置时间
*
* @param key 键
* @param map 对应多个键值
* @param time 时间(秒)
* @return true成功 false失败
*/
public boolean hmset(String key, Map<String, Object> map, long time) {
try {
redisTemplate.opsForHash().putAll(key, map);
if (time > 0) {
expire(key, time);
}
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 向一张hash表中放入数据,如果不存在将创建
*
* @param key 键
* @param item 项
* @param value 值
* @return true 成功 false失败
*/
public boolean hset(String key, String item, Object value) {
try {
redisTemplate.opsForHash().put(key, item, value);
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 向一张hash表中放入数据,如果不存在将创建
*
* @param key 键
* @param item 项
* @param value 值
* @param time 时间(秒) 注意:如果已存在的hash表有时间,这里将会替换原有的时间
* @return true 成功 false失败
*/
public boolean hset(String key, String item, Object value, long time) {
try {
redisTemplate.opsForHash().put(key, item, value);
if (time > 0) {
expire(key, time);
}
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 删除hash表中的值
*
* @param key 键 不能为null
* @param item 项 可以使多个 不能为null
*/
public void hdel(String key, Object... item) {
redisTemplate.opsForHash().delete(key, item);
}
/**
* 判断hash表中是否有该项的值
*
* @param key 键 不能为null
* @param item 项 不能为null
* @return true 存在 false不存在
*/
public boolean hHasKey(String key, String item) {
return redisTemplate.opsForHash().hasKey(key, item);
}
/**
* hash递增 如果不存在,就会创建一个 并把新增后的值返回
*
* @param key 键
* @param item 项
* @param by 要增加几(大于0)
* @return
*/
public double hincr(String key, String item, double by) {
return redisTemplate.opsForHash().increment(key, item, by);
}
/**
* hash递减
*
* @param key 键
* @param item 项
* @param by 要减少记(小于0)
* @return
*/
public double hdecr(String key, String item, double by) {
return redisTemplate.opsForHash().increment(key, item, -by);
}
//============================set=============================
/**
* 根据key获取Set中的所有值
*
* @param key 键
* @return
*/
public Set<Object> sGet(String key) {
try {
return redisTemplate.opsForSet().members(key);
} catch (Exception e) {
e.printStackTrace();
return null;
}
}
/**
* 根据value从一个set中查询,是否存在
*
* @param key 键
* @param value 值
* @return true 存在 false不存在
*/
public boolean sHasKey(String key, Object value) {
try {
return redisTemplate.opsForSet().isMember(key, value);
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 将数据放入set缓存
*
* @param key 键
* @param values 值 可以是多个
* @return 成功个数
*/
public long sSet(String key, Object... values) {
try {
return redisTemplate.opsForSet().add(key, values);
} catch (Exception e) {
e.printStackTrace();
return 0;
}
}
/**
* 将set数据放入缓存
*
* @param key 键
* @param time 时间(秒)
* @param values 值 可以是多个
* @return 成功个数
*/
public long sSetAndTime(String key, long time, Object... values) {
try {
Long count = redisTemplate.opsForSet().add(key, values);
if (time > 0) expire(key, time);
return count;
} catch (Exception e) {
e.printStackTrace();
return 0;
}
}
/**
* 获取set缓存的长度
*
* @param key 键
* @return
*/
public long sGetSetSize(String key) {
try {
return redisTemplate.opsForSet().size(key);
} catch (Exception e) {
e.printStackTrace();
return 0;
}
}
/**
* 移除值为value的
*
* @param key 键
* @param values 值 可以是多个
* @return 移除的个数
*/
public long setRemove(String key, Object... values) {
try {
Long count = redisTemplate.opsForSet().remove(key, values);
return count;
} catch (Exception e) {
e.printStackTrace();
return 0;
}
}
//===============================list=================================
/**
* 获取list缓存的内容
*
* @param key 键
* @param start 开始
* @param end 结束 0 到 -1代表所有值
* @return
*/
public List<Object> lGet(String key, long start, long end) {
try {
return redisTemplate.opsForList().range(key, start, end);
} catch (Exception e) {
e.printStackTrace();
return null;
}
}
/**
* 获取list缓存的长度
*
* @param key 键
* @return
*/
public long lGetListSize(String key) {
try {
return redisTemplate.opsForList().size(key);
} catch (Exception e) {
e.printStackTrace();
return 0;
}
}
/**
* 通过索引 获取list中的值
*
* @param key 键
* @param index 索引 index>=0时, 0 表头,1 第二个元素,依次类推;index<0时,-1,表尾,-2倒数第二个元素,依次类推
* @return
*/
public Object lGetIndex(String key, long index) {
try {
return redisTemplate.opsForList().index(key, index);
} catch (Exception e) {
e.printStackTrace();
return null;
}
}
/**
* 将list放入缓存
*
* @param key 键
* @param value 值
* @return
*/
public boolean lSet(String key, Object value) {
try {
redisTemplate.opsForList().rightPush(key, value);
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 将list放入缓存
*
* @param key 键
* @param value 值
* @param time 时间(秒)
* @return
*/
public boolean lSet(String key, Object value, long time) {
try {
redisTemplate.opsForList().rightPush(key, value);
if (time > 0) expire(key, time);
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 将list放入缓存
*
* @param key 键
* @param value 值
* @return
*/
public boolean lSet(String key, List<Object> value) {
try {
redisTemplate.opsForList().rightPushAll(key, value);
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 将list放入缓存
*
* @param key 键
* @param value 值
* @param time 时间(秒)
* @return
*/
public boolean lSet(String key, List<Object> value, long time) {
try {
redisTemplate.opsForList().rightPushAll(key, value);
if (time > 0) expire(key, time);
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 根据索引修改list中的某条数据
*
* @param key 键
* @param index 索引
* @param value 值
* @return
*/
public boolean lUpdateIndex(String key, long index, Object value) {
try {
redisTemplate.opsForList().set(key, index, value);
return true;
} catch (Exception e) {
e.printStackTrace();
return false;
}
}
/**
* 移除N个值为value
*
* @param key 键
* @param count 移除多少个
* @param value 值
* @return 移除的个数
*/
public long lRemove(String key, long count, Object value) {
try {
Long remove = redisTemplate.opsForList().remove(key, count, value);
return remove;
} catch (Exception e) {
e.printStackTrace();
return 0;
}
}
//================有序集合 sort set===================
/**
* 有序set添加元素
*
* @param key
* @param value
* @param score
* @return
*/
public boolean zSet(String key, Object value, double score) {
return redisTemplate.opsForZSet().add(key, value, score);
}
public long batchZSet(String key, Set<ZSetOperations.TypedTuple> typles) {
return redisTemplate.opsForZSet().add(key, typles);
}
public void zIncrementScore(String key, Object value, long delta) {
redisTemplate.opsForZSet().incrementScore(key, value, delta);
}
public void zUnionAndStore(String key, Collection otherKeys, String destKey) {
redisTemplate.opsForZSet().unionAndStore(key, otherKeys, destKey);
}
/**
* 获取zset数量
* @param key
* @param value
* @return
*/
public long getZsetScore(String key, Object value) {
Double score = redisTemplate.opsForZSet().score(key, value);
if(score==null){
return 0;
}else{
return score.longValue();
}
}
/**
* 获取有序集 key 中成员 member 的排名 。
* 其中有序集成员按 score 值递减 (从大到小) 排序。
* @param key
* @param start
* @param end
* @return
*/
public Set<ZSetOperations.TypedTuple> getZSetRank(String key, long start, long end) {
return redisTemplate.opsForZSet().reverseRangeWithScores(key, start, end);
}
}
@@ -0,0 +1,160 @@
package com.admin.common.utils;
import com.admin.entity.Node;
import com.admin.service.NodeService;
import com.alibaba.fastjson.JSONObject;
import lombok.SneakyThrows;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils;
import org.springframework.web.socket.CloseStatus;
import org.springframework.web.socket.TextMessage;
import org.springframework.web.socket.WebSocketSession;
import org.springframework.web.socket.handler.TextWebSocketHandler;
import javax.annotation.Resource;
import java.util.Objects;
import java.util.concurrent.CopyOnWriteArraySet;
import java.util.concurrent.ConcurrentHashMap;
@Slf4j
public class WebSocketServer extends TextWebSocketHandler {
@Resource
NodeService nodeService;
// 存储所有活跃的 WebSocket 连接
private static final CopyOnWriteArraySet<WebSocketSession> activeSessions = new CopyOnWriteArraySet<>();
// 为每个session提供锁对象,防止并发发送消息
private static final ConcurrentHashMap<String, Object> sessionLocks = new ConcurrentHashMap<>();
//接受客户端消息
@Override
public void handleTextMessage(WebSocketSession session, TextMessage message) {
try {
if (StringUtils.isNoneBlank(message.getPayload())) {
//log.info("收到消息: {}", message.getPayload());
String id = session.getAttributes().get("id").toString();
String type = session.getAttributes().get("type").toString();
// 先发送确认消息
sendToUser(session, "ok");
// 如果是节点类型,转发消息给其他会话
if (Objects.equals(type, "1")) {
JSONObject jsonObject = new JSONObject();
jsonObject.put("id", id);
jsonObject.put("type", "info");
jsonObject.put("data", message.getPayload());
String broadcastMessage = jsonObject.toJSONString();
// 异步处理广播消息,避免阻塞当前线程
for (WebSocketSession targetSession : activeSessions) {
if (targetSession != null && targetSession.isOpen() && !targetSession.equals(session)) {
sendToUser(targetSession, broadcastMessage);
}
}
}
}
} catch (Exception e) {
log.error("处理WebSocket消息时发生异常: {}", e.getMessage(), e);
}
}
// 建立连接
@Override
public void afterConnectionEstablished(WebSocketSession session) {
try {
String id = session.getAttributes().get("id").toString();
String type = session.getAttributes().get("type").toString();
if (!Objects.equals(type, "1")) {
activeSessions.add(session);
}else {
Node byId = nodeService.getById(id);
if (byId != null) {
byId.setStatus(1);
nodeService.updateById(byId);
JSONObject res = new JSONObject();
res.put("id", id);
res.put("type", "status");
res.put("data", 1);
broadcastMessage(res.toJSONString());
}
}
log.info("WebSocket 连接建立成功 - id: {}, type: {}, 当前连接数: {}", id, type, activeSessions.size());
} catch (Exception e) {
log.error("建立连接时发生异常: {}", e.getMessage(), e);
}
}
// 连接关闭后
@Override
public void afterConnectionClosed(WebSocketSession session, CloseStatus status) {
try {
String id = session.getAttributes().get("id").toString();
String type = session.getAttributes().get("type").toString();
String sessionId = session.getId();
if (!Objects.equals(type, "1")) {
activeSessions.remove(session);
}else {
Node byId = nodeService.getById(id);
if (byId != null) {
byId.setStatus(0);
nodeService.updateById(byId);
JSONObject res = new JSONObject();
res.put("id", id);
res.put("type", "status");
res.put("data", 0);
broadcastMessage(res.toJSONString());
}
}
// 清理session锁对象
sessionLocks.remove(sessionId);
log.info("WebSocket 连接关闭 - id: {}, sessionId: {}, 关闭状态: {}, 当前连接数: {}",
id, sessionId, status, activeSessions.size());
} catch (Exception e) {
log.error("关闭连接时发生异常: {}", e.getMessage(), e);
}
}
// 点对点发送消息
@SneakyThrows
public static void sendToUser(WebSocketSession socketSession, String message) {
if (socketSession != null && socketSession.isOpen()) {
String sessionId = socketSession.getId();
Object lock = sessionLocks.computeIfAbsent(sessionId, k -> new Object());
synchronized (lock) {
try {
if (socketSession.isOpen()) {
socketSession.sendMessage(new TextMessage(message));
}
} catch (Exception e) {
log.error("发送WebSocket消息失败 [sessionId={}]: {}", sessionId, e.getMessage());
activeSessions.remove(socketSession);
sessionLocks.remove(sessionId);
}
}
} else {
activeSessions.remove(socketSession);
if (socketSession != null) {
sessionLocks.remove(socketSession.getId());
}
}
}
// 广播消息
public static void broadcastMessage(String message) {
for (WebSocketSession session : activeSessions) {
sendToUser(session, message);
}
}
}
@@ -0,0 +1,27 @@
package com.admin.config;
import com.baomidou.mybatisplus.autoconfigure.ConfigurationCustomizer;
import com.baomidou.mybatisplus.extension.plugins.MybatisPlusInterceptor;
import com.baomidou.mybatisplus.extension.plugins.inner.BlockAttackInnerInterceptor;
import com.baomidou.mybatisplus.extension.plugins.inner.PaginationInnerInterceptor;
import org.mybatis.spring.annotation.MapperScan;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
@Configuration
@MapperScan("com.admin.mapper")
public class MybatisPlusConfig {
@Bean
public MybatisPlusInterceptor mybatisPlusInterceptor() {
MybatisPlusInterceptor interceptor = new MybatisPlusInterceptor();
interceptor.addInnerInterceptor(new PaginationInnerInterceptor()); // 分页插件
interceptor.addInnerInterceptor(new BlockAttackInnerInterceptor()); // 防止全表更新插件
return interceptor;
}
@Bean
public ConfigurationCustomizer configurationCustomizer() {
return configuration -> configuration.setUseDeprecatedExecutor(false);
}
}
@@ -0,0 +1,34 @@
package com.admin.config;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.data.redis.connection.RedisConnectionFactory;
import org.springframework.data.redis.core.RedisTemplate;
import org.springframework.data.redis.serializer.Jackson2JsonRedisSerializer;
import org.springframework.data.redis.serializer.StringRedisSerializer;
@Configuration
public class RedisConfig {
// 序列化redis
@Bean
RedisTemplate redisTemplate(RedisConnectionFactory redisConnectionFactory) {
RedisTemplate redisTemplate = new RedisTemplate();
redisTemplate.setConnectionFactory(redisConnectionFactory);
Jackson2JsonRedisSerializer jackson2JsonRedisSerializer = new Jackson2JsonRedisSerializer(Object.class);
jackson2JsonRedisSerializer.setObjectMapper(new ObjectMapper());
redisTemplate.setKeySerializer(new StringRedisSerializer());
redisTemplate.setValueSerializer(jackson2JsonRedisSerializer);
redisTemplate.setHashKeySerializer(new StringRedisSerializer());
redisTemplate.setHashValueSerializer(jackson2JsonRedisSerializer);
return redisTemplate;
}
}
@@ -0,0 +1,68 @@
package com.admin.config;
import org.apache.http.HttpHost;
import org.apache.http.conn.ssl.NoopHostnameVerifier;
import org.apache.http.conn.ssl.SSLConnectionSocketFactory;
import org.apache.http.impl.client.CloseableHttpClient;
import org.apache.http.impl.client.HttpClientBuilder;
import org.apache.http.impl.client.HttpClients;
import org.apache.http.ssl.SSLContexts;
import org.apache.http.ssl.TrustStrategy;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.http.client.ClientHttpRequestFactory;
import org.springframework.http.client.HttpComponentsClientHttpRequestFactory;
import org.springframework.web.client.RestTemplate;
import javax.net.ssl.SSLContext;
import java.security.KeyManagementException;
import java.security.KeyStoreException;
import java.security.NoSuchAlgorithmException;
@Configuration
public class RestTemplateConfig {
@Bean
public RestTemplate restTemplate(ClientHttpRequestFactory factory){
return new RestTemplate(factory);
}
@Bean
public ClientHttpRequestFactory simpleClientHttpRequestFactory(){
HttpComponentsClientHttpRequestFactory factory = new HttpComponentsClientHttpRequestFactory();
factory.setConnectTimeout(15000);
factory.setReadTimeout(5000);
return factory;
}
// 禁用ssl证书校验
public static HttpComponentsClientHttpRequestFactory generateHttpRequestFactory() throws NoSuchAlgorithmException, KeyManagementException, KeyStoreException{
TrustStrategy acceptingTrustStrategy = (x509Certificates, authType) -> true;
SSLContext sslContext = SSLContexts.custom().loadTrustMaterial(null, acceptingTrustStrategy).build();
SSLConnectionSocketFactory connectionSocketFactory = new SSLConnectionSocketFactory(sslContext, new NoopHostnameVerifier());
HttpClientBuilder httpClientBuilder = HttpClients.custom();
httpClientBuilder.setSSLSocketFactory(connectionSocketFactory);
CloseableHttpClient httpClient = httpClientBuilder.build();
HttpComponentsClientHttpRequestFactory factory = new HttpComponentsClientHttpRequestFactory();
factory.setHttpClient(httpClient);
return factory;
}
// 禁用ssl证书校验 添加代理ip post请求
public static HttpComponentsClientHttpRequestFactory post(String url, Integer port) throws NoSuchAlgorithmException, KeyManagementException, KeyStoreException{
TrustStrategy acceptingTrustStrategy = (x509Certificates, authType) -> true;
SSLContext sslContext = SSLContexts.custom().loadTrustMaterial(null, acceptingTrustStrategy).build();
SSLConnectionSocketFactory connectionSocketFactory = new SSLConnectionSocketFactory(sslContext, new NoopHostnameVerifier());
HttpClientBuilder httpClientBuilder = HttpClients.custom();
httpClientBuilder.setProxy(new HttpHost(url, port, "http"));
httpClientBuilder.setSSLSocketFactory(connectionSocketFactory);
CloseableHttpClient httpClient = httpClientBuilder.build();
HttpComponentsClientHttpRequestFactory factory = new HttpComponentsClientHttpRequestFactory();
factory.setHttpClient(httpClient);
return factory;
}
}
@@ -0,0 +1,62 @@
package com.admin.config;
import com.admin.common.interceptor.JwtInterceptor;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.web.cors.CorsConfiguration;
import org.springframework.web.cors.UrlBasedCorsConfigurationSource;
import org.springframework.web.filter.CorsFilter;
import org.springframework.web.servlet.config.annotation.CorsRegistry;
import org.springframework.web.servlet.config.annotation.EnableWebMvc;
import org.springframework.web.servlet.config.annotation.InterceptorRegistry;
import org.springframework.web.servlet.config.annotation.WebMvcConfigurer;
@Configuration
@EnableWebMvc
public class WebMvcConfig implements WebMvcConfigurer {
private CorsConfiguration buildConfig() {
CorsConfiguration corsConfiguration = new CorsConfiguration();
corsConfiguration.addAllowedOrigin("*");
corsConfiguration.addAllowedHeader("*");
corsConfiguration.addAllowedMethod("*");
corsConfiguration.addExposedHeader("Authorization");
return corsConfiguration;
}
@Bean
public CorsFilter corsFilter() {
UrlBasedCorsConfigurationSource source = new UrlBasedCorsConfigurationSource();
source.registerCorsConfiguration("/**", buildConfig());
return new CorsFilter(source);
}
@Override
public void addCorsMappings(CorsRegistry registry) {
registry.addMapping("/**")
.allowedOrigins("*")
.allowedMethods("GET", "POST", "DELETE", "PUT")
.maxAge(3600);
}
/**
* JWT拦截器
*/
@Bean
public JwtInterceptor jwtInterceptor() {
return new JwtInterceptor();
}
/**
* 添加JWT拦截器
*/
@Override
public void addInterceptors(InterceptorRegistry registry) {
// 添加JWT拦截器,不拦截登录接口
registry.addInterceptor(jwtInterceptor())
.addPathPatterns("/api/**")
.excludePathPatterns("/flow/**")
.excludePathPatterns("/api/v1/user/login");
}
}
@@ -0,0 +1,36 @@
package com.admin.config;
import com.admin.common.utils.WebSocketServer;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.web.socket.WebSocketHandler;
import org.springframework.web.socket.config.annotation.EnableWebSocket;
import org.springframework.web.socket.config.annotation.WebSocketConfigurer;
import org.springframework.web.socket.config.annotation.WebSocketHandlerRegistry;
import javax.annotation.Resource;
@Configuration
@EnableWebSocket
public class WebSocketConfig implements WebSocketConfigurer {
@Resource
private WebSocketInterceptor webSocketInterceptor;
@Override
public void registerWebSocketHandlers(WebSocketHandlerRegistry webSocketHandlerRegistry) {
webSocketHandlerRegistry
.addHandler(myHandler(), "/system-info")
.setAllowedOrigins("*")
.addInterceptors(webSocketInterceptor);
}
@Bean
public WebSocketHandler myHandler() {
return new WebSocketServer();
}
}
@@ -0,0 +1,63 @@
package com.admin.config;
import com.admin.common.utils.JwtUtil;
import com.admin.common.utils.RedisUtil;
import com.admin.common.utils.WebSocketServer;
import com.admin.entity.Node;
import com.admin.service.NodeService;
import com.admin.service.UserService;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import lombok.extern.slf4j.Slf4j;
import org.springframework.context.annotation.Configuration;
import org.springframework.http.server.ServerHttpRequest;
import org.springframework.http.server.ServerHttpResponse;
import org.springframework.http.server.ServletServerHttpRequest;
import org.springframework.web.socket.CloseStatus;
import org.springframework.web.socket.WebSocketHandler;
import org.springframework.web.socket.WebSocketSession;
import org.springframework.web.socket.server.support.HttpSessionHandshakeInterceptor;
import javax.annotation.Resource;
import java.util.List;
import java.util.Map;
import java.util.Objects;
@Configuration
@Slf4j
public class WebSocketInterceptor extends HttpSessionHandshakeInterceptor {
@Resource
NodeService nodeService;
@Override
public void afterHandshake(ServerHttpRequest request, ServerHttpResponse response, WebSocketHandler wsHandler, Exception ex) {
}
@Override
public boolean beforeHandshake(ServerHttpRequest request, ServerHttpResponse response, WebSocketHandler wsHandler, Map<String, Object> attributes) throws Exception {
ServletServerHttpRequest serverHttpRequest = (ServletServerHttpRequest) request;
String secret = serverHttpRequest.getServletRequest().getParameter("secret");
String type = serverHttpRequest.getServletRequest().getParameter("type");
if (Objects.equals(type, "1")) {
String client_ip = serverHttpRequest.getServletRequest().getParameter("client_ip");
Node node = nodeService.getOne(new QueryWrapper<Node>().eq("secret", secret));
if (node == null) return false;
attributes.put("id", node.getId());
node.setStatus(1);
if (!Objects.equals(node.getIp(), client_ip)){
node.setIp(client_ip);
}
nodeService.updateById(node);
}else {
boolean b = JwtUtil.validateToken(secret);
if (!b) return false;
attributes.put("id", JwtUtil.getUserIdFromToken(secret));
}
attributes.put("type", type);
return true;
}
}
@@ -0,0 +1,26 @@
package com.admin.controller;
import com.admin.service.*;
import org.springframework.beans.factory.annotation.Autowired;
public class BaseController {
@Autowired
UserService userService;
@Autowired
NodeService nodeService;
@Autowired
UserTunnelService userTunnelService;
@Autowired
TunnelService tunnelService;
@Autowired
ForwardService forwardService;
@Autowired
SpeedLimitService speedLimitService;
}
@@ -0,0 +1,349 @@
package com.admin.controller;
import com.admin.common.aop.LogAnnotation;
import com.admin.common.dto.FlowDto;
import com.admin.common.lang.R;
import com.admin.common.utils.GostUtil;
import com.admin.entity.*;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper;
import org.springframework.web.bind.annotation.CrossOrigin;
import org.springframework.web.bind.annotation.RequestBody;
import org.springframework.web.bind.annotation.RequestMapping;
import org.springframework.web.bind.annotation.RestController;
import java.util.List;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
import java.util.stream.Collectors;
/**
* 流量上报控制器
* 处理节点上报的流量数据,更新用户和隧道的流量统计
*
* 并发安全解决方案:
* 1. 使用UpdateWrapper进行数据库层面的原子更新操作,避免读取-修改-写入的竞态条件
* 2. 使用synchronized锁确保同一用户/隧道的流量更新串行执行
* 3. 这样可以避免相同用户相同隧道不同转发同时上报时流量统计丢失的问题
*/
@RestController
@RequestMapping("/flow")
@CrossOrigin
public class FlowController extends BaseController {
// 常量定义
private static final String SUCCESS_RESPONSE = "ok";
private static final String ERROR_RESPONSE = "err1";
private static final String DEFAULT_USER_TUNNEL_ID = "0";
private static final int FLOW_TYPE_UPLOAD_ONLY = 1;
private static final int FLOW_TYPE_BIDIRECTIONAL = 2;
private static final long BYTES_TO_GB = 1024L * 1024L * 1024L;
// 用于同步相同用户和隧道的流量更新操作
private static final ConcurrentHashMap<String, Object> USER_LOCKS = new ConcurrentHashMap<>();
private static final ConcurrentHashMap<String, Object> TUNNEL_LOCKS = new ConcurrentHashMap<>();
@RequestMapping("/test")
@LogAnnotation
public String test() {
return "test";
}
/**
* 处理流量数据上报
* @param flowDataList 流量数据列表
* @param secret 节点密钥
* @return 处理结果
*/
@RequestMapping("/upload")
@LogAnnotation
public String uploadFlowData(@RequestBody List<FlowDto> flowDataList, String secret) {
// 1. 验证节点权限
if (!isValidNode(secret)) {
return ERROR_RESPONSE;
}
// 2. 过滤有效流量数据
List<FlowDto> validFlowData = filterValidFlowData(flowDataList);
if (validFlowData.isEmpty()) {
return SUCCESS_RESPONSE;
}
// 3. 解析服务名称获取ID信息
String[] serviceIds = parseServiceName(validFlowData.get(0).getN());
String forwardId = serviceIds[0];
String userId = serviceIds[1];
String userTunnelId = serviceIds[2];
// 4. 计算总流量
FlowStatistics flowStats = calculateTotalFlow(validFlowData);
// 5. 获取流量计费类型
int flowType = getFlowType(forwardId);
// 6. 更新各项流量统计
updateForwardFlow(forwardId, flowStats);
updateUserFlow(userId, flowStats, flowType);
updateUserTunnelFlow(userTunnelId, flowStats, flowType, forwardId, userId);
// 7. 检查用户总流量限制
checkUserTotalFlowLimit(userId, userTunnelId);
return SUCCESS_RESPONSE;
}
/**
* 验证节点是否有效
*/
private boolean isValidNode(String secret) {
int nodeCount = nodeService.count(new QueryWrapper<Node>().eq("secret", secret));
return nodeCount > 0;
}
/**
* 过滤有效的流量数据
*/
private List<FlowDto> filterValidFlowData(List<FlowDto> flowDataList) {
return flowDataList.stream()
.filter(flow -> flow.getU() != null && flow.getD() != null)
.filter(flow -> flow.getU() > 0 && flow.getD() > 0)
.collect(Collectors.toList());
}
/**
* 解析服务名称获取ID信息
*/
private String[] parseServiceName(String serviceName) {
return serviceName.split("_");
}
/**
* 计算总流量统计
*/
private FlowStatistics calculateTotalFlow(List<FlowDto> validFlowData) {
long totalUpload = 0L;
long totalDownload = 0L;
for (FlowDto flow : validFlowData) {
totalUpload += flow.getU();
totalDownload += flow.getD();
}
return new FlowStatistics(totalUpload, totalDownload);
}
/**
* 获取流量计费类型
*/
private int getFlowType(String forwardId) {
int defaultFlowType = FLOW_TYPE_BIDIRECTIONAL;
Forward forward = forwardService.getById(forwardId);
if (forward != null) {
Tunnel tunnel = tunnelService.getById(forward.getTunnelId());
if (tunnel != null) {
return tunnel.getFlow();
}
}
return defaultFlowType;
}
/**
* 更新转发流量统计 - 使用原子操作避免并发问题
*/
private void updateForwardFlow(String forwardId, FlowStatistics flowStats) {
UpdateWrapper<Forward> updateWrapper = new UpdateWrapper<>();
updateWrapper.eq("id", forwardId);
updateWrapper.setSql("in_flow = in_flow + " + flowStats.getDownload());
updateWrapper.setSql("out_flow = out_flow + " + flowStats.getUpload());
forwardService.update(null, updateWrapper);
}
/**
* 更新用户流量统计 - 使用原子操作避免并发问题
*/
private void updateUserFlow(String userId, FlowStatistics flowStats, int flowType) {
// 对相同用户的流量更新进行同步,避免并发覆盖
synchronized (getUserLock(userId)) {
UpdateWrapper<User> updateWrapper = new UpdateWrapper<>();
updateWrapper.eq("id", userId);
// 使用SQL的原子更新操作,避免读取-修改-写入的并发问题
if (flowType == FLOW_TYPE_BIDIRECTIONAL) {
// 双向计费:同时更新上传和下载流量
updateWrapper.setSql("in_flow = in_flow + " + flowStats.getDownload());
updateWrapper.setSql("out_flow = out_flow + " + flowStats.getUpload());
} else {
// 仅上传计费:只更新上传流量
updateWrapper.setSql("out_flow = out_flow + " + flowStats.getUpload());
}
userService.update(null, updateWrapper);
}
}
/**
* 更新用户隧道流量统计并检查限制
*/
private void updateUserTunnelFlow(String userTunnelId, FlowStatistics flowStats,
int flowType, String forwardId, String userId) {
if (Objects.equals(userTunnelId, DEFAULT_USER_TUNNEL_ID)) {
return;
}
// 对相同用户隧道的流量更新进行同步,避免并发覆盖
synchronized (getTunnelLock(userTunnelId)) {
UpdateWrapper<UserTunnel> updateWrapper = new UpdateWrapper<>();
updateWrapper.eq("id", userTunnelId);
updateWrapper.setSql("in_flow = in_flow + " + flowStats.getDownload());
updateWrapper.setSql("out_flow = out_flow + " + flowStats.getUpload());
boolean updateSuccess = userTunnelService.update(null, updateWrapper);
if (!updateSuccess) {
return; // 更新失败,可能记录不存在
}
}
// 重新获取最新的流量数据进行限制检查
UserTunnel userTunnel = userTunnelService.getById(userTunnelId);
if (userTunnel != null) {
checkUserTunnelFlowLimit(userTunnel, flowType, forwardId, userId, userTunnelId);
}
}
/**
* 检查用户隧道流量限制
*/
private void checkUserTunnelFlowLimit(UserTunnel userTunnel, int flowType,
String forwardId, String userId, String userTunnelId) {
long currentFlow = (flowType == FLOW_TYPE_UPLOAD_ONLY) ?
userTunnel.getOutFlow() :
userTunnel.getInFlow() + userTunnel.getOutFlow();
long flowLimit = userTunnel.getFlow() * BYTES_TO_GB;
if (flowLimit < currentFlow) {
pauseServiceDueToTunnelLimit(userTunnel.getTunnelId(), forwardId, userId, userTunnelId);
}
}
/**
* 因隧道流量超限暂停服务
*/
private void pauseServiceDueToTunnelLimit(Integer tunnelId, String forwardId,
String userId, String userTunnelId) {
Tunnel tunnel = tunnelService.getById(tunnelId);
if (tunnel != null) {
Node node = nodeService.getNodeById(tunnel.getInNodeId());
if (node != null) {
String serviceName = buildServiceName(forwardId, userId, userTunnelId);
GostUtil.PauseService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret());
// 隧道转发需要同时暂停远端服务
if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD
Node outNode = nodeService.getNodeById(tunnel.getOutNodeId());
if (outNode != null) {
GostUtil.PauseRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret());
}
}
}
}
// 更新转发状态为暂停
Forward forward = forwardService.getById(forwardId);
if (forward != null) {
forward.setStatus(0);
forwardService.updateById(forward);
}
}
/**
* 检查用户总流量限制
*/
private void checkUserTotalFlowLimit(String userId, String userTunnelId) {
User user = userService.getById(userId);
if (user == null) {
return;
}
long userFlowLimit = user.getFlow() * BYTES_TO_GB;
long userCurrentFlow = user.getInFlow() + user.getOutFlow();
if (userFlowLimit < userCurrentFlow) {
pauseAllUserServices(userId, userTunnelId);
}
}
/**
* 暂停用户所有服务
*/
private void pauseAllUserServices(String userId, String userTunnelId) {
List<Forward> userForwards = forwardService.list(new QueryWrapper<Forward>().eq("user_id", userId));
for (Forward forward : userForwards) {
Tunnel tunnel = tunnelService.getById(forward.getTunnelId());
if (tunnel != null) {
Node node = nodeService.getNodeById(tunnel.getInNodeId());
if (node != null) {
String serviceName = buildServiceName(String.valueOf(forward.getId()), userId, userTunnelId);
GostUtil.PauseService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret());
// 隧道转发需要同时暂停远端服务
if (tunnel.getType() == 2) { // TUNNEL_TYPE_TUNNEL_FORWARD
Node outNode = nodeService.getNodeById(tunnel.getOutNodeId());
if (outNode != null) {
GostUtil.PauseRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret());
}
}
}
}
forward.setStatus(0);
forwardService.updateById(forward);
}
}
/**
* 构建服务名称
*/
private String buildServiceName(String forwardId, String userId, String userTunnelId) {
return forwardId + "_" + userId + "_" + userTunnelId;
}
/**
* 获取用户锁对象
*/
private Object getUserLock(String userId) {
return USER_LOCKS.computeIfAbsent(userId, k -> new Object());
}
/**
* 获取隧道锁对象
*/
private Object getTunnelLock(String userTunnelId) {
return TUNNEL_LOCKS.computeIfAbsent(userTunnelId, k -> new Object());
}
/**
* 流量统计数据类
*/
private static class FlowStatistics {
private final long upload;
private final long download;
public FlowStatistics(long upload, long download) {
this.upload = upload;
this.download = download;
}
public long getUpload() {
return upload;
}
public long getDownload() {
return download;
}
}
}
@@ -0,0 +1,70 @@
package com.admin.controller;
import com.admin.common.aop.LogAnnotation;
import com.admin.common.annotation.RequireRole;
import com.admin.common.dto.ForwardDto;
import com.admin.common.dto.ForwardUpdateDto;
import com.admin.common.dto.PageDto;
import com.admin.common.lang.R;
import com.admin.service.ForwardService;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.validation.annotation.Validated;
import org.springframework.web.bind.annotation.*;
import java.util.Map;
/**
* <p>
* 前端控制器
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@RestController
@CrossOrigin
@RequestMapping("/api/v1/forward")
public class ForwardController extends BaseController {
@LogAnnotation
@PostMapping("/create")
public R create(@Validated @RequestBody ForwardDto forwardDto) {
return forwardService.createForward(forwardDto);
}
@LogAnnotation
@PostMapping("/list")
public R readAll() {
return forwardService.getAllForwards();
}
@LogAnnotation
@PostMapping("/update")
public R update(@Validated @RequestBody ForwardUpdateDto forwardUpdateDto) {
return forwardService.updateForward(forwardUpdateDto);
}
@LogAnnotation
@PostMapping("/delete")
public R delete(@RequestBody Map<String, Object> params) {
Long id = Long.valueOf(params.get("id").toString());
return forwardService.deleteForward(id);
}
@LogAnnotation
@PostMapping("/pause")
public R pause(@RequestBody Map<String, Object> params) {
Long id = Long.valueOf(params.get("id").toString());
return forwardService.pauseForward(id);
}
@LogAnnotation
@PostMapping("/resume")
public R resume(@RequestBody Map<String, Object> params) {
Long id = Long.valueOf(params.get("id").toString());
return forwardService.resumeForward(id);
}
}
@@ -0,0 +1,65 @@
package com.admin.controller;
import com.admin.common.annotation.RequireRole;
import com.admin.common.aop.LogAnnotation;
import com.admin.common.dto.NodeDto;
import com.admin.common.dto.NodeUpdateDto;
import com.admin.common.dto.PageDto;
import com.admin.common.lang.R;
import org.springframework.validation.annotation.Validated;
import org.springframework.web.bind.annotation.*;
import java.util.Map;
/**
* <p>
* 前端控制器
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@RestController
@CrossOrigin
@RequestMapping("/api/v1/node")
public class NodeController extends BaseController {
@LogAnnotation
@RequireRole
@PostMapping("/create")
public R create(@Validated @RequestBody NodeDto nodeDto) {
return nodeService.createNode(nodeDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/list")
public R list() {
return nodeService.getAllNodes();
}
@LogAnnotation
@RequireRole
@PostMapping("/update")
public R update(@Validated @RequestBody NodeUpdateDto nodeUpdateDto) {
return nodeService.updateNode(nodeUpdateDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/delete")
public R delete(@RequestBody Map<String, Object> params) {
Long id = Long.valueOf(params.get("id").toString());
return nodeService.deleteNode(id);
}
@LogAnnotation
@RequireRole
@PostMapping("/install")
public R getInstallCommand(@RequestBody Map<String, Object> params) {
Long id = Long.valueOf(params.get("id").toString());
return nodeService.getInstallCommand(id);
}
}
@@ -0,0 +1,70 @@
package com.admin.controller;
import com.admin.common.aop.LogAnnotation;
import com.admin.common.annotation.RequireRole;
import com.admin.common.dto.SpeedLimitDto;
import com.admin.common.dto.SpeedLimitUpdateDto;
import com.admin.common.lang.R;
import com.admin.service.SpeedLimitService;
import com.admin.service.TunnelService;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.validation.annotation.Validated;
import org.springframework.web.bind.annotation.*;
import java.util.Map;
/**
* <p>
* 限速规则前端控制器
* </p>
*
* @author QAQ
* @since 2025-06-04
*/
@RestController
@RequestMapping("/api/v1/speed-limit")
@CrossOrigin
public class SpeedLimitController extends BaseController {
@Autowired
private SpeedLimitService speedLimitService;
@Autowired
private TunnelService tunnelService;
@LogAnnotation
@RequireRole
@PostMapping("/create")
public R create(@Validated @RequestBody SpeedLimitDto speedLimitDto) {
return speedLimitService.createSpeedLimit(speedLimitDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/list")
public R list() {
return speedLimitService.getAllSpeedLimits();
}
@LogAnnotation
@RequireRole
@PostMapping("/update")
public R update(@Validated @RequestBody SpeedLimitUpdateDto speedLimitUpdateDto) {
return speedLimitService.updateSpeedLimit(speedLimitUpdateDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/delete")
public R delete(@RequestBody Map<String, Object> params) {
Long id = Long.valueOf(params.get("id").toString());
return speedLimitService.deleteSpeedLimit(id);
}
@LogAnnotation
@RequireRole
@PostMapping("/tunnels")
public R getTunnels() {
return tunnelService.getAllTunnels();
}
}
@@ -0,0 +1,135 @@
package com.admin.controller;
import com.admin.common.aop.LogAnnotation;
import com.admin.common.annotation.RequireRole;
import com.admin.common.dto.PageDto;
import com.admin.common.dto.TunnelDto;
import com.admin.common.dto.UserTunnelDto;
import com.admin.common.dto.UserTunnelQueryDto;
import com.admin.common.dto.UserTunnelUpdateDto;
import com.admin.common.lang.R;
import com.admin.service.TunnelService;
import com.admin.service.UserTunnelService;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.validation.annotation.Validated;
import org.springframework.web.bind.annotation.*;
import java.util.Map;
/**
* <p>
* 隧道前端控制器
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@RestController
@CrossOrigin
@RequestMapping("/api/v1/tunnel")
public class TunnelController extends BaseController {
@Autowired
private TunnelService tunnelService;
@Autowired
private UserTunnelService userTunnelService;
@LogAnnotation
@RequireRole
@PostMapping("/create")
public R create(@Validated @RequestBody TunnelDto tunnelDto) {
return tunnelService.createTunnel(tunnelDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/list")
public R readAll(@RequestBody(required = false) PageDto pageDto) {
return tunnelService.getAllTunnels();
}
@LogAnnotation
@RequireRole
@PostMapping("/delete")
public R delete(@RequestBody Map<String, Object> params) {
Long id = Long.valueOf(params.get("id").toString());
return tunnelService.deleteTunnel(id);
}
// ============ 用户隧道权限管理相关方法 ============
/**
* 分配用户隧道权限
* @param userTunnelDto 用户隧道权限数据
* @return 操作结果
*/
@LogAnnotation
@RequireRole
@PostMapping("/user/assign")
public R assignUserTunnel(@Validated @RequestBody UserTunnelDto userTunnelDto) {
return userTunnelService.assignUserTunnel(userTunnelDto);
}
/**
* 查询用户隧道权限列表
* @param queryDto 查询条件
* @return 用户隧道权限列表
*/
@LogAnnotation
@RequireRole
@PostMapping("/user/list")
public R getUserTunnelList(@RequestBody @Validated UserTunnelQueryDto queryDto) {
return userTunnelService.getUserTunnelList(queryDto);
}
/**
* 删除用户隧道权限
* @param params 包含userId和tunnelId的参数
* @return 操作结果
*/
@LogAnnotation
@RequireRole
@PostMapping("/user/remove")
public R removeUserTunnel(@RequestBody Map<String, Object> params) {
Integer id = Integer.valueOf(params.get("id").toString());
return userTunnelService.removeUserTunnel(id);
}
/**
* 更新用户隧道流量限制
* @param params 包含userId、tunnelId和flow的参数
* @return 操作结果
*/
@LogAnnotation
@RequireRole
@PostMapping("/user/updateFlow")
public R updateUserTunnelFlow(@RequestBody Map<String, Object> params) {
Integer id = Integer.valueOf(params.get("id").toString());
Long flow = Long.valueOf(params.get("flow").toString());
return userTunnelService.updateUserTunnelFlow(id, flow);
}
/**
* 更新用户隧道权限(包含流量、流量重置时间、到期时间)
* @param updateDto 更新数据
* @return 操作结果
*/
@LogAnnotation
@RequireRole
@PostMapping("/user/update")
public R updateUserTunnel(@Validated @RequestBody UserTunnelUpdateDto updateDto) {
return userTunnelService.updateUserTunnel(updateDto);
}
@LogAnnotation
@PostMapping("/user/tunnel")
public R userTunnel() {
return tunnelService.userTunnel();
}
}
@@ -0,0 +1,81 @@
package com.admin.controller;
import com.admin.common.aop.LogAnnotation;
import com.admin.common.annotation.RequireRole;
import com.admin.common.dto.ChangePasswordDto;
import com.admin.common.dto.LoginDto;
import com.admin.common.dto.PageDto;
import com.admin.common.dto.UserDto;
import com.admin.common.dto.UserUpdateDto;
import com.admin.common.lang.R;
import org.springframework.validation.annotation.Validated;
import org.springframework.web.bind.annotation.*;
import java.util.Map;
/**
* <p>
* 前端控制器
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@RestController
@CrossOrigin
@RequestMapping("/api/v1/user")
public class UserController extends BaseController {
@LogAnnotation
@PostMapping("/login")
public R login(@Validated @RequestBody LoginDto loginDto) {
return userService.login(loginDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/create")
public R create(@Validated @RequestBody UserDto userDto) {
return userService.createUser(userDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/list")
public R readAll(@RequestBody(required = false) PageDto pageDto) {
// 如果没有传分页参数,使用默认值
if (pageDto == null) {
pageDto = new PageDto();
}
return userService.getAllUsers(pageDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/update")
public R update(@Validated @RequestBody UserUpdateDto userUpdateDto) {
return userService.updateUser(userUpdateDto);
}
@LogAnnotation
@RequireRole
@PostMapping("/delete")
public R delete(@RequestBody Map<String, Object> params) {
Long id = Long.valueOf(params.get("id").toString());
return userService.deleteUser(id);
}
@LogAnnotation
@PostMapping("/package")
public R getUserPackageInfo() {
return userService.getUserPackageInfo();
}
@LogAnnotation
@PostMapping("/updatePassword")
public R updatePassword(@Validated @RequestBody ChangePasswordDto changePasswordDto) {
return userService.updatePassword(changePasswordDto);
}
}
@@ -0,0 +1,37 @@
package com.admin.entity;
import com.baomidou.mybatisplus.annotation.IdType;
import com.baomidou.mybatisplus.annotation.TableId;
import lombok.Data;
import java.io.Serializable;
/**
* 基础实体类,包含公共字段
*/
@Data
public class BaseEntity implements Serializable {
private static final long serialVersionUID = 1L;
/**
* 主键ID
*/
@TableId(value = "id", type = IdType.AUTO)
private Long id;
/**
* 创建时间(时间戳)
*/
private Long createdTime;
/**
* 更新时间(时间戳)
*/
private Long updatedTime;
/**
* 状态(0:正常,1:删除)
*/
private Integer status;
}
@@ -0,0 +1,40 @@
package com.admin.entity;
import java.io.Serializable;
import lombok.Data;
import lombok.EqualsAndHashCode;
/**
* <p>
*
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Data
@EqualsAndHashCode(callSuper = false)
public class Forward extends BaseEntity{
private static final long serialVersionUID = 1L;
private Integer userId;
private String userName;
private String name;
private Integer tunnelId;
private Integer inPort;
private Integer outPort;
private String remoteAddr;
private Long inFlow;
private Long outFlow;
}
@@ -0,0 +1,30 @@
package com.admin.entity;
import java.io.Serializable;
import lombok.Data;
import lombok.EqualsAndHashCode;
/**
* <p>
*
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Data
@EqualsAndHashCode(callSuper = true)
public class Node extends BaseEntity {
private static final long serialVersionUID = 1L;
private String name;
private String secret;
private String ip;
private Integer port;
}
@@ -0,0 +1,53 @@
package com.admin.entity;
import com.admin.entity.BaseEntity;
import com.baomidou.mybatisplus.annotation.IdType;
import com.baomidou.mybatisplus.annotation.TableId;
import lombok.Data;
import lombok.EqualsAndHashCode;
import java.io.Serializable;
/**
* <p>
*
* </p>
*
* @author QAQ
* @since 2025-06-04
*/
@Data
public class SpeedLimit implements Serializable {
private static final long serialVersionUID = 1L;
/**
* 主键ID
*/
@TableId(value = "id", type = IdType.AUTO)
private Long id;
/**
* 创建时间(时间戳)
*/
private Long createdTime;
/**
* 更新时间(时间戳)
*/
private Long updatedTime;
/**
* 状态(0:正常,1:删除)
*/
private Integer status;
private String name;
private Integer speed;
private Long tunnelId;
private String tunnelName;
}
@@ -0,0 +1,75 @@
package com.admin.entity;
import java.io.Serializable;
import lombok.Data;
import lombok.EqualsAndHashCode;
/**
* <p>
* 隧道实体类
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Data
@EqualsAndHashCode(callSuper = true)
public class Tunnel extends BaseEntity {
private static final long serialVersionUID = 1L;
/**
* 隧道名称
*/
private String name;
/**
* 入口节点ID
*/
private Long inNodeId;
/**
* 入口IP (兼容字段)
*/
private String inIp;
/**
* 入口端口开始
*/
private Integer inPortSta;
/**
* 入口端口结束
*/
private Integer inPortEnd;
/**
* 出口节点ID
*/
private Long outNodeId;
/**
* 出口IP (兼容字段)
*/
private String outIp;
/**
* 出口端口开始
*/
private Integer outIpSta;
/**
* 出口端口结束
*/
private Integer outIpEnd;
/**
* 隧道类型(1-端口转发,2-隧道转发)
*/
private Integer type;
/**
* 流量计算类型(1 单向计算上传。2 双向)
*/
private int flow;
}
@@ -0,0 +1,42 @@
package com.admin.entity;
import lombok.Data;
import lombok.EqualsAndHashCode;
/**
* <p>
*
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Data
@EqualsAndHashCode(callSuper = true)
public class User extends BaseEntity {
private static final long serialVersionUID = 1L;
private String name;
private String user;
private String pwd;
private Integer roleId;
private Long expTime;
private Long flow;
private Long inFlow;
private Long outFlow;
private Integer num;
private Long flowResetTime;
}
@@ -0,0 +1,53 @@
package com.admin.entity;
import java.io.Serializable;
import com.baomidou.mybatisplus.annotation.FieldStrategy;
import com.baomidou.mybatisplus.annotation.TableField;
import lombok.Data;
import lombok.EqualsAndHashCode;
import com.baomidou.mybatisplus.annotation.IdType;
import com.baomidou.mybatisplus.annotation.TableId;
/**
* <p>
*
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Data
@EqualsAndHashCode(callSuper = false)
public class UserTunnel implements Serializable {
private static final long serialVersionUID = 1L;
/**
* 主键ID
*/
@TableId(value = "id", type = IdType.AUTO)
private Integer id;
private Integer userId;
private Integer tunnelId;
private Long flow;
private Long inFlow;
private Long outFlow;
private Long flowResetTime;
private Long expTime;
@TableField(updateStrategy = FieldStrategy.IGNORED)
private Integer speedId;
private Integer num;
private Integer status;
}
@@ -0,0 +1,33 @@
package com.admin.mapper;
import com.admin.entity.Forward;
import com.admin.common.dto.ForwardWithTunnelDto;
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
import org.apache.ibatis.annotations.Param;
import java.util.List;
/**
* <p>
* Mapper 接口
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface ForwardMapper extends BaseMapper<Forward> {
/**
* 查询所有转发信息(包含隧道信息)
* @return 转发信息列表
*/
List<ForwardWithTunnelDto> selectAllForwardsWithTunnel();
/**
* 根据用户ID查询转发信息(包含隧道信息)
* @param userId 用户ID
* @return 转发信息列表
*/
List<ForwardWithTunnelDto> selectForwardsWithTunnelByUserId(@Param("userId") Integer userId);
}
@@ -0,0 +1,16 @@
package com.admin.mapper;
import com.admin.entity.Node;
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
/**
* <p>
* Mapper 接口
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface NodeMapper extends BaseMapper<Node> {
}
@@ -0,0 +1,16 @@
package com.admin.mapper;
import com.admin.entity.SpeedLimit;
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
/**
* <p>
* Mapper 接口
* </p>
*
* @author QAQ
* @since 2025-06-04
*/
public interface SpeedLimitMapper extends BaseMapper<SpeedLimit> {
}
@@ -0,0 +1,16 @@
package com.admin.mapper;
import com.admin.entity.Tunnel;
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
/**
* <p>
* Mapper 接口
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface TunnelMapper extends BaseMapper<Tunnel> {
}
@@ -0,0 +1,39 @@
package com.admin.mapper;
import com.admin.entity.User;
import com.admin.common.dto.UserPackageDto;
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
import org.apache.ibatis.annotations.Param;
import java.util.List;
/**
* <p>
* Mapper 接口
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface UserMapper extends BaseMapper<User> {
/**
* 查询用户隧道权限详情
* @param userId 用户ID
* @return 隧道权限列表
*/
List<UserPackageDto.UserTunnelDetailDto> getUserTunnelDetails(@Param("userId") Integer userId);
/**
* 查询用户转发详情
* @param userId 用户ID
* @return 转发列表
*/
List<UserPackageDto.UserForwardDetailDto> getUserForwardDetails(@Param("userId") Integer userId);
/**
* 管理员查询所有隧道(流量和转发设置为99999)
* @return 隧道列表
*/
List<UserPackageDto.UserTunnelDetailDto> getAllTunnelsForAdmin();
}
@@ -0,0 +1,24 @@
package com.admin.mapper;
import com.admin.entity.UserTunnel;
import com.admin.common.dto.UserTunnelWithDetailDto;
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
import org.apache.ibatis.annotations.Param;
import org.apache.ibatis.annotations.Select;
import java.util.List;
/**
* <p>
* Mapper 接口
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface UserTunnelMapper extends BaseMapper<UserTunnel> {
List<UserTunnelWithDetailDto> getUserTunnelWithDetails(@Param("userId") Integer userId);
}
@@ -0,0 +1,59 @@
package com.admin.service;
import com.admin.common.dto.ForwardDto;
import com.admin.common.dto.ForwardUpdateDto;
import com.admin.common.lang.R;
import com.admin.entity.Forward;
import com.baomidou.mybatisplus.extension.service.IService;
/**
* <p>
* 服务类
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface ForwardService extends IService<Forward> {
/**
* 创建端口转发
* @param forwardDto 转发数据
* @return 结果
*/
R createForward(ForwardDto forwardDto);
/**
* 获取端口转发列表
* @return 结果
*/
R getAllForwards();
/**
* 更新端口转发
* @param forwardUpdateDto 更新数据
* @return 结果
*/
R updateForward(ForwardUpdateDto forwardUpdateDto);
/**
* 删除端口转发
* @param id 转发ID
* @return 结果
*/
R deleteForward(Long id);
/**
* 暂停转发服务
* @param id 转发ID
* @return 结果
*/
R pauseForward(Long id);
/**
* 恢复转发服务
* @param id 转发ID
* @return 结果
*/
R resumeForward(Long id);
}
@@ -0,0 +1,31 @@
package com.admin.service;
import com.admin.common.dto.NodeDto;
import com.admin.common.dto.NodeUpdateDto;
import com.admin.common.dto.PageDto;
import com.admin.common.lang.R;
import com.admin.entity.Node;
import com.baomidou.mybatisplus.extension.service.IService;
/**
* <p>
* 服务类
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface NodeService extends IService<Node> {
R createNode(NodeDto nodeDto);
R getAllNodes();
R updateNode(NodeUpdateDto nodeUpdateDto);
R deleteNode(Long id);
Node getNodeById(Long id);
R getInstallCommand(Long id);
}
@@ -0,0 +1,45 @@
package com.admin.service;
import com.admin.common.dto.SpeedLimitDto;
import com.admin.common.dto.SpeedLimitUpdateDto;
import com.admin.common.lang.R;
import com.admin.entity.SpeedLimit;
import com.baomidou.mybatisplus.extension.service.IService;
/**
* <p>
* 限速规则服务类
* </p>
*
* @author QAQ
* @since 2025-06-04
*/
public interface SpeedLimitService extends IService<SpeedLimit> {
/**
* 创建限速规则
* @param speedLimitDto 限速规则数据
* @return 结果
*/
R createSpeedLimit(SpeedLimitDto speedLimitDto);
/**
* 获取所有限速规则
* @return 结果
*/
R getAllSpeedLimits();
/**
* 更新限速规则
* @param speedLimitUpdateDto 更新数据
* @return 结果
*/
R updateSpeedLimit(SpeedLimitUpdateDto speedLimitUpdateDto);
/**
* 删除限速规则
* @param id 限速规则ID
* @return 结果
*/
R deleteSpeedLimit(Long id);
}
@@ -0,0 +1,43 @@
package com.admin.service;
import com.admin.common.dto.PageDto;
import com.admin.common.dto.TunnelDto;
import com.admin.common.lang.R;
import com.admin.entity.Tunnel;
import com.baomidou.mybatisplus.extension.service.IService;
/**
* <p>
* 隧道服务类
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface TunnelService extends IService<Tunnel> {
/**
* 创建隧道
* @param tunnelDto 隧道数据
* @return 结果
*/
R createTunnel(TunnelDto tunnelDto);
/**
* 获取隧道列表
* @return 结果
*/
R getAllTunnels();
/**
* 删除隧道
* @param id 隧道ID
* @return 结果
*/
R deleteTunnel(Long id);
R userTunnel();
}
@@ -0,0 +1,35 @@
package com.admin.service;
import com.admin.common.dto.ChangePasswordDto;
import com.admin.common.dto.LoginDto;
import com.admin.common.dto.PageDto;
import com.admin.common.dto.UserDto;
import com.admin.common.dto.UserUpdateDto;
import com.admin.common.lang.R;
import com.admin.entity.User;
import com.baomidou.mybatisplus.extension.service.IService;
/**
* <p>
* 服务类
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface UserService extends IService<User> {
R login(LoginDto loginDto);
R createUser(UserDto userDto);
R getAllUsers(PageDto pageDto);
R updateUser(UserUpdateDto userUpdateDto);
R deleteUser(Long id);
R getUserPackageInfo();
R updatePassword(ChangePasswordDto changePasswordDto);
}
@@ -0,0 +1,56 @@
package com.admin.service;
import com.admin.common.dto.UserTunnelDto;
import com.admin.common.dto.UserTunnelQueryDto;
import com.admin.common.dto.UserTunnelUpdateDto;
import com.admin.common.lang.R;
import com.admin.entity.UserTunnel;
import com.baomidou.mybatisplus.extension.service.IService;
/**
* <p>
* 用户隧道权限服务类
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
public interface UserTunnelService extends IService<UserTunnel> {
/**
* 分配用户隧道权限
* @param userTunnelDto 用户隧道权限数据
* @return 结果
*/
R assignUserTunnel(UserTunnelDto userTunnelDto);
/**
* 查询用户隧道权限列表
* @param queryDto 查询条件
* @return 结果
*/
R getUserTunnelList(UserTunnelQueryDto queryDto);
/**
* 删除用户隧道权限
* @param id ID
* @return 结果
*/
R removeUserTunnel(Integer id);
/**
* 更新用户隧道流量限制
* @param id ID
* @param flow 新的流量限制
* @return 结果
*/
R updateUserTunnelFlow(Integer id, Long flow);
/**
* 更新用户隧道权限(包含流量、流量重置时间、到期时间)
* @param updateDto 更新数据
* @return 结果
*/
R updateUserTunnel(UserTunnelUpdateDto updateDto);
}
@@ -0,0 +1,856 @@
package com.admin.service.impl;
import cn.hutool.core.util.StrUtil;
import com.admin.common.dto.ForwardDto;
import com.admin.common.dto.ForwardUpdateDto;
import com.admin.common.dto.ForwardWithTunnelDto;
import com.admin.common.dto.GostDto;
import com.admin.common.lang.R;
import com.admin.common.utils.GostUtil;
import com.admin.common.utils.JwtUtil;
import com.admin.entity.*;
import com.admin.mapper.ForwardMapper;
import com.admin.service.*;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
import lombok.Data;
import org.springframework.beans.BeanUtils;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.context.annotation.Lazy;
import org.springframework.stereotype.Service;
import javax.annotation.Resource;
import javax.swing.*;
import java.util.List;
import java.util.Objects;
import java.util.Set;
import java.util.stream.Collectors;
/**
* <p>
* 端口转发服务实现类
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Service
public class ForwardServiceImpl extends ServiceImpl<ForwardMapper, Forward> implements ForwardService {
// 常量定义
private static final String GOST_SUCCESS_MSG = "OK";
private static final String GOST_NOT_FOUND_MSG = "not found";
private static final int ADMIN_ROLE_ID = 0;
private static final int TUNNEL_TYPE_PORT_FORWARD = 1;
private static final int TUNNEL_TYPE_TUNNEL_FORWARD = 2;
private static final int FORWARD_STATUS_ACTIVE = 1;
private static final int FORWARD_STATUS_PAUSED = 0;
private static final int FORWARD_STATUS_ERROR = -1;
private static final int TUNNEL_STATUS_ACTIVE = 1;
private static final int FLOW_TYPE_UPLOAD_ONLY = 1;
private static final long BYTES_TO_GB = 1024L * 1024L * 1024L;
@Resource
@Lazy
private TunnelService tunnelService;
@Resource
UserTunnelService userTunnelService;
@Resource
UserService userService;
@Resource
NodeService nodeService;
@Override
public R createForward(ForwardDto forwardDto) {
// 1. 获取当前用户信息
UserInfo currentUser = getCurrentUserInfo();
// 2. 检查隧道是否存在和可用
Tunnel tunnel = validateTunnel(forwardDto.getTunnelId());
if (tunnel == null) {
return R.err("隧道不存在");
}
if (tunnel.getStatus() != TUNNEL_STATUS_ACTIVE) {
return R.err("隧道已禁用,无法创建转发");
}
// 3. 普通用户权限和限制检查
UserPermissionResult permissionResult = checkUserPermissions(currentUser, tunnel, null);
if (permissionResult.isHasError()) {
return R.err(permissionResult.getErrorMessage());
}
// 4. 分配端口
PortAllocation portAllocation = allocatePorts(tunnel);
if (portAllocation.isHasError()) {
return R.err(portAllocation.getErrorMessage());
}
// 5. 创建并保存Forward对象
Forward forward = createForwardEntity(forwardDto, currentUser, portAllocation);
if (!this.save(forward)) {
return R.err("端口转发创建失败");
}
// 6. 调用Gost服务创建转发
R gostResult = createGostServices(forward, tunnel, permissionResult.getLimiter());
if (gostResult.getCode() != 0) {
return gostResult;
}
return R.ok();
}
@Override
public R getAllForwards() {
UserInfo currentUser = getCurrentUserInfo();
List<ForwardWithTunnelDto> forwardList;
if (currentUser.getRoleId() != ADMIN_ROLE_ID) {
forwardList = baseMapper.selectForwardsWithTunnelByUserId(currentUser.getUserId());
} else {
forwardList = baseMapper.selectAllForwardsWithTunnel();
}
return R.ok(forwardList);
}
@Override
public R updateForward(ForwardUpdateDto forwardUpdateDto) {
// 1. 获取当前用户信息
UserInfo currentUser = getCurrentUserInfo();
// 2. 检查转发是否存在
Forward existForward = validateForwardExists(forwardUpdateDto.getId(), currentUser);
if (existForward == null) {
return R.err("转发不存在");
}
// 3. 检查隧道是否存在和可用
Tunnel tunnel = validateTunnel(forwardUpdateDto.getTunnelId());
if (tunnel == null) {
return R.err("隧道不存在");
}
if (tunnel.getStatus() != TUNNEL_STATUS_ACTIVE) {
return R.err("隧道已禁用,无法更新转发");
}
// 4. 检查权限和限制(仅当隧道发生变化时)
UserPermissionResult permissionResult = null;
if (isTunnelChanged(existForward, forwardUpdateDto)) {
permissionResult = checkUserPermissions(currentUser, tunnel, forwardUpdateDto.getId());
if (permissionResult.isHasError()) {
return R.err(permissionResult.getErrorMessage());
}
}
// 5. 更新Forward对象
Forward updatedForward = updateForwardEntity(forwardUpdateDto, existForward, tunnel);
// 6. 调用Gost服务更新转发
R gostResult = updateGostServices(updatedForward, tunnel,
permissionResult != null ? permissionResult.getLimiter() : null);
if (gostResult.getCode() != 0) {
return gostResult;
}
// 7. 保存更新
boolean result = this.updateById(updatedForward);
return result ? R.ok("端口转发更新成功") : R.err("端口转发更新失败");
}
@Override
public R deleteForward(Long id) {
// 1. 获取当前用户信息
UserInfo currentUser = getCurrentUserInfo();
// 2. 检查转发是否存在
Forward forward = validateForwardExists(id, currentUser);
if (forward == null) {
return R.err("端口转发不存在");
}
// 3. 获取隧道信息
Tunnel tunnel = validateTunnel(forward.getTunnelId());
if (tunnel == null) {
return R.err("隧道不存在");
}
// 4. 权限检查(仅普通用户需要)
if (currentUser.getRoleId() != ADMIN_ROLE_ID) {
if (!hasUserTunnelPermission(currentUser.getUserId(), tunnel.getId().intValue())) {
return R.err("你没有该隧道权限");
}
}
// 5. 调用Gost服务删除转发
R gostResult = deleteGostServices(forward, tunnel);
if (gostResult.getCode() != 0) {
return gostResult;
}
// 6. 删除转发记录
boolean result = this.removeById(id);
if (result) {
// 归还用户转发条数(普通用户才需要归还)
returnUserForwardQuota(currentUser);
return R.ok("端口转发删除成功");
} else {
return R.err("端口转发删除失败");
}
}
@Override
public R pauseForward(Long id) {
return changeForwardStatus(id, FORWARD_STATUS_PAUSED, "暂停", "PauseService");
}
@Override
public R resumeForward(Long id) {
return changeForwardStatus(id, FORWARD_STATUS_ACTIVE, "恢复", "ResumeService");
}
/**
* 改变转发状态(暂停/恢复)
*/
private R changeForwardStatus(Long id, int targetStatus, String operation, String gostMethod) {
// 1. 获取当前用户信息
UserInfo currentUser = getCurrentUserInfo();
// 2. 检查转发是否存在
Forward forward = validateForwardExists(id, currentUser);
if (forward == null) {
return R.err("转发不存在");
}
// 3. 获取隧道信息
Tunnel tunnel = validateTunnel(forward.getTunnelId());
if (tunnel == null) {
return R.err("隧道不存在");
}
// 4. 恢复服务时需要额外检查
if (targetStatus == FORWARD_STATUS_ACTIVE) {
if (tunnel.getStatus() != TUNNEL_STATUS_ACTIVE) {
return R.err("隧道已禁用,无法恢复服务");
}
// 普通用户需要检查流量和账户状态
if (currentUser.getRoleId() != ADMIN_ROLE_ID) {
R flowCheckResult = checkUserFlowLimits(currentUser.getUserId(), tunnel);
if (flowCheckResult.getCode() != 0) {
return flowCheckResult;
}
}
}
// 5. 权限检查(仅普通用户需要)
if (currentUser.getRoleId() != ADMIN_ROLE_ID) {
if (!hasUserTunnelPermission(currentUser.getUserId(), tunnel.getId().intValue())) {
return R.err("你没有该隧道权限");
}
}
// 6. 调用Gost服务
Node node = nodeService.getNodeById(tunnel.getInNodeId());
if (node == null) {
return R.err("节点不存在");
}
String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId());
GostDto gostResult;
if ("PauseService".equals(gostMethod)) {
gostResult = GostUtil.PauseService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret());
// 隧道转发需要同时暂停远端服务
if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) {
Node outNode = nodeService.getNodeById(tunnel.getOutNodeId());
if (outNode != null) {
GostDto remoteResult = GostUtil.PauseRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret());
if (!isGostOperationSuccess(remoteResult)) {
return R.err(operation + "远端服务失败:" + remoteResult.getMsg());
}
}
}
} else {
gostResult = GostUtil.ResumeService(node.getIp() + ":" + node.getPort(), serviceName, node.getSecret());
// 隧道转发需要同时恢复远端服务
if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) {
Node outNode = nodeService.getNodeById(tunnel.getOutNodeId());
if (outNode != null) {
GostDto remoteResult = GostUtil.ResumeRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret());
if (!isGostOperationSuccess(remoteResult)) {
return R.err(operation + "远端服务失败:" + remoteResult.getMsg());
}
}
}
}
if (!isGostOperationSuccess(gostResult)) {
return R.err(operation + "服务失败:" + gostResult.getMsg());
}
// 7. 更新转发状态
forward.setStatus(targetStatus);
forward.setUpdatedTime(System.currentTimeMillis());
boolean result = this.updateById(forward);
return result ? R.ok("服务已" + operation) : R.err("更新状态失败");
}
/**
* 获取当前用户信息
*/
private UserInfo getCurrentUserInfo() {
Integer userId = JwtUtil.getUserIdFromToken();
Integer roleId = JwtUtil.getRoleIdFromToken();
String userName = JwtUtil.getNameFromToken();
return new UserInfo(userId, roleId, userName);
}
/**
* 验证隧道是否存在
*/
private Tunnel validateTunnel(Integer tunnelId) {
return tunnelService.getById(tunnelId);
}
/**
* 验证转发是否存在且用户有权限访问
*/
private Forward validateForwardExists(Long forwardId, UserInfo currentUser) {
Forward forward = this.getById(forwardId);
if (forward == null) {
return null;
}
// 普通用户只能操作自己的转发
if (currentUser.getRoleId() != ADMIN_ROLE_ID &&
!Objects.equals(currentUser.getUserId(), forward.getUserId())) {
return null;
}
return forward;
}
/**
* 检查用户权限和限制
*/
private UserPermissionResult checkUserPermissions(UserInfo currentUser, Tunnel tunnel, Long excludeForwardId) {
if (currentUser.getRoleId() == ADMIN_ROLE_ID) {
return UserPermissionResult.success(null);
}
// 获取用户信息
User userInfo = userService.getById(currentUser.getUserId());
if (userInfo.getExpTime() != null && userInfo.getExpTime() <= System.currentTimeMillis()) {
return UserPermissionResult.error("当前账号已到期");
}
// 检查用户隧道权限
UserTunnel userTunnel = getUserTunnel(currentUser.getUserId(), tunnel.getId().intValue());
if (userTunnel == null) {
return UserPermissionResult.error("你没有该隧道权限");
}
// 检查隧道权限到期时间
if (userTunnel.getExpTime() != null && userTunnel.getExpTime() <= System.currentTimeMillis()) {
return UserPermissionResult.error("该隧道权限已到期");
}
// 流量限制检查
if (userInfo.getFlow() <= 0) {
return UserPermissionResult.error("用户总流量已用完");
}
if (userTunnel.getFlow() <= 0) {
return UserPermissionResult.error("该隧道流量已用完");
}
// 转发数量限制检查
R quotaCheckResult = checkForwardQuota(currentUser.getUserId(), tunnel.getId().intValue(), userTunnel, userInfo, excludeForwardId);
if (quotaCheckResult.getCode() != 0) {
return UserPermissionResult.error(quotaCheckResult.getMsg());
}
return UserPermissionResult.success(userTunnel.getSpeedId());
}
/**
* 检查用户转发数量限制
*/
private R checkForwardQuota(Integer userId, Integer tunnelId, UserTunnel userTunnel, User userInfo, Long excludeForwardId) {
// 检查用户总转发数量限制
long userForwardCount = this.count(new QueryWrapper<Forward>().eq("user_id", userId));
if (userForwardCount >= userInfo.getNum()) {
return R.err("用户总转发数量已达上限,当前限制:" + userInfo.getNum() + "个");
}
// 检查用户在该隧道的转发数量限制
QueryWrapper<Forward> tunnelQuery = new QueryWrapper<Forward>()
.eq("user_id", userId)
.eq("tunnel_id", tunnelId);
if (excludeForwardId != null) {
tunnelQuery.ne("id", excludeForwardId);
}
long tunnelForwardCount = this.count(tunnelQuery);
if (tunnelForwardCount >= userTunnel.getNum()) {
return R.err("该隧道转发数量已达上限,当前限制:" + userTunnel.getNum() + "个");
}
return R.ok();
}
/**
* 检查用户流量限制
*/
private R checkUserFlowLimits(Integer userId, Tunnel tunnel) {
User userInfo = userService.getById(userId);
if (userInfo.getExpTime() != null && userInfo.getExpTime() <= System.currentTimeMillis()) {
return R.err("当前账号已到期");
}
UserTunnel userTunnel = getUserTunnel(userId, tunnel.getId().intValue());
if (userTunnel == null) {
return R.err("你没有该隧道权限");
}
// 检查隧道权限到期时间
if (userTunnel.getExpTime() != null && userTunnel.getExpTime() <= System.currentTimeMillis()) {
return R.err("该隧道权限已到期,无法恢复服务");
}
// 检查用户总流量限制
if (userInfo.getFlow() * BYTES_TO_GB <= userInfo.getInFlow() + userInfo.getOutFlow()) {
return R.err("用户总流量已用完,无法恢复服务");
}
// 检查隧道流量限制
long tunnelFlow = (tunnel.getFlow() == FLOW_TYPE_UPLOAD_ONLY) ?
userTunnel.getOutFlow() :
userTunnel.getInFlow() + userTunnel.getOutFlow();
if (userTunnel.getFlow() * BYTES_TO_GB <= tunnelFlow) {
return R.err("该隧道流量已用完,无法恢复服务");
}
return R.ok();
}
/**
* 分配端口
*/
private PortAllocation allocatePorts(Tunnel tunnel) {
Integer inPort = allocateInPort(tunnel);
if (inPort == null) {
return PortAllocation.error("隧道入口端口已满,无法分配新端口");
}
Integer outPort = null;
if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) {
outPort = allocateOutPort(tunnel);
if (outPort == null) {
return PortAllocation.error("隧道出口端口已满,无法分配新端口");
}
}
return PortAllocation.success(inPort, outPort);
}
/**
* 创建Forward实体对象
*/
private Forward createForwardEntity(ForwardDto forwardDto, UserInfo currentUser, PortAllocation portAllocation) {
Forward forward = new Forward();
forward.setStatus(FORWARD_STATUS_ACTIVE);
forward.setInPort(portAllocation.getInPort());
forward.setOutPort(portAllocation.getOutPort());
forward.setUserId(currentUser.getUserId());
forward.setUserName(currentUser.getUserName());
forward.setCreatedTime(System.currentTimeMillis());
forward.setUpdatedTime(System.currentTimeMillis());
BeanUtils.copyProperties(forwardDto, forward);
return forward;
}
/**
* 更新Forward实体对象
*/
private Forward updateForwardEntity(ForwardUpdateDto forwardUpdateDto, Forward existForward, Tunnel tunnel) {
Forward forward = new Forward();
BeanUtils.copyProperties(forwardUpdateDto, forward);
// 如果隧道ID发生变化,需要重新分配端口
if (!existForward.getTunnelId().equals(forwardUpdateDto.getTunnelId())) {
PortAllocation portAllocation = allocatePorts(tunnel);
forward.setInPort(portAllocation.getInPort());
forward.setOutPort(portAllocation.getOutPort());
} else {
// 隧道未变化,保持原端口
forward.setInPort(existForward.getInPort());
forward.setOutPort(existForward.getOutPort());
}
forward.setUpdatedTime(System.currentTimeMillis());
return forward;
}
/**
* 创建Gost服务
*/
private R createGostServices(Forward forward, Tunnel tunnel, Integer limiter) {
String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId());
Node inNode = nodeService.getNodeById(tunnel.getInNodeId());
// 隧道转发需要创建链和远程服务
if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) {
R chainResult = createChainService(inNode, serviceName, tunnel.getOutIp(), forward.getOutPort());
if (chainResult.getCode() != 0) {
updateForwardStatusToError(forward);
return chainResult;
}
R remoteResult = createRemoteService(tunnel.getOutNodeId().intValue(), serviceName, forward);
if (remoteResult.getCode() != 0) {
updateForwardStatusToError(forward);
return remoteResult;
}
}
// 创建主服务
R serviceResult = createMainService(inNode, serviceName, forward, limiter, tunnel.getType());
if (serviceResult.getCode() != 0) {
updateForwardStatusToError(forward);
return serviceResult;
}
return R.ok();
}
/**
* 更新Gost服务
*/
private R updateGostServices(Forward forward, Tunnel tunnel, Integer limiter) {
String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId());
Node inNode = nodeService.getNodeById(tunnel.getInNodeId());
// 隧道转发需要更新链和远程服务
if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) {
R chainResult = updateChainService(inNode, serviceName, tunnel.getOutIp(), forward.getOutPort());
if (chainResult.getCode() != 0) {
updateForwardStatusToError(forward);
return chainResult;
}
R remoteResult = updateRemoteService(tunnel.getOutNodeId().intValue(), serviceName, forward);
if (remoteResult.getCode() != 0) {
updateForwardStatusToError(forward);
return remoteResult;
}
}
// 更新主服务
R serviceResult = updateMainService(inNode, serviceName, forward, limiter, tunnel.getType());
if (serviceResult.getCode() != 0) {
updateForwardStatusToError(forward);
return serviceResult;
}
return R.ok();
}
/**
* 删除Gost服务
*/
private R deleteGostServices(Forward forward, Tunnel tunnel) {
String serviceName = buildServiceName(forward.getId(), forward.getUserId(), forward.getTunnelId());
Node inNode = nodeService.getNodeById(tunnel.getInNodeId());
// 删除主服务
GostDto serviceResult = GostUtil.DeleteService(inNode.getIp() + ":" + inNode.getPort(), serviceName, inNode.getSecret());
if (!isGostOperationSuccess(serviceResult)) {
return R.err(serviceResult.getMsg());
}
// 隧道转发需要删除链和远程服务
if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) {
GostDto chainResult = GostUtil.DeleteChains(inNode.getIp() + ":" + inNode.getPort(), serviceName, inNode.getSecret());
if (!isGostOperationSuccess(chainResult)) {
return R.err(chainResult.getMsg());
}
Node outNode = nodeService.getNodeById(tunnel.getOutNodeId());
GostDto remoteResult = GostUtil.DeleteRemoteService(outNode.getIp() + ":" + outNode.getPort(), serviceName, outNode.getSecret());
if (!isGostOperationSuccess(remoteResult)) {
return R.err(remoteResult.getMsg());
}
}
return R.ok();
}
/**
* 创建链服务
*/
private R createChainService(Node inNode, String serviceName, String outIp, Integer outPort) {
String remoteAddr = outIp + ":" + outPort;
GostDto result = GostUtil.AddChains(inNode.getIp() + ":" + inNode.getPort(), serviceName, remoteAddr, inNode.getSecret());
return isGostOperationSuccess(result) ? R.ok() : R.err(result.getMsg());
}
/**
* 创建远程服务
*/
private R createRemoteService(Integer outNodeId, String serviceName, Forward forward) {
Node outNode = nodeService.getNodeById(outNodeId.longValue());
GostDto result = GostUtil.AddRemoteService(outNode.getIp() + ":" + outNode.getPort(),
serviceName, forward.getOutPort(),
forward.getRemoteAddr(), outNode.getSecret());
return isGostOperationSuccess(result) ? R.ok() : R.err(result.getMsg());
}
/**
* 创建主服务
*/
private R createMainService(Node inNode, String serviceName, Forward forward, Integer limiter, Integer tunnelType) {
GostDto result = GostUtil.AddService(inNode.getIp() + ":" + inNode.getPort(), serviceName,
forward.getInPort(), limiter, forward.getRemoteAddr(),
inNode.getSecret(), tunnelType);
return isGostOperationSuccess(result) ? R.ok() : R.err(result.getMsg());
}
/**
* 更新链服务
*/
private R updateChainService(Node inNode, String serviceName, String outIp, Integer outPort) {
// 创建新链
String remoteAddr = outIp + ":" + outPort;
GostDto createResult = GostUtil.UpdateChains(inNode.getIp() + ":" + inNode.getPort(), serviceName, remoteAddr, inNode.getSecret());
if (createResult.getMsg().contains(GOST_NOT_FOUND_MSG)) {
createResult = GostUtil.AddChains(inNode.getIp() + ":" + inNode.getPort(), serviceName, remoteAddr, inNode.getSecret());
}
return isGostOperationSuccess(createResult) ? R.ok() : R.err(createResult.getMsg());
}
/**
* 更新远程服务
*/
private R updateRemoteService(Integer outNodeId, String serviceName, Forward forward) {
Node outNode = nodeService.getNodeById(outNodeId.longValue());
// 创建新远程服务
GostDto createResult = GostUtil.UpdateRemoteService(outNode.getIp() + ":" + outNode.getPort(),
serviceName, forward.getOutPort(),
forward.getRemoteAddr(), outNode.getSecret());
if (createResult.getMsg().contains(GOST_NOT_FOUND_MSG)) {
createResult = GostUtil.AddRemoteService(outNode.getIp() + ":" + outNode.getPort(),
serviceName, forward.getOutPort(),
forward.getRemoteAddr(), outNode.getSecret());
}
return isGostOperationSuccess(createResult) ? R.ok() : R.err(createResult.getMsg());
}
/**
* 更新主服务
*/
private R updateMainService(Node inNode, String serviceName, Forward forward, Integer limiter, Integer tunnelType) {
GostDto result = GostUtil.UpdateService(inNode.getIp() + ":" + inNode.getPort(), serviceName,
forward.getInPort(), limiter, forward.getRemoteAddr(),
inNode.getSecret(), tunnelType);
if (result.getMsg().contains(GOST_NOT_FOUND_MSG)) {
result = GostUtil.AddService(inNode.getIp() + ":" + inNode.getPort(), serviceName,
forward.getInPort(), limiter, forward.getRemoteAddr(),
inNode.getSecret(), tunnelType);
}
return isGostOperationSuccess(result) ? R.ok() : R.err(result.getMsg());
}
/**
* 更新转发状态为错误
*/
private void updateForwardStatusToError(Forward forward) {
forward.setStatus(FORWARD_STATUS_ERROR);
this.updateById(forward);
}
/**
* 检查是否有用户隧道权限
*/
private boolean hasUserTunnelPermission(Integer userId, Integer tunnelId) {
return getUserTunnel(userId, tunnelId) != null;
}
/**
* 获取用户隧道关系
*/
private UserTunnel getUserTunnel(Integer userId, Integer tunnelId) {
return userTunnelService.getOne(new QueryWrapper<UserTunnel>()
.eq("user_id", userId)
.eq("tunnel_id", tunnelId));
}
/**
* 检查隧道是否发生变化
*/
private boolean isTunnelChanged(Forward existForward, ForwardUpdateDto updateDto) {
return !existForward.getTunnelId().equals(updateDto.getTunnelId());
}
/**
* 归还用户转发配额
*/
private void returnUserForwardQuota(UserInfo currentUser) {
if (currentUser.getRoleId() != ADMIN_ROLE_ID) {
User user = userService.getById(currentUser.getUserId());
if (user != null) {
user.setNum(user.getNum() + 1);
user.setUpdatedTime(System.currentTimeMillis());
userService.updateById(user);
}
}
}
/**
* 检查Gost操作是否成功
*/
private boolean isGostOperationSuccess(GostDto gostResult) {
return Objects.equals(gostResult.getMsg(), GOST_SUCCESS_MSG);
}
/**
* 为隧道分配一个可用的入口端口
*/
private Integer allocateInPort(Tunnel tunnel) {
// 获取所有使用相同入口节点的隧道
List<Tunnel> tunnelsWithSameInNode = tunnelService.list(new QueryWrapper<Tunnel>().eq("in_node_id", tunnel.getInNodeId()));
Set<Long> tunnelIds = tunnelsWithSameInNode.stream()
.map(Tunnel::getId)
.collect(Collectors.toSet());
// 获取这些隧道的所有转发已使用的入口端口
List<Forward> usedForwards = this.list(new QueryWrapper<Forward>().in("tunnel_id", tunnelIds));
Set<Integer> usedInPorts = usedForwards.stream()
.map(Forward::getInPort)
.filter(port -> port != null)
.collect(Collectors.toSet());
// 在隧道端口范围内寻找未使用的端口
for (int port = tunnel.getInPortSta(); port <= tunnel.getInPortEnd(); port++) {
if (!usedInPorts.contains(port)) {
return port;
}
}
return null;
}
/**
* 为隧道分配一个可用的出口端口
*/
private Integer allocateOutPort(Tunnel tunnel) {
// 获取所有使用相同出口节点的隧道
List<Tunnel> tunnelsWithSameOutNode = tunnelService.list(new QueryWrapper<Tunnel>().eq("out_node_id", tunnel.getOutNodeId()));
Set<Long> tunnelIds = tunnelsWithSameOutNode.stream()
.map(Tunnel::getId)
.collect(Collectors.toSet());
// 获取这些隧道的所有转发已使用的出口端口
List<Forward> usedForwards = this.list(new QueryWrapper<Forward>().in("tunnel_id", tunnelIds));
Set<Integer> usedOutPorts = usedForwards.stream()
.map(Forward::getOutPort)
.filter(port -> port != null)
.collect(Collectors.toSet());
// 在隧道出口端口范围内寻找未使用的端口
for (int port = tunnel.getOutIpSta(); port <= tunnel.getOutIpEnd(); port++) {
if (!usedOutPorts.contains(port)) {
return port;
}
}
return null;
}
/**
* 构建服务名称,确保管理员和用户操作的一致性
*/
private String buildServiceName(Long forwardId, Integer userId, Integer tunnelId) {
// 根据userId和tunnelId查询UserTunnel获取正确的user_tunnel_id
UserTunnel userTunnel = userTunnelService.getOne(new QueryWrapper<UserTunnel>()
.eq("user_id", userId)
.eq("tunnel_id", tunnelId));
int userTunnelId = (userTunnel != null) ? userTunnel.getId() : 0;
return forwardId + "_" + userId + "_" + userTunnelId;
}
// ========== 内部数据类 ==========
/**
* 用户信息封装类
*/
@Data
private static class UserInfo {
private final Integer userId;
private final Integer roleId;
private final String userName;
}
/**
* 用户权限检查结果
*/
@Data
private static class UserPermissionResult {
private final boolean hasError;
private final String errorMessage;
private final Integer limiter;
private UserPermissionResult(boolean hasError, String errorMessage, Integer limiter) {
this.hasError = hasError;
this.errorMessage = errorMessage;
this.limiter = limiter;
}
public static UserPermissionResult success(Integer limiter) {
return new UserPermissionResult(false, null, limiter);
}
public static UserPermissionResult error(String errorMessage) {
return new UserPermissionResult(true, errorMessage, null);
}
}
/**
* 端口分配结果
*/
@Data
private static class PortAllocation {
private final boolean hasError;
private final String errorMessage;
private final Integer inPort;
private final Integer outPort;
private PortAllocation(boolean hasError, String errorMessage, Integer inPort, Integer outPort) {
this.hasError = hasError;
this.errorMessage = errorMessage;
this.inPort = inPort;
this.outPort = outPort;
}
public static PortAllocation success(Integer inPort, Integer outPort) {
return new PortAllocation(false, null, inPort, outPort);
}
public static PortAllocation error(String errorMessage) {
return new PortAllocation(true, errorMessage, null, null);
}
}
}
@@ -0,0 +1,311 @@
package com.admin.service.impl;
import cn.hutool.core.util.IdUtil;
import cn.hutool.core.util.StrUtil;
import com.admin.common.dto.NodeDto;
import com.admin.common.dto.NodeUpdateDto;
import com.admin.common.dto.PageDto;
import com.admin.common.lang.R;
import com.admin.entity.Node;
import com.admin.entity.Tunnel;
import com.admin.mapper.NodeMapper;
import com.admin.mapper.TunnelMapper;
import com.admin.service.NodeService;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
import org.springframework.beans.BeanUtils;
import org.springframework.stereotype.Service;
import javax.annotation.Resource;
import java.util.List;
import org.springframework.beans.factory.annotation.Value;
/**
* <p>
* 节点服务实现类
* 提供节点的增删改查功能,包括节点创建、更新、删除和查询操作
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Service
public class NodeServiceImpl extends ServiceImpl<NodeMapper, Node> implements NodeService {
// ========== 常量定义 ==========
/** 节点默认状态:启用 */
private static final int NODE_STATUS_ACTIVE = 1;
/** 成功响应消息 */
private static final String SUCCESS_CREATE_MSG = "节点创建成功";
private static final String SUCCESS_UPDATE_MSG = "节点更新成功";
private static final String SUCCESS_DELETE_MSG = "节点删除成功";
/** 错误响应消息 */
private static final String ERROR_CREATE_MSG = "节点创建失败";
private static final String ERROR_UPDATE_MSG = "节点更新失败";
private static final String ERROR_DELETE_MSG = "节点删除失败";
private static final String ERROR_NODE_NOT_FOUND = "节点不存在";
/** 隧道使用检查相关消息 */
private static final String ERROR_IN_NODE_IN_USE = "该节点还有 %d 个隧道作为入口节点在使用,请先删除相关隧道";
private static final String ERROR_OUT_NODE_IN_USE = "该节点还有 %d 个隧道作为出口节点在使用,请先删除相关隧道";
// ========== 依赖注入 ==========
@Resource
private TunnelMapper tunnelMapper;
@Value("${server-addr}")
private String serverAddr;
// ========== 公共接口实现 ==========
/**
* 创建新节点
*
* @param nodeDto 节点创建数据传输对象
* @return 创建结果响应
*/
@Override
public R createNode(NodeDto nodeDto) {
Node node = buildNewNode(nodeDto);
boolean result = this.save(node);
return result ? R.ok(SUCCESS_CREATE_MSG) : R.err(ERROR_CREATE_MSG);
}
/**
* 获取所有节点列表
* 注意:返回结果中会隐藏节点密钥信息
*
* @return 包含所有节点的响应对象
*/
@Override
public R getAllNodes() {
List<Node> nodeList = this.list();
hideNodeSecrets(nodeList);
return R.ok(nodeList);
}
/**
* 更新节点信息
*
* @param nodeUpdateDto 节点更新数据传输对象
* @return 更新结果响应
*/
@Override
public R updateNode(NodeUpdateDto nodeUpdateDto) {
// 1. 验证节点是否存在
if (!isNodeExists(nodeUpdateDto.getId())) {
return R.err(ERROR_NODE_NOT_FOUND);
}
// 2. 构建更新对象并执行更新
Node updateNode = buildUpdateNode(nodeUpdateDto);
boolean result = this.updateById(updateNode);
return result ? R.ok(SUCCESS_UPDATE_MSG) : R.err(ERROR_UPDATE_MSG);
}
/**
* 删除节点
* 删除前会检查是否有隧道正在使用该节点
*
* @param id 节点ID
* @return 删除结果响应
*/
@Override
public R deleteNode(Long id) {
// 1. 验证节点是否存在
if (!isNodeExists(id)) {
return R.err(ERROR_NODE_NOT_FOUND);
}
// 2. 检查节点使用情况
R usageCheckResult = checkNodeUsage(id);
if (usageCheckResult.getCode() != 0) {
return usageCheckResult;
}
// 3. 执行删除操作
boolean result = this.removeById(id);
return result ? R.ok(SUCCESS_DELETE_MSG) : R.err(ERROR_DELETE_MSG);
}
/**
* 根据ID获取节点信息
*
* @param id 节点ID
* @return 节点对象
* @throws RuntimeException 当节点不存在时抛出异常
*/
@Override
public Node getNodeById(Long id) {
Node node = this.getById(id);
if (node == null) {
throw new RuntimeException(ERROR_NODE_NOT_FOUND);
}
return node;
}
// ========== 私有辅助方法 ==========
/**
* 构建新节点对象
*
* @param nodeDto 节点创建DTO
* @return 构建完成的节点对象
*/
private Node buildNewNode(NodeDto nodeDto) {
Node node = new Node();
BeanUtils.copyProperties(nodeDto, node);
// 设置默认属性
node.setSecret(IdUtil.simpleUUID());
node.setStatus(NODE_STATUS_ACTIVE);
// 设置时间戳
long currentTime = System.currentTimeMillis();
node.setCreatedTime(currentTime);
node.setUpdatedTime(currentTime);
return node;
}
/**
* 构建节点更新对象
*
* @param nodeUpdateDto 节点更新DTO
* @return 构建完成的更新对象
*/
private Node buildUpdateNode(NodeUpdateDto nodeUpdateDto) {
Node node = new Node();
node.setId(nodeUpdateDto.getId());
node.setName(nodeUpdateDto.getName());
node.setPort(nodeUpdateDto.getPort());
node.setUpdatedTime(System.currentTimeMillis());
return node;
}
/**
* 隐藏节点列表中的密钥信息
*
* @param nodeList 节点列表
*/
private void hideNodeSecrets(List<Node> nodeList) {
nodeList.forEach(node -> node.setSecret(null));
}
/**
* 检查节点是否存在
*
* @param nodeId 节点ID
* @return 节点是否存在
*/
private boolean isNodeExists(Long nodeId) {
return this.getById(nodeId) != null;
}
/**
* 检查节点使用情况
* 验证是否有隧道正在使用该节点作为入口或出口节点
*
* @param nodeId 节点ID
* @return 检查结果响应
*/
private R checkNodeUsage(Long nodeId) {
// 检查入口节点使用情况
R inNodeCheckResult = checkInNodeUsage(nodeId);
if (inNodeCheckResult.getCode() != 0) {
return inNodeCheckResult;
}
// 检查出口节点使用情况
return checkOutNodeUsage(nodeId);
}
/**
* 检查节点作为入口节点的使用情况
*
* @param nodeId 节点ID
* @return 检查结果响应
*/
private R checkInNodeUsage(Long nodeId) {
QueryWrapper<Tunnel> query = new QueryWrapper<>();
query.eq("in_node_id", nodeId);
long tunnelCount = tunnelMapper.selectCount(query);
if (tunnelCount > 0) {
String errorMsg = String.format(ERROR_IN_NODE_IN_USE, tunnelCount);
return R.err(errorMsg);
}
return R.ok();
}
/**
* 检查节点作为出口节点的使用情况
*
* @param nodeId 节点ID
* @return 检查结果响应
*/
private R checkOutNodeUsage(Long nodeId) {
QueryWrapper<Tunnel> query = new QueryWrapper<>();
query.eq("out_node_id", nodeId);
long tunnelCount = tunnelMapper.selectCount(query);
if (tunnelCount > 0) {
String errorMsg = String.format(ERROR_OUT_NODE_IN_USE, tunnelCount);
return R.err(errorMsg);
}
return R.ok();
}
/**
* 获取节点安装命令
* 根据节点信息生成对应的安装命令
*
* @param id 节点ID
* @return 包含安装命令的响应对象
*/
@Override
public R getInstallCommand(Long id) {
// 1. 验证节点是否存在
Node node = this.getById(id);
if (node == null) {
return R.err(ERROR_NODE_NOT_FOUND);
}
// 2. 构建安装命令
String installCommand = buildInstallCommand(node);
return R.ok(installCommand);
}
/**
* 构建节点安装命令
*
* @param node 节点对象
* @return 格式化的安装命令
*/
private String buildInstallCommand(Node node) {
StringBuilder command = new StringBuilder();
// 第一部分:下载安装脚本
command.append("curl -L https://raw.githubusercontent.com/bqlpfy/forward-panel/refs/heads/main/install.sh")
.append(" -o ./install.sh && chmod +x ./install.sh && ");
// 第二部分:执行安装脚本(去掉-u参数)
command.append("./install.sh")
.append(" -a ").append(serverAddr) // 服务器地址
.append(" -p ").append(node.getPort()) // 节点端口
.append(" -s ").append(node.getSecret()); // 节点密钥
return command.toString();
}
}
@@ -0,0 +1,406 @@
package com.admin.service.impl;
import com.admin.common.dto.GostDto;
import com.admin.common.dto.SpeedLimitDto;
import com.admin.common.dto.SpeedLimitUpdateDto;
import com.admin.common.lang.R;
import com.admin.common.utils.GostUtil;
import com.admin.entity.Node;
import com.admin.entity.SpeedLimit;
import com.admin.entity.Tunnel;
import com.admin.entity.UserTunnel;
import com.admin.mapper.SpeedLimitMapper;
import com.admin.service.NodeService;
import com.admin.service.SpeedLimitService;
import com.admin.service.TunnelService;
import com.admin.service.UserTunnelService;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
import lombok.Data;
import org.springframework.beans.BeanUtils;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.context.annotation.Lazy;
import org.springframework.stereotype.Service;
import java.math.BigDecimal;
import java.math.RoundingMode;
import java.util.List;
import java.util.Objects;
import java.util.UUID;
/**
* <p>
* 限速规则服务实现类
* 提供限速规则的增删改查功能,包括与Gost服务的集成
* 支持限速器的创建、更新、删除和查询操作
* </p>
*
* @author QAQ
* @since 2025-06-04
*/
@Service
public class SpeedLimitServiceImpl extends ServiceImpl<SpeedLimitMapper, SpeedLimit> implements SpeedLimitService {
// ========== 常量定义 ==========
/** Gost操作成功响应消息 */
private static final String GOST_SUCCESS_MSG = "OK";
/** Gost未找到资源响应消息 */
private static final String GOST_NOT_FOUND_MSG = "not found";
/** 限速规则状态 */
private static final int SPEED_LIMIT_ACTIVE_STATUS = 1;
private static final int SPEED_LIMIT_INACTIVE_STATUS = 0;
/** 速度转换比率:比特到字节 */
private static final double BITS_TO_BYTES_RATIO = 8.0;
/** 成功响应消息 */
private static final String SUCCESS_UPDATE_MSG = "限速规则更新成功";
private static final String SUCCESS_DELETE_MSG = "限速规则删除成功";
/** 错误响应消息 */
private static final String ERROR_CREATE_MSG = "限速规则创建失败";
private static final String ERROR_UPDATE_MSG = "限速规则更新失败";
private static final String ERROR_DELETE_MSG = "限速规则删除失败";
private static final String ERROR_SPEED_LIMIT_NOT_FOUND = "限速规则不存在";
private static final String ERROR_TUNNEL_NOT_FOUND = "指定的隧道不存在";
private static final String ERROR_TUNNEL_NOT_EXISTS = "隧道不存在";
private static final String ERROR_TUNNEL_NAME_MISMATCH = "隧道名称与隧道ID不匹配";
private static final String ERROR_SPEED_LIMIT_IN_USE = "该限速规则还有用户在使用 请先取消分配";
// ========== 依赖注入 ==========
@Autowired
private TunnelService tunnelService;
@Autowired
private NodeService nodeService;
@Autowired
private UserTunnelService userTunnelService;
@Autowired
@Lazy
private SpeedLimitService speedLimitService;
// ========== 公共接口实现 ==========
/**
* 创建限速规则
*
* @param speedLimitDto 限速规则创建数据传输对象
* @return 创建结果响应
*/
@Override
public R createSpeedLimit(SpeedLimitDto speedLimitDto) {
// 1. 验证隧道
TunnelValidationResult tunnelValidation = validateTunnelWithResult(speedLimitDto.getTunnelId(), speedLimitDto.getTunnelName());
if (tunnelValidation.isHasError()) {
return R.err(tunnelValidation.getErrorMessage());
}
// 2. 创建限速规则实体
SpeedLimit speedLimit = createSpeedLimitEntity(speedLimitDto);
if (!this.save(speedLimit)) {
return R.err(ERROR_CREATE_MSG);
}
// 3. 调用Gost API添加限速器
R gostResult = addGostLimiter(speedLimit, tunnelValidation.getTunnel());
if (gostResult.getCode() != 0) {
handleGostOperationFailure(speedLimit);
return gostResult;
}
return R.ok();
}
/**
* 获取所有限速规则
*
* @return 包含所有限速规则的响应对象
*/
@Override
public R getAllSpeedLimits() {
List<SpeedLimit> speedLimits = this.list();
return R.ok(speedLimits);
}
/**
* 更新限速规则
*
* @param speedLimitUpdateDto 限速规则更新数据传输对象
* @return 更新结果响应
*/
@Override
public R updateSpeedLimit(SpeedLimitUpdateDto speedLimitUpdateDto) {
// 1. 验证限速规则是否存在
SpeedLimit speedLimit = this.getById(speedLimitUpdateDto.getId());
if (speedLimit == null) {
return R.err(ERROR_SPEED_LIMIT_NOT_FOUND);
}
// 2. 验证隧道
TunnelValidationResult tunnelValidation = validateTunnelWithResult(speedLimitUpdateDto.getTunnelId(), speedLimitUpdateDto.getTunnelName());
if (tunnelValidation.isHasError()) {
return R.err(tunnelValidation.getErrorMessage());
}
// 3. 更新限速规则数据
updateSpeedLimitEntity(speedLimitUpdateDto, speedLimit);
// 4. 调用Gost API更新限速器
R gostResult = updateGostLimiter(speedLimit, tunnelValidation.getTunnel());
if (gostResult.getCode() != 0) {
return gostResult;
}
// 5. 保存更新
boolean result = this.updateById(speedLimit);
return result ? R.ok(SUCCESS_UPDATE_MSG) : R.err(ERROR_UPDATE_MSG);
}
/**
* 删除限速规则
* 删除前会检查是否有用户正在使用该限速规则
*
* @param id 限速规则ID
* @return 删除结果响应
*/
@Override
public R deleteSpeedLimit(Long id) {
// 1. 验证限速规则是否存在
SpeedLimit speedLimit = this.getById(id);
if (speedLimit == null) {
return R.err(ERROR_SPEED_LIMIT_NOT_FOUND);
}
// 2. 检查使用情况
R usageCheckResult = checkSpeedLimitUsage(id);
if (usageCheckResult.getCode() != 0) {
return usageCheckResult;
}
// 3. 获取隧道信息
Tunnel tunnel = tunnelService.getById(speedLimit.getTunnelId());
if (tunnel == null) {
return R.err(ERROR_TUNNEL_NOT_EXISTS);
}
// 4. 调用Gost API删除限速器
R gostResult = deleteGostLimiter(id, tunnel);
if (gostResult.getCode() != 0) {
return gostResult;
}
// 5. 删除限速规则
boolean result = this.removeById(id);
return result ? R.ok(SUCCESS_DELETE_MSG) : R.err(ERROR_DELETE_MSG);
}
// ========== 私有辅助方法 ==========
/**
* 验证隧道是否存在且名称匹配(返回详细结果)
*
* @param tunnelId 隧道ID
* @param tunnelName 隧道名称
* @return 隧道验证结果
*/
private TunnelValidationResult validateTunnelWithResult(Long tunnelId, String tunnelName) {
Tunnel tunnel = tunnelService.getById(tunnelId);
if (tunnel == null) {
return TunnelValidationResult.error(ERROR_TUNNEL_NOT_FOUND);
}
if (!tunnel.getName().equals(tunnelName)) {
return TunnelValidationResult.error(ERROR_TUNNEL_NAME_MISMATCH);
}
return TunnelValidationResult.success(tunnel);
}
/**
* 验证隧道是否存在且名称匹配(兼容原有方法)
*
* @param tunnelId 隧道ID
* @param tunnelName 隧道名称
* @return 验证结果响应
*/
private R validateTunnel(Long tunnelId, String tunnelName) {
TunnelValidationResult result = validateTunnelWithResult(tunnelId, tunnelName);
return result.isHasError() ? R.err(result.getErrorMessage()) : R.ok(result.getTunnel());
}
/**
* 创建限速规则实体对象
*
* @param speedLimitDto 限速规则创建DTO
* @return 构建完成的限速规则对象
*/
private SpeedLimit createSpeedLimitEntity(SpeedLimitDto speedLimitDto) {
SpeedLimit speedLimit = new SpeedLimit();
BeanUtils.copyProperties(speedLimitDto, speedLimit);
// 设置默认属性
long currentTime = System.currentTimeMillis();
speedLimit.setCreatedTime(currentTime);
speedLimit.setUpdatedTime(currentTime);
speedLimit.setStatus(SPEED_LIMIT_ACTIVE_STATUS);
return speedLimit;
}
/**
* 更新限速规则实体对象
*
* @param speedLimitUpdateDto 限速规则更新DTO
* @param speedLimit 待更新的限速规则对象
*/
private void updateSpeedLimitEntity(SpeedLimitUpdateDto speedLimitUpdateDto, SpeedLimit speedLimit) {
BeanUtils.copyProperties(speedLimitUpdateDto, speedLimit);
speedLimit.setUpdatedTime(System.currentTimeMillis());
}
/**
* 检查限速规则使用情况
*
* @param speedLimitId 限速规则ID
* @return 检查结果响应
*/
private R checkSpeedLimitUsage(Long speedLimitId) {
int userCount = userTunnelService.count(new QueryWrapper<UserTunnel>().eq("speed_id", speedLimitId));
if (userCount != 0) {
return R.err(ERROR_SPEED_LIMIT_IN_USE);
}
return R.ok();
}
/**
* 添加Gost限速器
*
* @param speedLimit 限速规则对象
* @param tunnel 隧道对象
* @return 操作结果响应
*/
private R addGostLimiter(SpeedLimit speedLimit, Tunnel tunnel) {
String speedInMBps = convertBitsToMBps(speedLimit.getSpeed());
Node node = nodeService.getNodeById(tunnel.getInNodeId());
GostDto gostResult = GostUtil.AddLimiters(
buildNodeAddress(node),
speedLimit.getId(),
speedInMBps,
node.getSecret()
);
return isGostOperationSuccess(gostResult) ? R.ok() : R.err(gostResult.getMsg());
}
/**
* 更新Gost限速器
*
* @param speedLimit 限速规则对象
* @param tunnel 隧道对象
* @return 操作结果响应
*/
private R updateGostLimiter(SpeedLimit speedLimit, Tunnel tunnel) {
String speedInMBps = convertBitsToMBps(speedLimit.getSpeed());
Node node = nodeService.getNodeById(tunnel.getInNodeId());
String nodeAddress = buildNodeAddress(node);
// 尝试更新限速器
GostDto gostResult = GostUtil.UpdateLimiters(nodeAddress, speedLimit.getId(), speedInMBps, node.getSecret());
// 如果限速器不存在,则创建新的
if (gostResult.getMsg().contains(GOST_NOT_FOUND_MSG)) {
gostResult = GostUtil.AddLimiters(nodeAddress, speedLimit.getId(), speedInMBps, node.getSecret());
}
return isGostOperationSuccess(gostResult) ? R.ok() : R.err(gostResult.getMsg());
}
/**
* 删除Gost限速器
*
* @param speedLimitId 限速规则ID
* @param tunnel 隧道对象
* @return 操作结果响应
*/
private R deleteGostLimiter(Long speedLimitId, Tunnel tunnel) {
Node node = nodeService.getNodeById(tunnel.getInNodeId());
GostDto gostResult = GostUtil.DeleteLimiters(buildNodeAddress(node), speedLimitId, node.getSecret());
return isGostOperationSuccess(gostResult) ? R.ok() : R.err(gostResult.getMsg());
}
/**
* 处理Gost操作失败的情况
*
* @param speedLimit 限速规则对象
*/
private void handleGostOperationFailure(SpeedLimit speedLimit) {
speedLimit.setStatus(SPEED_LIMIT_INACTIVE_STATUS);
speedLimitService.updateById(speedLimit);
}
/**
* 构建节点地址
*
* @param node 节点对象
* @return 节点地址字符串
*/
private String buildNodeAddress(Node node) {
return node.getIp() + ":" + node.getPort();
}
/**
* 将比特率转换为兆字节每秒
*
* @param speedInBits 比特率速度
* @return 兆字节每秒字符串
*/
private String convertBitsToMBps(Integer speedInBits) {
double mbs = speedInBits / BITS_TO_BYTES_RATIO;
BigDecimal bd = new BigDecimal(mbs).setScale(1, RoundingMode.HALF_UP);
return bd.doubleValue() + "";
}
/**
* 检查Gost操作是否成功
*
* @param gostResult Gost操作结果
* @return 是否成功
*/
private boolean isGostOperationSuccess(GostDto gostResult) {
return Objects.equals(gostResult.getMsg(), GOST_SUCCESS_MSG);
}
// ========== 内部数据类 ==========
/**
* 隧道验证结果封装类
*/
@Data
private static class TunnelValidationResult {
private final boolean hasError;
private final String errorMessage;
private final Tunnel tunnel;
private TunnelValidationResult(boolean hasError, String errorMessage, Tunnel tunnel) {
this.hasError = hasError;
this.errorMessage = errorMessage;
this.tunnel = tunnel;
}
public static TunnelValidationResult success(Tunnel tunnel) {
return new TunnelValidationResult(false, null, tunnel);
}
public static TunnelValidationResult error(String errorMessage) {
return new TunnelValidationResult(true, errorMessage, null);
}
}
}
@@ -0,0 +1,509 @@
package com.admin.service.impl;
import cn.hutool.core.util.StrUtil;
import com.admin.common.dto.TunnelDto;
import com.admin.common.dto.TunnelListDto;
import com.admin.common.lang.R;
import com.admin.common.utils.JwtUtil;
import com.admin.entity.Forward;
import com.admin.entity.Node;
import com.admin.entity.Tunnel;
import com.admin.entity.User;
import com.admin.entity.UserTunnel;
import com.admin.mapper.TunnelMapper;
import com.admin.mapper.UserTunnelMapper;
import com.admin.service.ForwardService;
import com.admin.service.NodeService;
import com.admin.service.TunnelService;
import com.admin.service.UserTunnelService;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
import lombok.Data;
import org.springframework.beans.BeanUtils;
import org.springframework.stereotype.Service;
import javax.annotation.Resource;
import java.util.List;
import java.util.stream.Collectors;
/**
* <p>
* 隧道服务实现类
* 提供隧道的增删改查功能,包括隧道创建、删除和用户权限管理
* 支持端口转发和隧道转发两种模式
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Service
public class TunnelServiceImpl extends ServiceImpl<TunnelMapper, Tunnel> implements TunnelService {
// ========== 常量定义 ==========
/** 隧道类型常量 */
private static final int TUNNEL_TYPE_PORT_FORWARD = 1; // 端口转发
private static final int TUNNEL_TYPE_TUNNEL_FORWARD = 2; // 隧道转发
/** 隧道状态常量 */
private static final int TUNNEL_STATUS_ACTIVE = 1; // 启用状态
/** 用户角色常量 */
private static final int ADMIN_ROLE_ID = 0; // 管理员角色ID
/** 成功响应消息 */
private static final String SUCCESS_CREATE_MSG = "隧道创建成功";
private static final String SUCCESS_DELETE_MSG = "隧道删除成功";
/** 错误响应消息 */
private static final String ERROR_CREATE_MSG = "隧道创建失败";
private static final String ERROR_DELETE_MSG = "隧道删除失败";
private static final String ERROR_TUNNEL_NOT_FOUND = "隧道不存在";
private static final String ERROR_TUNNEL_NAME_EXISTS = "隧道名称已存在";
private static final String ERROR_IN_NODE_NOT_FOUND = "入口节点不存在";
private static final String ERROR_OUT_NODE_NOT_FOUND = "出口节点不存在";
private static final String ERROR_OUT_NODE_REQUIRED = "出口节点不能为空";
private static final String ERROR_OUT_PORT_REQUIRED = "出口端口不能为空";
private static final String ERROR_SAME_NODE_NOT_ALLOWED = "隧道转发模式下,入口和出口不能是同一个节点";
private static final String ERROR_IN_PORT_RANGE_INVALID = "入口端口开始不能大于结束端口";
private static final String ERROR_OUT_PORT_RANGE_INVALID = "出口端口开始不能大于结束端口";
private static final String ERROR_NO_AVAILABLE_TUNNELS = "暂无可用隧道";
/** 使用检查相关消息 */
private static final String ERROR_FORWARDS_IN_USE = "该隧道还有 %d 个转发在使用,请先删除相关转发";
private static final String ERROR_USER_PERMISSIONS_IN_USE = "该隧道还有 %d 个用户权限关联,请先取消用户权限分配";
// ========== 依赖注入 ==========
@Resource
UserTunnelMapper userTunnelMapper;
@Resource
NodeService nodeService;
@Resource
ForwardService forwardService;
@Resource
UserTunnelService userTunnelService;
// ========== 公共接口实现 ==========
/**
* 创建隧道
* 支持端口转发和隧道转发两种模式
*
* @param tunnelDto 隧道创建数据传输对象
* @return 创建结果响应
*/
@Override
public R createTunnel(TunnelDto tunnelDto) {
// 1. 验证隧道名称唯一性
R nameValidationResult = validateTunnelNameUniqueness(tunnelDto.getName());
if (nameValidationResult.getCode() != 0) {
return nameValidationResult;
}
// 2. 验证入口节点和端口
NodeValidationResult inNodeValidation = validateInNode(tunnelDto);
if (inNodeValidation.isHasError()) {
return R.err(inNodeValidation.getErrorMessage());
}
// 3. 构建隧道实体
Tunnel tunnel = buildTunnelEntity(tunnelDto, inNodeValidation.getNode());
// 4. 根据隧道类型设置出口参数
R outNodeSetupResult = setupOutNodeParameters(tunnel, tunnelDto);
if (outNodeSetupResult.getCode() != 0) {
return outNodeSetupResult;
}
// 5. 设置默认属性并保存
setDefaultTunnelProperties(tunnel);
boolean result = this.save(tunnel);
return result ? R.ok(SUCCESS_CREATE_MSG) : R.err(ERROR_CREATE_MSG);
}
/**
* 获取所有隧道列表
*
* @return 包含所有隧道的响应对象
*/
@Override
public R getAllTunnels() {
List<Tunnel> tunnelList = this.list();
return R.ok(tunnelList);
}
/**
* 删除隧道
* 删除前会检查是否有转发或用户权限在使用该隧道
*
* @param id 隧道ID
* @return 删除结果响应
*/
@Override
public R deleteTunnel(Long id) {
// 1. 验证隧道是否存在
if (!isTunnelExists(id)) {
return R.err(ERROR_TUNNEL_NOT_FOUND);
}
// 2. 检查隧道使用情况
R usageCheckResult = checkTunnelUsage(id);
if (usageCheckResult.getCode() != 0) {
return usageCheckResult;
}
// 3. 执行删除操作
boolean result = this.removeById(id);
return result ? R.ok(SUCCESS_DELETE_MSG) : R.err(ERROR_DELETE_MSG);
}
/**
* 获取用户可用的隧道列表
* 管理员可以看到所有启用的隧道,普通用户只能看到有权限的启用隧道
*
* @return 用户可用隧道列表响应
*/
@Override
public R userTunnel() {
UserInfo currentUser = getCurrentUserInfo();
// 根据用户角色获取隧道列表
List<Tunnel> tunnelEntities = getUserAccessibleTunnels(currentUser);
// 转换为DTO并返回
List<TunnelListDto> tunnelDtos = convertToTunnelListDtos(tunnelEntities);
return R.ok(tunnelDtos);
}
// ========== 私有辅助方法 ==========
/**
* 获取当前用户信息
*
* @return 用户信息对象
*/
private UserInfo getCurrentUserInfo() {
Integer roleId = JwtUtil.getRoleIdFromToken();
Integer userId = JwtUtil.getUserIdFromToken();
return new UserInfo(userId, roleId);
}
/**
* 验证隧道名称唯一性
*
* @param tunnelName 隧道名称
* @return 验证结果响应
*/
private R validateTunnelNameUniqueness(String tunnelName) {
Tunnel existTunnel = this.getOne(new QueryWrapper<Tunnel>().eq("name", tunnelName));
if (existTunnel != null) {
return R.err(ERROR_TUNNEL_NAME_EXISTS);
}
return R.ok();
}
/**
* 验证入口节点和端口
*
* @param tunnelDto 隧道创建DTO
* @return 节点验证结果
*/
private NodeValidationResult validateInNode(TunnelDto tunnelDto) {
// 验证入口节点是否存在
Node inNode = nodeService.getById(tunnelDto.getInNodeId());
if (inNode == null) {
return NodeValidationResult.error(ERROR_IN_NODE_NOT_FOUND);
}
// 验证入口端口范围
if (tunnelDto.getInPortSta() > tunnelDto.getInPortEnd()) {
return NodeValidationResult.error(ERROR_IN_PORT_RANGE_INVALID);
}
return NodeValidationResult.success(inNode);
}
/**
* 构建隧道实体对象
*
* @param tunnelDto 隧道创建DTO
* @param inNode 入口节点
* @return 构建完成的隧道对象
*/
private Tunnel buildTunnelEntity(TunnelDto tunnelDto, Node inNode) {
Tunnel tunnel = new Tunnel();
BeanUtils.copyProperties(tunnelDto, tunnel);
// 设置入口节点信息
tunnel.setInNodeId(tunnelDto.getInNodeId());
tunnel.setInIp(inNode.getIp());
// 设置流量计算类型
tunnel.setFlow(tunnelDto.getFlow());
return tunnel;
}
/**
* 设置出口节点参数
*
* @param tunnel 隧道对象
* @param tunnelDto 隧道创建DTO
* @return 设置结果响应
*/
private R setupOutNodeParameters(Tunnel tunnel, TunnelDto tunnelDto) {
if (tunnelDto.getType() == TUNNEL_TYPE_PORT_FORWARD) {
// 端口转发:出口参数使用入口参数
return setupPortForwardOutParameters(tunnel, tunnelDto);
} else {
// 隧道转发:需要验证出口参数
return setupTunnelForwardOutParameters(tunnel, tunnelDto);
}
}
/**
* 设置端口转发的出口参数
*
* @param tunnel 隧道对象
* @param tunnelDto 隧道创建DTO
* @return 设置结果响应
*/
private R setupPortForwardOutParameters(Tunnel tunnel, TunnelDto tunnelDto) {
tunnel.setOutNodeId(tunnelDto.getInNodeId());
tunnel.setOutIp(tunnel.getInIp());
tunnel.setOutIpSta(tunnelDto.getInPortSta());
tunnel.setOutIpEnd(tunnelDto.getInPortEnd());
return R.ok();
}
/**
* 设置隧道转发的出口参数
*
* @param tunnel 隧道对象
* @param tunnelDto 隧道创建DTO
* @return 设置结果响应
*/
private R setupTunnelForwardOutParameters(Tunnel tunnel, TunnelDto tunnelDto) {
// 验证出口节点不能为空
if (tunnelDto.getOutNodeId() == null) {
return R.err(ERROR_OUT_NODE_REQUIRED);
}
// 验证入口和出口不能是同一个节点
if (tunnelDto.getInNodeId().equals(tunnelDto.getOutNodeId())) {
return R.err(ERROR_SAME_NODE_NOT_ALLOWED);
}
// 验证出口节点是否存在
Node outNode = nodeService.getById(tunnelDto.getOutNodeId());
if (outNode == null) {
return R.err(ERROR_OUT_NODE_NOT_FOUND);
}
// 验证出口端口参数
if (tunnelDto.getOutIpSta() == null || tunnelDto.getOutIpEnd() == null) {
return R.err(ERROR_OUT_PORT_REQUIRED);
}
if (tunnelDto.getOutIpSta() > tunnelDto.getOutIpEnd()) {
return R.err(ERROR_OUT_PORT_RANGE_INVALID);
}
// 设置出口参数
tunnel.setOutNodeId(tunnelDto.getOutNodeId());
tunnel.setOutIp(outNode.getIp());
return R.ok();
}
/**
* 设置隧道默认属性
*
* @param tunnel 隧道对象
*/
private void setDefaultTunnelProperties(Tunnel tunnel) {
tunnel.setStatus(TUNNEL_STATUS_ACTIVE);
long currentTime = System.currentTimeMillis();
tunnel.setCreatedTime(currentTime);
tunnel.setUpdatedTime(currentTime);
}
/**
* 检查隧道是否存在
*
* @param tunnelId 隧道ID
* @return 隧道是否存在
*/
private boolean isTunnelExists(Long tunnelId) {
return this.getById(tunnelId) != null;
}
/**
* 检查隧道使用情况
*
* @param tunnelId 隧道ID
* @return 检查结果响应
*/
private R checkTunnelUsage(Long tunnelId) {
// 检查转发使用情况
R forwardCheckResult = checkForwardUsage(tunnelId);
if (forwardCheckResult.getCode() != 0) {
return forwardCheckResult;
}
// 检查用户权限使用情况
return checkUserPermissionUsage(tunnelId);
}
/**
* 检查转发使用情况
*
* @param tunnelId 隧道ID
* @return 检查结果响应
*/
private R checkForwardUsage(Long tunnelId) {
QueryWrapper<Forward> forwardQuery = new QueryWrapper<>();
forwardQuery.eq("tunnel_id", tunnelId);
long forwardCount = forwardService.count(forwardQuery);
if (forwardCount > 0) {
String errorMsg = String.format(ERROR_FORWARDS_IN_USE, forwardCount);
return R.err(errorMsg);
}
return R.ok();
}
/**
* 检查用户权限使用情况
*
* @param tunnelId 隧道ID
* @return 检查结果响应
*/
private R checkUserPermissionUsage(Long tunnelId) {
QueryWrapper<UserTunnel> userTunnelQuery = new QueryWrapper<>();
userTunnelQuery.eq("tunnel_id", tunnelId);
long userTunnelCount = userTunnelService.count(userTunnelQuery);
if (userTunnelCount > 0) {
String errorMsg = String.format(ERROR_USER_PERMISSIONS_IN_USE, userTunnelCount);
return R.err(errorMsg);
}
return R.ok();
}
/**
* 获取用户可访问的隧道列表
*
* @param userInfo 用户信息
* @return 隧道列表
*/
private List<Tunnel> getUserAccessibleTunnels(UserInfo userInfo) {
if (userInfo.getRoleId() == ADMIN_ROLE_ID) {
// 管理员:获取所有启用状态的隧道
return getActiveTunnels();
} else {
// 普通用户:根据权限获取启用状态的隧道
return getUserAuthorizedTunnels(userInfo.getUserId());
}
}
/**
* 获取所有启用状态的隧道
*
* @return 启用状态的隧道列表
*/
private List<Tunnel> getActiveTunnels() {
return this.list(new QueryWrapper<Tunnel>().eq("status", TUNNEL_STATUS_ACTIVE));
}
/**
* 获取用户有权限的启用隧道
*
* @param userId 用户ID
* @return 用户有权限的隧道列表
*/
private List<Tunnel> getUserAuthorizedTunnels(Integer userId) {
List<UserTunnel> userTunnels = userTunnelMapper.selectList(
new QueryWrapper<UserTunnel>().eq("user_id", userId)
);
if (userTunnels.isEmpty()) {
return java.util.Collections.emptyList(); // 返回空列表
}
List<Integer> tunnelIds = userTunnels.stream()
.map(UserTunnel::getTunnelId)
.collect(Collectors.toList());
return this.list(new QueryWrapper<Tunnel>()
.in("id", tunnelIds)
.eq("status", TUNNEL_STATUS_ACTIVE));
}
/**
* 将隧道实体列表转换为DTO列表
*
* @param tunnelEntities 隧道实体列表
* @return 隧道DTO列表
*/
private List<TunnelListDto> convertToTunnelListDtos(List<Tunnel> tunnelEntities) {
return tunnelEntities.stream()
.map(this::convertToTunnelListDto)
.collect(Collectors.toList());
}
/**
* 将Tunnel实体转换为TunnelListDto
*
* @param tunnel 隧道实体
* @return 隧道列表DTO
*/
private TunnelListDto convertToTunnelListDto(Tunnel tunnel) {
TunnelListDto dto = new TunnelListDto();
dto.setId(tunnel.getId().intValue());
dto.setName(tunnel.getName());
return dto;
}
// ========== 内部数据类 ==========
/**
* 用户信息封装类
*/
@Data
private static class UserInfo {
private final Integer userId;
private final Integer roleId;
}
/**
* 节点验证结果封装类
*/
@Data
private static class NodeValidationResult {
private final boolean hasError;
private final String errorMessage;
private final Node node;
private NodeValidationResult(boolean hasError, String errorMessage, Node node) {
this.hasError = hasError;
this.errorMessage = errorMessage;
this.node = node;
}
public static NodeValidationResult success(Node node) {
return new NodeValidationResult(false, null, node);
}
public static NodeValidationResult error(String errorMessage) {
return new NodeValidationResult(true, errorMessage, null);
}
}
}
@@ -0,0 +1,793 @@
package com.admin.service.impl;
import cn.hutool.core.map.MapUtil;
import cn.hutool.core.util.StrUtil;
import com.admin.common.dto.ChangePasswordDto;
import com.admin.common.dto.LoginDto;
import com.admin.common.dto.PageDto;
import com.admin.common.dto.UserDto;
import com.admin.common.dto.UserUpdateDto;
import com.admin.common.dto.UserPackageDto;
import com.admin.common.dto.GostDto;
import com.admin.common.lang.R;
import com.admin.common.task.DelayQueueManager;
import com.admin.common.task.DelayTask;
import com.admin.common.task.TaskBase;
import com.admin.common.utils.GostUtil;
import com.admin.common.utils.JwtUtil;
import com.admin.common.utils.Md5Util;
import com.admin.entity.Forward;
import com.admin.entity.Node;
import com.admin.entity.Tunnel;
import com.admin.entity.User;
import com.admin.entity.UserTunnel;
import com.admin.mapper.ForwardMapper;
import com.admin.mapper.UserMapper;
import com.admin.mapper.UserTunnelMapper;
import com.admin.service.NodeService;
import com.admin.service.TunnelService;
import com.admin.service.UserService;
import com.admin.service.UserTunnelService;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
import lombok.Data;
import lombok.extern.slf4j.Slf4j;
import org.springframework.beans.BeanUtils;
import org.springframework.context.annotation.Lazy;
import org.springframework.stereotype.Service;
import javax.annotation.Resource;
import java.util.List;
/**
* <p>
* 用户服务实现类
* 提供用户的增删改查功能,包括用户登录、创建、更新、删除和套餐信息查询
* 支持用户关联数据的级联删除,包括转发和Gost服务的清理
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Slf4j
@Service
public class UserServiceImpl extends ServiceImpl<UserMapper, User> implements UserService {
// ========== 常量定义 ==========
/** 用户角色常量 */
private static final int ADMIN_ROLE_ID = 0; // 管理员角色ID
private static final int USER_ROLE_ID = 1; // 普通用户角色ID
private static final long ADMIN_USER_ID = 1L; // 管理员用户ID
/** 用户状态常量 */
private static final int USER_STATUS_ACTIVE = 1; // 用户启用状态
private static final int USER_STATUS_DISABLED = 0; // 用户停用状态
/** 隧道类型常量 */
private static final int TUNNEL_TYPE_TUNNEL_FORWARD = 2; // 隧道转发类型
/** 成功响应消息 */
private static final String SUCCESS_CREATE_MSG = "用户创建成功";
private static final String SUCCESS_UPDATE_MSG = "用户更新成功";
private static final String SUCCESS_DELETE_MSG = "用户及关联数据删除成功";
/** 错误响应消息 */
private static final String ERROR_LOGIN_CREDENTIALS = "账号或密码错误";
private static final String ERROR_ACCOUNT_DISABLED = "账户停用";
private static final String ERROR_CREATE_FAILED = "用户创建失败";
private static final String ERROR_UPDATE_FAILED = "用户更新失败";
private static final String ERROR_DELETE_FAILED = "用户删除失败";
private static final String ERROR_USER_NOT_FOUND = "用户不存在";
private static final String ERROR_USERNAME_EXISTS = "用户名已存在";
private static final String ERROR_USERNAME_TAKEN = "用户名已被其他用户使用";
private static final String ERROR_CANNOT_DELETE_ADMIN = "不能删除管理员用户";
private static final String ERROR_USER_NOT_LOGGED_IN = "用户未登录或token无效";
private static final String ERROR_GET_PACKAGE_INFO_FAILED = "获取套餐信息失败";
private static final String ERROR_CURRENT_PASSWORD_WRONG = "当前密码错误";
private static final String ERROR_PASSWORD_NOT_MATCH = "新密码和确认密码不匹配";
private static final String SUCCESS_PASSWORD_UPDATE = "密码修改成功";
/** 登录响应字段名 */
private static final String LOGIN_TOKEN_FIELD = "token";
private static final String LOGIN_NAME_FIELD = "name";
private static final String LOGIN_ROLE_ID_FIELD = "role_id";
// ========== 依赖注入 ==========
@Resource
private UserMapper userMapper;
@Resource
@Lazy
private ForwardMapper forwardMapper;
@Resource
private UserTunnelMapper userTunnelMapper;
@Resource
@Lazy
private TunnelService tunnelService;
@Resource
@Lazy
private NodeService nodeService;
@Resource
UserTunnelService userTunnelService;
@Resource
private DelayQueueManager delayQueueManager;
// ========== 公共接口实现 ==========
/**
* 用户登录
* 验证用户名密码,检查账户状态,生成JWT令牌
*
* @param loginDto 登录数据传输对象
* @return 登录结果响应,包含令牌和用户信息
*/
@Override
public R login(LoginDto loginDto) {
// 1. 验证用户凭据
LoginValidationResult validationResult = validateUserCredentials(loginDto);
if (validationResult.isHasError()) {
return R.err(validationResult.getErrorMessage());
}
// 2. 生成令牌并返回用户信息
User user = validationResult.getUser();
String token = JwtUtil.generateToken(user);
return R.ok(MapUtil.builder()
.put(LOGIN_TOKEN_FIELD, token)
.put(LOGIN_NAME_FIELD, user.getName())
.put(LOGIN_ROLE_ID_FIELD, user.getRoleId())
.build());
}
/**
* 创建用户
* 检查用户名唯一性,设置默认属性,加密密码
*
* @param userDto 用户创建数据传输对象
* @return 创建结果响应
*/
@Override
public R createUser(UserDto userDto) {
// 1. 验证用户名唯一性
R usernameValidationResult = validateUsernameUniqueness(userDto.getUser(), null);
if (usernameValidationResult.getCode() != 0) {
return usernameValidationResult;
}
// 2. 构建用户实体并保存
User user = buildNewUserEntity(userDto);
boolean result = this.save(user);
if (result) {
// 3. 添加到期时间延时任务
scheduleUserExpirationTask(user);
return R.ok(SUCCESS_CREATE_MSG);
} else {
return R.err(ERROR_CREATE_FAILED);
}
}
/**
* 获取所有用户(分页)
* 支持关键字搜索,排除管理员用户,清除密码信息
*
* @param pageDto 分页查询数据传输对象
* @return 分页用户列表响应
*/
@Override
public R getAllUsers(PageDto pageDto) {
// 1. 构建分页查询
Page<User> page = new Page<>(pageDto.getCurrent(), pageDto.getSize());
QueryWrapper<User> queryWrapper = buildUserQueryWrapper(pageDto);
// 2. 执行查询并处理结果
Page<User> userPage = this.page(page, queryWrapper);
clearUserPasswords(userPage.getRecords());
return R.ok(userPage);
}
/**
* 更新用户信息
* 验证用户存在性和用户名唯一性,处理密码加密
*
* @param userUpdateDto 用户更新数据传输对象
* @return 更新结果响应
*/
@Override
public R updateUser(UserUpdateDto userUpdateDto) {
// 1. 验证用户是否存在
if (!isUserExists(userUpdateDto.getId())) {
return R.err(ERROR_USER_NOT_FOUND);
}
// 2. 验证用户名唯一性
R usernameValidationResult = validateUsernameUniqueness(userUpdateDto.getUser(), userUpdateDto.getId());
if (usernameValidationResult.getCode() != 0) {
return usernameValidationResult;
}
// 4. 构建更新实体并保存
User user = buildUpdateUserEntity(userUpdateDto);
boolean result = this.updateById(user);
if (result) {
// 5. 处理到期时间延时任务
handleUserExpirationTaskUpdate(user);
return R.ok(SUCCESS_UPDATE_MSG);
} else {
return R.err(ERROR_UPDATE_FAILED);
}
}
/**
* 删除用户
* 级联删除用户相关的所有数据,包括转发、Gost服务和隧道权限
*
* @param id 用户ID
* @return 删除结果响应
*/
@Override
public R deleteUser(Long id) {
// 1. 验证删除条件
R deleteValidationResult = validateUserDeletion(id);
if (deleteValidationResult.getCode() != 0) {
return deleteValidationResult;
}
try {
// 2. 级联删除用户相关数据
deleteUserRelatedData(id);
delayQueueManager.remove("user_exp_" + id);
// 3. 删除用户
boolean result = this.removeById(id);
return result ? R.ok(SUCCESS_DELETE_MSG) : R.err(ERROR_DELETE_FAILED);
} catch (Exception e) {
e.printStackTrace();
return R.err("删除用户时发生错误:" + e.getMessage());
}
}
/**
* 获取用户套餐信息
* 包括用户基本信息、隧道权限详情和转发详情
*
* @return 用户套餐信息响应
*/
@Override
public R getUserPackageInfo() {
try {
// 1. 获取当前用户信息
CurrentUserInfo currentUser = getCurrentUserInfo();
if (currentUser.isHasError()) {
return R.err(currentUser.getErrorMessage());
}
// 2. 构建套餐信息
UserPackageDto packageDto = buildUserPackageDto(currentUser);
return R.ok(packageDto);
} catch (Exception e) {
e.printStackTrace();
return R.err(ERROR_GET_PACKAGE_INFO_FAILED);
}
}
/**
* 修改密码
* 验证当前密码、新密码确认、更新用户密码
*
* @param changePasswordDto 修改密码数据传输对象
* @return 修改结果响应
*/
@Override
public R updatePassword(ChangePasswordDto changePasswordDto) {
try {
// 1. 获取当前用户信息
CurrentUserInfo currentUser = getCurrentUserInfo();
if (currentUser.isHasError()) {
return R.err(currentUser.getErrorMessage());
}
// 2. 验证新密码和确认密码是否匹配
if (!changePasswordDto.getNewPassword().equals(changePasswordDto.getConfirmPassword())) {
return R.err(ERROR_PASSWORD_NOT_MATCH);
}
// 3. 验证当前密码是否正确
User user = currentUser.getUser();
String currentPasswordMd5 = Md5Util.md5(changePasswordDto.getCurrentPassword());
if (!user.getPwd().equals(currentPasswordMd5)) {
return R.err(ERROR_CURRENT_PASSWORD_WRONG);
}
// 4. 更新密码
User updateUser = new User();
updateUser.setId(user.getId());
updateUser.setPwd(Md5Util.md5(changePasswordDto.getNewPassword()));
updateUser.setUpdatedTime(System.currentTimeMillis());
boolean result = this.updateById(updateUser);
return result ? R.ok(SUCCESS_PASSWORD_UPDATE) : R.err(ERROR_UPDATE_FAILED);
} catch (Exception e) {
e.printStackTrace();
return R.err("修改密码时发生错误:" + e.getMessage());
}
}
// ========== 私有辅助方法 ==========
/**
* 验证用户登录凭据
*
* @param loginDto 登录数据传输对象
* @return 登录验证结果
*/
private LoginValidationResult validateUserCredentials(LoginDto loginDto) {
User user = this.getOne(new QueryWrapper<User>().eq("user", loginDto.getUsername()));
if (user == null) {
return LoginValidationResult.error(ERROR_LOGIN_CREDENTIALS);
}
if (!user.getPwd().equals(Md5Util.md5(loginDto.getPassword()))) {
return LoginValidationResult.error(ERROR_LOGIN_CREDENTIALS);
}
if (user.getStatus() == USER_STATUS_DISABLED) {
return LoginValidationResult.error(ERROR_ACCOUNT_DISABLED);
}
return LoginValidationResult.success(user);
}
/**
* 验证用户名唯一性
*
* @param username 用户名
* @param excludeUserId 排除的用户ID(用于更新时排除自己)
* @return 验证结果响应
*/
private R validateUsernameUniqueness(String username, Long excludeUserId) {
QueryWrapper<User> queryWrapper = new QueryWrapper<User>().eq("user", username);
if (excludeUserId != null) {
queryWrapper.ne("id", excludeUserId);
}
User existUser = this.getOne(queryWrapper);
if (existUser != null) {
String errorMsg = excludeUserId != null ? ERROR_USERNAME_TAKEN : ERROR_USERNAME_EXISTS;
return R.err(errorMsg);
}
return R.ok();
}
/**
* 构建新用户实体对象
*
* @param userDto 用户创建DTO
* @return 构建完成的用户对象
*/
private User buildNewUserEntity(UserDto userDto) {
User user = new User();
BeanUtils.copyProperties(userDto, user);
// 设置加密密码
user.setPwd(Md5Util.md5(userDto.getPwd()));
// 设置默认属性
user.setStatus(userDto.getStatus() != null ? userDto.getStatus() : USER_STATUS_ACTIVE);
user.setRoleId(USER_ROLE_ID);
// 设置时间戳
long currentTime = System.currentTimeMillis();
user.setCreatedTime(currentTime);
user.setUpdatedTime(currentTime);
return user;
}
/**
* 构建用户查询条件
*
* @param pageDto 分页查询DTO
* @return 查询条件包装器
*/
private QueryWrapper<User> buildUserQueryWrapper(PageDto pageDto) {
QueryWrapper<User> queryWrapper = new QueryWrapper<>();
// 关键字搜索
if (StrUtil.isNotBlank(pageDto.getKeyword())) {
queryWrapper.and(wrapper -> wrapper
.like("name", pageDto.getKeyword())
.or()
.like("user", pageDto.getKeyword())
);
}
// 排除管理员用户
queryWrapper.ne("id", ADMIN_USER_ID);
// 按更新时间降序排列
queryWrapper.orderByDesc("updated_time");
return queryWrapper;
}
/**
* 清除用户列表中的密码信息
*
* @param users 用户列表
*/
private void clearUserPasswords(List<User> users) {
users.forEach(user -> user.setPwd(null));
}
/**
* 检查用户是否存在
*
* @param userId 用户ID
* @return 用户是否存在
*/
private boolean isUserExists(Long userId) {
return this.getById(userId) != null;
}
/**
* 构建用户更新实体对象
*
* @param userUpdateDto 用户更新DTO
* @return 构建完成的更新对象
*/
private User buildUpdateUserEntity(UserUpdateDto userUpdateDto) {
User user = new User();
BeanUtils.copyProperties(userUpdateDto, user);
// 处理密码更新
if (StrUtil.isNotBlank(userUpdateDto.getPwd())) {
user.setPwd(Md5Util.md5(userUpdateDto.getPwd()));
} else {
user.setPwd(null); // 不更新密码字段
}
// 设置更新时间
user.setUpdatedTime(System.currentTimeMillis());
return user;
}
/**
* 验证用户删除条件
*
* @param userId 用户ID
* @return 验证结果响应
*/
private R validateUserDeletion(Long userId) {
User user = this.getById(userId);
if (user == null) {
return R.err(ERROR_USER_NOT_FOUND);
}
if (user.getRoleId() == ADMIN_ROLE_ID) {
return R.err(ERROR_CANNOT_DELETE_ADMIN);
}
return R.ok();
}
/**
* 删除用户相关的所有数据
*
* @param userId 用户ID
*/
private void deleteUserRelatedData(Long userId) {
// 1. 删除用户的所有转发和对应的Gost服务
deleteUserForwardsAndGostServices(userId);
// 2. 删除用户隧道权限
deleteUserTunnelPermissions(userId);
}
/**
* 删除用户转发和对应的Gost服务
*
* @param userId 用户ID
*/
private void deleteUserForwardsAndGostServices(Long userId) {
QueryWrapper<Forward> forwardQuery = new QueryWrapper<>();
forwardQuery.eq("user_id", userId);
List<Forward> userForwards = forwardMapper.selectList(forwardQuery);
for (Forward forward : userForwards) {
try {
// 删除Gost服务
deleteGostServicesForForward(forward, userId);
} catch (Exception e) {
// 记录错误但继续删除,避免因为Gost服务删除失败而阻断用户删除
System.err.println("删除用户转发对应的Gost服务失败,转发ID: " + forward.getId() + ", 错误: " + e.getMessage());
}
// 删除数据库中的转发记录
forwardMapper.deleteById(forward.getId());
}
}
/**
* 删除转发对应的Gost服务
*
* @param forward 转发对象
* @param userId 用户ID
*/
private void deleteGostServicesForForward(Forward forward, Long userId) {
Tunnel tunnel = tunnelService.getById(forward.getTunnelId());
if (tunnel == null) return;
Node inNode = nodeService.getNodeById(tunnel.getInNodeId());
if (inNode == null) return;
// 获取用户隧道关系
UserTunnel userTunnel = getUserTunnelRelation(userId, tunnel.getId());
if (userTunnel == null) return;
String serviceName = buildServiceName(forward.getId(), userId, userTunnel.getId());
// 删除主服务
GostUtil.DeleteService(buildNodeAddress(inNode), serviceName, inNode.getSecret());
// 如果是隧道转发,还需要删除链和远程服务
if (tunnel.getType() == TUNNEL_TYPE_TUNNEL_FORWARD) {
deleteGostTunnelForwardServices(tunnel, serviceName, inNode);
}
}
/**
* 删除隧道转发相关的Gost服务
*
* @param tunnel 隧道对象
* @param serviceName 服务名称
* @param inNode 入口节点
*/
private void deleteGostTunnelForwardServices(Tunnel tunnel, String serviceName, Node inNode) {
Node outNode = nodeService.getNodeById(tunnel.getOutNodeId());
if (outNode != null) {
GostUtil.DeleteChains(buildNodeAddress(inNode), serviceName, inNode.getSecret());
GostUtil.DeleteRemoteService(buildNodeAddress(outNode), serviceName, outNode.getSecret());
}
}
/**
* 获取用户隧道关系
*
* @param userId 用户ID
* @param tunnelId 隧道ID
* @return 用户隧道关系对象
*/
private UserTunnel getUserTunnelRelation(Long userId, Long tunnelId) {
return userTunnelService.getOne(new QueryWrapper<UserTunnel>()
.eq("user_id", userId)
.eq("tunnel_id", tunnelId));
}
/**
* 构建服务名称
*
* @param forwardId 转发ID
* @param userId 用户ID
* @param userTunnelId 用户隧道ID
* @return 服务名称
*/
private String buildServiceName(Long forwardId, Long userId, Integer userTunnelId) {
return forwardId + "_" + userId + "_" + userTunnelId;
}
/**
* 构建节点地址
*
* @param node 节点对象
* @return 节点地址字符串
*/
private String buildNodeAddress(Node node) {
return node.getIp() + ":" + node.getPort();
}
/**
* 删除用户隧道权限
*
* @param userId 用户ID
*/
private void deleteUserTunnelPermissions(Long userId) {
QueryWrapper<UserTunnel> userTunnelQuery = new QueryWrapper<>();
userTunnelQuery.eq("user_id", userId);
userTunnelMapper.delete(userTunnelQuery);
}
/**
* 获取当前用户信息
*
* @return 当前用户信息结果
*/
private CurrentUserInfo getCurrentUserInfo() {
Integer userId = JwtUtil.getUserIdFromToken();
Integer roleId = JwtUtil.getRoleIdFromToken();
if (userId == null) {
return CurrentUserInfo.error(ERROR_USER_NOT_LOGGED_IN);
}
User user = this.getById(userId);
if (user == null) {
return CurrentUserInfo.error(ERROR_USER_NOT_FOUND);
}
return CurrentUserInfo.success(user, roleId);
}
/**
* 构建用户套餐信息DTO
*
* @param currentUser 当前用户信息
* @return 用户套餐信息DTO
*/
private UserPackageDto buildUserPackageDto(CurrentUserInfo currentUser) {
User user = currentUser.getUser();
Integer roleId = currentUser.getRoleId();
// 1. 构造用户基本信息
UserPackageDto.UserInfoDto userInfo = buildUserInfoDto(user);
// 2. 获取隧道权限详情
List<UserPackageDto.UserTunnelDetailDto> tunnelPermissions = getTunnelPermissions(user.getId(), roleId);
// 3. 获取转发详情
List<UserPackageDto.UserForwardDetailDto> forwards = userMapper.getUserForwardDetails(user.getId().intValue());
// 4. 构造返回结果
UserPackageDto packageDto = new UserPackageDto();
packageDto.setUserInfo(userInfo);
packageDto.setTunnelPermissions(tunnelPermissions);
packageDto.setForwards(forwards);
return packageDto;
}
/**
* 构建用户基本信息DTO
*
* @param user 用户对象
* @return 用户基本信息DTO
*/
private UserPackageDto.UserInfoDto buildUserInfoDto(User user) {
UserPackageDto.UserInfoDto userInfo = new UserPackageDto.UserInfoDto();
userInfo.setId(user.getId());
userInfo.setName(user.getName());
userInfo.setUser(user.getUser());
userInfo.setStatus(user.getStatus());
userInfo.setFlow(user.getFlow());
userInfo.setInFlow(user.getInFlow());
userInfo.setOutFlow(user.getOutFlow());
userInfo.setNum(user.getNum());
userInfo.setExpTime(user.getExpTime());
userInfo.setFlowResetTime(user.getFlowResetTime());
userInfo.setCreatedTime(user.getCreatedTime());
userInfo.setUpdatedTime(user.getUpdatedTime());
return userInfo;
}
/**
* 获取隧道权限详情
*
* @param userId 用户ID
* @param roleId 角色ID
* @return 隧道权限详情列表
*/
private List<UserPackageDto.UserTunnelDetailDto> getTunnelPermissions(Long userId, Integer roleId) {
if (roleId != null && roleId == ADMIN_ROLE_ID) {
return userMapper.getAllTunnelsForAdmin();
} else {
return userMapper.getUserTunnelDetails(userId.intValue());
}
}
/**
* 安排用户到期延时任务
*
* @param user 用户对象
*/
private void scheduleUserExpirationTask(User user) {
// 取消已存在的延时任务(如果有)
delayQueueManager.remove("user_exp_" + user.getId());
// 创建新的延时任务
TaskBase taskBase = new TaskBase(user.getId().toString());
taskBase.setType("1"); // 账号到期延迟任务
long delayTime = user.getExpTime() - System.currentTimeMillis();
DelayTask delayTask = new DelayTask(taskBase, delayTime);
delayQueueManager.put(delayTask);
}
/**
* 处理用户到期时间更新的延时任务
*
* @param newUser 新用户信息
*/
private void handleUserExpirationTaskUpdate(User newUser) {
String taskId = "user_exp_" + newUser.getId();
// 先取消原有的延时任务
delayQueueManager.remove(taskId);
TaskBase taskBase = new TaskBase(newUser.getId().toString());
taskBase.setType("1"); // 账号到期延迟任务
long delayTime = newUser.getExpTime() - System.currentTimeMillis();
DelayTask delayTask = new DelayTask(taskBase, delayTime);
delayQueueManager.put(delayTask);
}
// ========== 内部数据类 ==========
/**
* 登录验证结果封装类
*/
@Data
private static class LoginValidationResult {
private final boolean hasError;
private final String errorMessage;
private final User user;
private LoginValidationResult(boolean hasError, String errorMessage, User user) {
this.hasError = hasError;
this.errorMessage = errorMessage;
this.user = user;
}
public static LoginValidationResult success(User user) {
return new LoginValidationResult(false, null, user);
}
public static LoginValidationResult error(String errorMessage) {
return new LoginValidationResult(true, errorMessage, null);
}
}
/**
* 当前用户信息封装类
*/
@Data
private static class CurrentUserInfo {
private final boolean hasError;
private final String errorMessage;
private final User user;
private final Integer roleId;
private CurrentUserInfo(boolean hasError, String errorMessage, User user, Integer roleId) {
this.hasError = hasError;
this.errorMessage = errorMessage;
this.user = user;
this.roleId = roleId;
}
public static CurrentUserInfo success(User user, Integer roleId) {
return new CurrentUserInfo(false, null, user, roleId);
}
public static CurrentUserInfo error(String errorMessage) {
return new CurrentUserInfo(true, errorMessage, null, null);
}
}
}
@@ -0,0 +1,450 @@
package com.admin.service.impl;
import com.admin.common.dto.UserTunnelDto;
import com.admin.common.dto.UserTunnelQueryDto;
import com.admin.common.dto.UserTunnelUpdateDto;
import com.admin.common.dto.UserTunnelWithDetailDto;
import com.admin.common.lang.R;
import com.admin.entity.UserTunnel;
import com.admin.mapper.TunnelMapper;
import com.admin.mapper.UserTunnelMapper;
import com.admin.service.TunnelService;
import com.admin.service.UserTunnelService;
import com.admin.service.ForwardService;
import com.admin.service.NodeService;
import com.admin.common.task.DelayQueueManager;
import com.admin.common.utils.GostUtil;
import com.admin.entity.Forward;
import com.admin.entity.Tunnel;
import com.admin.entity.Node;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper;
import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
import org.springframework.beans.BeanUtils;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.context.annotation.Lazy;
import org.springframework.stereotype.Service;
import javax.annotation.Resource;
import java.util.List;
import java.util.Map;
/**
* <p>
* 用户隧道权限服务实现类
* 提供用户隧道权限的分配、查询、更新和删除功能
* 支持流量限制、数量限制、过期时间和限速规则的管理
* </p>
*
* @author QAQ
* @since 2025-06-03
*/
@Service
public class UserTunnelServiceImpl extends ServiceImpl<UserTunnelMapper, UserTunnel> implements UserTunnelService {
// ========== 常量定义 ==========
/** 成功响应消息 */
private static final String SUCCESS_ASSIGN_MSG = "用户隧道权限分配成功";
private static final String SUCCESS_REMOVE_MSG = "用户隧道权限删除成功";
private static final String SUCCESS_UPDATE_FLOW_MSG = "用户隧道流量限制更新成功";
private static final String SUCCESS_UPDATE_MSG = "用户隧道权限更新成功";
/** 错误响应消息 */
private static final String ERROR_ASSIGN_FAILED = "用户隧道权限分配失败";
private static final String ERROR_PERMISSION_EXISTS = "该用户已拥有此隧道权限";
private static final String ERROR_PERMISSION_NOT_FOUND = "未找到对应的用户隧道权限记录";
private static final String ERROR_USER_TUNNEL_NOT_EXISTS = "用户隧道权限不存在";
private static final String ERROR_NOT_EXISTS = "不存在";
private static final String ERROR_UPDATE_FAILED = "用户隧道权限更新失败";
// ========== 依赖注入 ==========
@Autowired
private DelayQueueManager delayQueueManager;
@Autowired
@Lazy
private ForwardService forwardService;
@Autowired
@Lazy
private TunnelService tunnelService;
@Autowired
private NodeService nodeService;
// ========== 公共接口实现 ==========
/**
* 分配用户隧道权限
* 检查权限是否已存在,避免重复分配
*
* @param userTunnelDto 用户隧道权限分配数据传输对象
* @return 分配结果响应
*/
@Override
public R assignUserTunnel(UserTunnelDto userTunnelDto) {
// 1. 检查权限是否已存在
if (isUserTunnelPermissionExists(userTunnelDto.getUserId(), userTunnelDto.getTunnelId())) {
return R.err(ERROR_PERMISSION_EXISTS);
}
// 2. 创建用户隧道权限实体并保存
UserTunnel userTunnel = buildUserTunnelEntity(userTunnelDto);
// 设置默认状态为启用
userTunnel.setStatus(1);
boolean success = this.save(userTunnel);
if (success) {
// 3. 如果是启用状态且有到期时间,添加延迟任务
if (isEnabledAndHasExpTime(userTunnel)) {
try {
delayQueueManager.addUserTunnelExpirationTask(userTunnel);
} catch (Exception e) {
// 延迟任务添加失败不影响主业务逻辑
// 可以考虑记录日志或其他处理方式
}
}
return R.ok(SUCCESS_ASSIGN_MSG);
}
return R.err(ERROR_ASSIGN_FAILED);
}
/**
* 获取用户隧道权限列表
* 通过连表查询获取用户隧道权限及隧道详细信息
*
* @param queryDto 用户隧道权限查询数据传输对象
* @return 用户隧道权限详情列表响应
*/
@Override
public R getUserTunnelList(UserTunnelQueryDto queryDto) {
List<UserTunnelWithDetailDto> userTunnelDetails = getUserTunnelDetailsFromDatabase(queryDto.getUserId());
return R.ok(userTunnelDetails);
}
/**
* 删除用户隧道权限
*
* @param id 用户隧道权限ID
* @return 删除结果响应
*/
@Override
public R removeUserTunnel(Integer id) {
// 1. 获取用户隧道权限信息
UserTunnel userTunnel = this.getById(id);
if (userTunnel == null) {
return R.err(ERROR_PERMISSION_NOT_FOUND);
}
// 2. 删除该用户在该隧道下的所有转发
try {
removeUserTunnelForwards(userTunnel.getUserId(), userTunnel.getTunnelId());
} catch (Exception e) {
// 转发删除失败,记录日志但不阻止权限删除
}
// 3. 移除延迟任务
try {
delayQueueManager.removeUserTunnelExpirationTask(id);
} catch (Exception e) {
// 延迟任务移除失败不影响主业务逻辑
}
// 4. 删除用户隧道权限记录
boolean success = this.removeById(id);
return success ? R.ok(SUCCESS_REMOVE_MSG) : R.err(ERROR_PERMISSION_NOT_FOUND);
}
/**
* 更新用户隧道流量限制
*
* @param id 用户隧道权限ID
* @param flow 流量限制值
* @return 更新结果响应
*/
@Override
public R updateUserTunnelFlow(Integer id, Long flow) {
// 1. 验证用户隧道权限是否存在
UserTunnel userTunnel = this.getById(id);
if (userTunnel == null) {
return R.err(ERROR_NOT_EXISTS);
}
// 2. 更新流量限制并保存
userTunnel.setFlow(flow);
boolean success = this.updateById(userTunnel);
return success ? R.ok(SUCCESS_UPDATE_FLOW_MSG) : R.err(ERROR_PERMISSION_NOT_FOUND);
}
/**
* 更新用户隧道权限
* 支持更新流量限制、数量限制、流量重置时间、过期时间和限速规则
*
* @param updateDto 用户隧道权限更新数据传输对象
* @return 更新结果响应
*/
@Override
public R updateUserTunnel(UserTunnelUpdateDto updateDto) {
// 1. 验证用户隧道权限是否存在
UserTunnel existingUserTunnel = this.getById(updateDto.getId());
if (existingUserTunnel == null) {
return R.err(ERROR_USER_TUNNEL_NOT_EXISTS);
}
// 2. 更新用户隧道权限属性
updateUserTunnelProperties(existingUserTunnel, updateDto);
// 3. 保存更新
boolean success = this.updateById(existingUserTunnel);
if (success) {
// 4. 处理延迟任务更新
handleDelayTaskUpdate(existingUserTunnel, updateDto);
return R.ok(SUCCESS_UPDATE_MSG);
}
return R.err(ERROR_UPDATE_FAILED);
}
// ========== 私有辅助方法 ==========
/**
* 检查用户隧道权限是否已存在
*
* @param userId 用户ID
* @param tunnelId 隧道ID
* @return 权限是否已存在
*/
private boolean isUserTunnelPermissionExists(Integer userId, Integer tunnelId) {
QueryWrapper<UserTunnel> queryWrapper = new QueryWrapper<>();
queryWrapper.eq("user_id", userId).eq("tunnel_id", tunnelId);
UserTunnel existingUserTunnel = this.getOne(queryWrapper);
return existingUserTunnel != null;
}
/**
* 构建用户隧道权限实体对象
*
* @param userTunnelDto 用户隧道权限DTO
* @return 构建完成的用户隧道权限对象
*/
private UserTunnel buildUserTunnelEntity(UserTunnelDto userTunnelDto) {
UserTunnel userTunnel = new UserTunnel();
BeanUtils.copyProperties(userTunnelDto, userTunnel);
return userTunnel;
}
/**
* 从数据库获取用户隧道权限详情
*
* @param userId 用户ID
* @return 用户隧道权限详情列表
*/
private List<UserTunnelWithDetailDto> getUserTunnelDetailsFromDatabase(Integer userId) {
return this.baseMapper.getUserTunnelWithDetails(userId);
}
/**
* 更新用户隧道权限属性
*
* @param existingUserTunnel 现有的用户隧道权限对象
* @param updateDto 更新数据传输对象
*/
private void updateUserTunnelProperties(UserTunnel existingUserTunnel, UserTunnelUpdateDto updateDto) {
// 更新基本属性
existingUserTunnel.setFlow(updateDto.getFlow());
existingUserTunnel.setNum(updateDto.getNum());
// 更新可选属性(仅在非空时更新)
updateOptionalProperty(existingUserTunnel::setFlowResetTime, updateDto.getFlowResetTime());
updateOptionalProperty(existingUserTunnel::setExpTime, updateDto.getExpTime());
// 更新限速规则ID(允许设置为null,表示不限速)
existingUserTunnel.setSpeedId(updateDto.getSpeedId());
}
/**
* 更新可选属性(仅在值非空时更新)
*
* @param setter 属性设置方法
* @param value 属性值
* @param <T> 属性类型
*/
private <T> void updateOptionalProperty(java.util.function.Consumer<T> setter, T value) {
if (value != null) {
setter.accept(value);
}
}
/**
* 处理延迟任务更新
*
* @param userTunnel 更新后的用户隧道对象
* @param updateDto 更新数据传输对象
*/
private void handleDelayTaskUpdate(UserTunnel userTunnel, UserTunnelUpdateDto updateDto) {
try {
// 如果更新了到期时间,需要重新处理延迟任务
if (updateDto.getExpTime() != null) {
// 先移除旧的延迟任务
delayQueueManager.removeUserTunnelExpirationTask(userTunnel.getId());
// 如果是启用状态且有到期时间,添加新的延迟任务
if (isEnabledAndHasExpTime(userTunnel)) {
delayQueueManager.addUserTunnelExpirationTask(userTunnel);
}
}
} catch (Exception e) {
// 延迟任务处理失败不影响主业务逻辑
// 可以考虑记录日志
}
}
/**
* 删除用户在指定隧道下的所有转发
*
* @param userId 用户ID
* @param tunnelId 隧道ID
*/
private void removeUserTunnelForwards(Integer userId, Integer tunnelId) {
try {
// 查询该用户在该隧道下的所有转发
QueryWrapper<Forward> queryWrapper = new QueryWrapper<>();
queryWrapper.eq("user_id", userId).eq("tunnel_id", tunnelId);
List<Forward> userTunnelForwards = forwardService.list(queryWrapper);
if (!userTunnelForwards.isEmpty()) {
// 获取用户隧道权限信息,用于构建服务名称
UserTunnel userTunnel = getUserTunnelByUserAndTunnel(userId, tunnelId);
for (Forward forward : userTunnelForwards) {
try {
// 先调用GostUtil删除/停止服务
stopForwardService(forward, userId, userTunnel != null ? userTunnel.getId() : 0);
// 然后删除数据库记录
forwardService.removeById(forward.getId());
} catch (Exception e) {
// 单个转发删除失败,记录错误但继续处理其他转发
}
}
}
} catch (Exception e) {
// 删除转发失败,抛出异常让上层处理
throw new RuntimeException("删除用户隧道转发失败:" + e.getMessage(), e);
}
}
/**
* 删除转发服务(按创建的反向顺序删除:主服务 -> 远端服务 -> 转发链)
*
* @param forward 转发对象
* @param userId 用户ID
* @param userTunnelId 用户隧道ID
*/
private void stopForwardService(Forward forward, Integer userId, Integer userTunnelId) {
try {
Tunnel tunnel = tunnelService.getById(forward.getTunnelId());
if (tunnel == null) {
return;
}
Node inNode = nodeService.getById(tunnel.getInNodeId());
Node outNode = nodeService.getById(tunnel.getOutNodeId());
String serviceName = buildServiceName(forward.getId(), Long.valueOf(userId), userTunnelId);
// 1. 先删除主服务
if (inNode != null) {
String inNodeAddress = buildNodeAddress(inNode);
try {
GostUtil.DeleteService(inNodeAddress, serviceName, inNode.getSecret());
} catch (Exception e) {
// 主服务删除失败,记录但继续
}
}
// 2. 如果是隧道转发,删除远端服务
if (tunnel.getType() == 1 && outNode != null && !outNode.getId().equals(inNode != null ? inNode.getId() : null)) {
String outNodeAddress = buildNodeAddress(outNode);
try {
GostUtil.DeleteRemoteService(outNodeAddress, serviceName, outNode.getSecret());
} catch (Exception e) {
// 远端服务删除失败,记录但继续
}
}
// 3. 如果是隧道转发,最后删除转发链
if (tunnel.getType() == 1 && inNode != null) {
String inNodeAddress = buildNodeAddress(inNode);
try {
GostUtil.DeleteChains(inNodeAddress, serviceName, inNode.getSecret());
} catch (Exception e) {
// 转发链删除失败,记录但继续
}
}
} catch (Exception e) {
// 服务删除失败,记录错误
throw new RuntimeException("删除转发服务失败,转发ID:" + forward.getId() + ",错误:" + e.getMessage(), e);
}
}
/**
* 根据用户ID和隧道ID获取用户隧道权限
*
* @param userId 用户ID
* @param tunnelId 隧道ID
* @return 用户隧道权限对象
*/
private UserTunnel getUserTunnelByUserAndTunnel(Integer userId, Integer tunnelId) {
try {
QueryWrapper<UserTunnel> queryWrapper = new QueryWrapper<>();
queryWrapper.eq("user_id", userId).eq("tunnel_id", tunnelId);
return this.getOne(queryWrapper);
} catch (Exception e) {
return null;
}
}
/**
* 构建服务名称
*
* @param forwardId 转发ID
* @param userId 用户ID
* @param userTunnelId 用户隧道ID
* @return 服务名称
*/
private String buildServiceName(Long forwardId, Long userId, Integer userTunnelId) {
return forwardId + "_" + userId + "_" + userTunnelId;
}
/**
* 构建节点地址
*
* @param node 节点对象
* @return 节点地址字符串
*/
private String buildNodeAddress(Node node) {
return node.getIp() + ":" + node.getPort();
}
/**
* 检查用户隧道是否启用且有到期时间
*
* @param userTunnel 用户隧道对象
* @return 是否启用且有到期时间
*/
private boolean isEnabledAndHasExpTime(UserTunnel userTunnel) {
return userTunnel.getStatus() != null && userTunnel.getStatus() == 1
&& userTunnel.getExpTime() != null;
}
}
+43
View File
@@ -0,0 +1,43 @@
spring:
datasource:
driver-class-name: com.mysql.cj.jdbc.Driver
url: jdbc:mysql://${DB_HOST}:3306/${DB_NAME}?useUnicode=true&useSSL=false&characterEncoding=utf8&serverTimezone=Asia/Shanghai&rewriteBatchedStatements=true
username: ${DB_USER}
password: ${DB_PASSWORD}
hikari:
max-lifetime: 500000
connection-timeout: 30000
idle-timeout: 600000
max-pool-size: 20
minimum-idle: 5
pool-name: HikariCP
auto-commit: true
connection-test-query: SELECT 1
redis:
host: 127.0.0.1
port: 6379
database: 0
password:
servlet:
multipart:
max-file-size: 50MB
max-request-size: 50MB
lifecycle:
timeout-per-shutdown-phase: 30s
server:
port: 6365
tomcat:
uri-encoding: UTF-8
max-thread: 800
max-connections: 2000
shutdown: graceful
mybatis-plus:
mapper-locations: classpath*:/mapper/**Mapper.xml
# configuration:
# log-impl: org.apache.ibatis.logging.stdout.StdOutImpl
jwt-secret: ${JWT_SECRET}
log-dir: ${LOG_DIR}
server-addr: ${SERVER_ADDR}
@@ -0,0 +1,48 @@
<?xml version="1.0" encoding="UTF-8"?>
<configuration scan="true" scanPeriod="60 seconds" debug="false">
<contextName>logback</contextName>
<springProperty scope="context" name="logDir" source="log-dir" defaultValue="logs" />
<property name="FILE_PATH" value="${logDir}/%d{yyyy-MM-dd}.log" />
<!--输出到控制台-->
<appender name="console" class="ch.qos.logback.core.ConsoleAppender">
<encoder>
<pattern>%d{HH:mm:ss} [%thread] %-5level %logger{36} - %msg%n</pattern>
</encoder>
</appender>
<!--按天生成日志,即一天只生成一个文件夹和一个日志文件-->
<appender name="logFile" class="ch.qos.logback.core.rolling.RollingFileAppender">
<Prudent>true</Prudent>
<rollingPolicy class="ch.qos.logback.core.rolling.TimeBasedRollingPolicy">
<FileNamePattern>${FILE_PATH}</FileNamePattern>
<maxHistory>30</maxHistory>
</rollingPolicy>
<layout class="ch.qos.logback.classic.PatternLayout">
<Pattern>
%d{yyyy-MM-dd HH:mm:ss} ^ %-5level ^ %logger{36} ^ %msg%n
</Pattern>
</layout>
</appender>
<!-- logger节点,可选节点,作用是指明具体的包或类的日志输出级别,
以及要使用的<appender>(可以把<appender>理解为一个日志模板)。
addtivity:非必写属性,是否向上级loger传递打印信息。默认是true-->
<logger name="com.framework.job" additivity="false">
<appender-ref ref="console"/>
<appender-ref ref="logFile"/>
</logger>
<!--项目的整体的日志打印级别为info-->
<root level="info">
<appender-ref ref="console"/>
<appender-ref ref="logFile"/>
</root>
</configuration>
@@ -0,0 +1,71 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
<mapper namespace="com.admin.mapper.ForwardMapper">
<!-- 查询所有转发信息(包含隧道信息) -->
<select id="selectAllForwardsWithTunnel" resultType="com.admin.common.dto.ForwardWithTunnelDto">
SELECT
f.id,
f.user_id AS userId,
f.name,
f.tunnel_id AS tunnelId,
f.in_port AS inPort,
f.out_port AS outPort,
f.remote_addr AS remoteAddr,
f.status,
f.created_time AS createdTime,
f.updated_time AS updatedTime,
f.user_name AS userName,
f.in_flow AS inFlow,
f.out_flow AS outFlow,
t.name AS tunnelName,
t.in_ip AS inIp,
t.in_port_sta AS inPortSta,
t.in_port_end AS inPortEnd,
t.out_ip AS outIp,
t.out_ip_sta AS outIpSta,
t.out_ip_end AS outIpEnd,
t.type
FROM
forward f
LEFT JOIN
tunnel t ON f.tunnel_id = t.id
ORDER BY
f.created_time DESC
</select>
<!-- 根据用户ID查询转发信息(包含隧道信息) -->
<select id="selectForwardsWithTunnelByUserId" resultType="com.admin.common.dto.ForwardWithTunnelDto">
SELECT
f.id,
f.user_id AS userId,
f.name,
f.tunnel_id AS tunnelId,
f.in_port AS inPort,
f.out_port AS outPort,
f.remote_addr AS remoteAddr,
f.status,
f.created_time AS createdTime,
f.updated_time AS updatedTime,
f.user_name AS userName,
f.in_flow AS inFlow,
f.out_flow AS outFlow,
t.name AS tunnelName,
t.in_ip AS inIp,
t.in_port_sta AS inPortSta,
t.in_port_end AS inPortEnd,
t.out_ip AS outIp,
t.out_ip_sta AS outIpSta,
t.out_ip_end AS outIpEnd,
t.type
FROM
forward f
LEFT JOIN
tunnel t ON f.tunnel_id = t.id
WHERE
f.user_id = #{userId}
ORDER BY
f.created_time DESC
</select>
</mapper>
@@ -0,0 +1,5 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
<mapper namespace="com.admin.mapper.NodeMapper">
</mapper>
@@ -0,0 +1,5 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
<mapper namespace="com.admin.mapper.SpeedLimitMapper">
</mapper>
@@ -0,0 +1,5 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
<mapper namespace="com.admin.mapper.TunnelMapper">
</mapper>
@@ -0,0 +1,72 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
<mapper namespace="com.admin.mapper.UserMapper">
<!-- 查询用户隧道权限详情(连表查询隧道和限速规则信息) -->
<select id="getUserTunnelDetails" resultType="com.admin.common.dto.UserPackageDto$UserTunnelDetailDto">
SELECT
ut.id,
ut.user_id as userId,
ut.tunnel_id as tunnelId,
t.name as tunnelName,
t.flow as tunnelFlow,
ut.flow,
ut.in_flow as inFlow,
ut.out_flow as outFlow,
ut.num,
ut.flow_reset_time as flowResetTime,
ut.exp_time as expTime,
ut.speed_id as speedId,
sl.name as speedLimitName,
sl.speed
FROM user_tunnel ut
LEFT JOIN tunnel t ON ut.tunnel_id = t.id
LEFT JOIN speed_limit sl ON ut.speed_id = sl.id
WHERE ut.user_id = #{userId}
ORDER BY ut.id
</select>
<!-- 查询用户转发详情(连表查询隧道信息) -->
<select id="getUserForwardDetails" resultType="com.admin.common.dto.UserPackageDto$UserForwardDetailDto">
SELECT
f.id,
f.name,
f.tunnel_id as tunnelId,
t.name as tunnelName,
t.in_ip as inIp,
f.in_port as inPort,
f.remote_addr as remoteAddr,
f.in_flow as inFlow,
f.out_flow as outFlow,
f.status,
f.created_time as createdTime
FROM forward f
LEFT JOIN tunnel t ON f.tunnel_id = t.id
WHERE f.user_id = #{userId}
ORDER BY f.created_time DESC
</select>
<!-- 管理员查询所有隧道(流量和转发设置为99999) -->
<select id="getAllTunnelsForAdmin" resultType="com.admin.common.dto.UserPackageDto$UserTunnelDetailDto">
SELECT
t.id,
0 as userId,
t.id as tunnelId,
t.name as tunnelName,
t.flow as tunnelFlow,
99999 as flow,
0 as inFlow,
0 as outFlow,
99999 as num,
null as flowResetTime,
null as expTime,
null as speedId,
'无限制' as speedLimitName,
null as speed
FROM tunnel t
WHERE t.status = 1
ORDER BY t.id
</select>
</mapper>
@@ -0,0 +1,36 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
<mapper namespace="com.admin.mapper.UserTunnelMapper">
<!-- 获取用户隧道权限及隧道详细信息 -->
<select id="getUserTunnelWithDetails" resultType="com.admin.common.dto.UserTunnelWithDetailDto">
SELECT
ut.id,
ut.user_id as userId,
ut.tunnel_id as tunnelId,
ut.flow,
ut.in_flow as inFlow,
ut.out_flow as outFlow,
ut.num,
ut.flow_reset_time as flowResetTime,
ut.exp_time as expTime,
ut.speed_id as speedId,
t.name as tunnelName,
t.flow as tunnelFlow,
t.in_ip as inIp,
t.in_port_sta as inPortSta,
t.in_port_end as inPortEnd,
t.out_ip as outIp,
t.out_ip_sta as outIpSta,
t.out_ip_end as outIpEnd,
t.type,
sl.name as speedLimitName,
sl.speed
FROM user_tunnel ut
LEFT JOIN tunnel t ON ut.tunnel_id = t.id
LEFT JOIN speed_limit sl ON ut.speed_id = sl.id
WHERE ut.user_id = #{userId}
ORDER BY ut.id
</select>
</mapper>
@@ -0,0 +1,8 @@
package com.admin;
import org.springframework.boot.test.context.SpringBootTest;
@SpringBootTest
class AdminApplicationTests {
}