域名映射功能重构:支持多域名HTTPS映射和SSL证书自动切换
This commit is contained in:
+25
-1
@@ -7,7 +7,10 @@ import io.netty.channel.nio.NioEventLoopGroup;
|
||||
import io.netty.channel.socket.SocketChannel;
|
||||
import io.netty.channel.socket.nio.NioServerSocketChannel;
|
||||
import io.netty.handler.logging.LoggingHandler;
|
||||
import io.netty.handler.ssl.SniHandler;
|
||||
import io.netty.handler.ssl.SslContext;
|
||||
import io.netty.handler.ssl.SslHandler;
|
||||
import io.netty.util.Mapping;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.apache.commons.lang3.StringUtils;
|
||||
import org.dromara.neutrinoproxy.core.util.FileUtil;
|
||||
@@ -35,6 +38,9 @@ import java.security.KeyStore;
|
||||
public class HttpsProxy implements EventListener<AppLoadEndEvent> {
|
||||
@Inject
|
||||
private ProxyConfig proxyConfig;
|
||||
@Inject
|
||||
private SslContextManager sslContextManager;
|
||||
|
||||
@Override
|
||||
public void onEvent(AppLoadEndEvent appLoadEndEvent) throws Throwable {
|
||||
if (null == proxyConfig.getServer().getTcp().getHttpsProxyPort() ||
|
||||
@@ -55,7 +61,7 @@ public class HttpsProxy implements EventListener<AppLoadEndEvent> {
|
||||
if (null != proxyConfig.getServer().getTcp().getTransferLogEnable() && proxyConfig.getServer().getTcp().getTransferLogEnable()) {
|
||||
ch.pipeline().addFirst(new LoggingHandler(HttpsProxy.class));
|
||||
}
|
||||
ch.pipeline().addLast(createSslHandler());
|
||||
ch.pipeline().addLast(createSniHandler());
|
||||
ch.pipeline().addFirst(new BytesMetricsHandler());
|
||||
ch.pipeline().addLast(new HttpVisitorSecurityChannelHandler(true));
|
||||
ch.pipeline().addLast("flowLimiter",new VisitorFlowLimiterChannelHandler());
|
||||
@@ -82,6 +88,7 @@ public class HttpsProxy implements EventListener<AppLoadEndEvent> {
|
||||
|
||||
serverContext.init(kmf.getKeyManagers(), trustManagers, null);
|
||||
|
||||
|
||||
SSLEngine sslEngine = serverContext.createSSLEngine();
|
||||
sslEngine.setUseClientMode(false);
|
||||
sslEngine.setNeedClientAuth(false);
|
||||
@@ -93,4 +100,21 @@ public class HttpsProxy implements EventListener<AppLoadEndEvent> {
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
public SniHandler createSniHandler() {
|
||||
try {
|
||||
return new SniHandler(domainName -> {
|
||||
SslContext sslContext = sslContextManager.getSslContextByFullDomain(domainName);
|
||||
if (sslContext != null) {
|
||||
return sslContext;
|
||||
} else {
|
||||
throw new IllegalArgumentException("No SSL context available for domain: " + domainName);
|
||||
}
|
||||
});
|
||||
} catch (Exception e) {
|
||||
log.info("create SSL handler failed", e);
|
||||
e.printStackTrace();
|
||||
}
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
+110
@@ -0,0 +1,110 @@
|
||||
package org.dromara.neutrinoproxy.server.proxy.enhance;
|
||||
|
||||
import ch.qos.logback.core.net.ssl.SSL;
|
||||
import com.baomidou.mybatisplus.core.toolkit.Wrappers;
|
||||
import io.netty.handler.codec.http.HttpUtil;
|
||||
import io.netty.handler.ssl.ClientAuth;
|
||||
import io.netty.handler.ssl.SniHandler;
|
||||
import io.netty.handler.ssl.SslContext;
|
||||
import io.netty.handler.ssl.SslContextBuilder;
|
||||
import io.netty.handler.ssl.util.SelfSignedCertificate;
|
||||
import io.netty.util.DomainWildcardMappingBuilder;
|
||||
import lombok.Data;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.apache.ibatis.solon.annotation.Db;
|
||||
import org.dromara.neutrinoproxy.core.util.FileUtil;
|
||||
import org.dromara.neutrinoproxy.server.base.proxy.ProxyConfig;
|
||||
import org.dromara.neutrinoproxy.server.dal.DomainMapper;
|
||||
import org.dromara.neutrinoproxy.server.dal.entity.DomainNameDO;
|
||||
import org.dromara.neutrinoproxy.server.util.ProxyUtil;
|
||||
import org.noear.solon.annotation.Component;
|
||||
import org.noear.solon.annotation.Http;
|
||||
import org.noear.solon.annotation.Init;
|
||||
import org.noear.solon.annotation.Inject;
|
||||
|
||||
import javax.net.ssl.KeyManagerFactory;
|
||||
import javax.net.ssl.SSLContext;
|
||||
import java.io.ByteArrayInputStream;
|
||||
import java.io.FileInputStream;
|
||||
import java.io.InputStream;
|
||||
import java.security.KeyStore;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
/**
|
||||
* 域名到SSL上下文
|
||||
* @author: Mirac
|
||||
* @date: 2024/8/25
|
||||
*/
|
||||
@Slf4j
|
||||
@Component
|
||||
@Data
|
||||
public class SslContextManager {
|
||||
|
||||
@Db
|
||||
private DomainMapper domainMapper;
|
||||
|
||||
// 维护域名到 SslContext 的映射
|
||||
private final ConcurrentHashMap<String, SslContext> domainSslContexts = new ConcurrentHashMap<>();
|
||||
|
||||
// 初始化时加载所有域名的 SSL 上下文
|
||||
@Init
|
||||
public void sslContextManagerInit() {
|
||||
try {
|
||||
initializeSslContexts();
|
||||
} catch (Exception e) {
|
||||
log.error("initialize SSL handler failed", e);
|
||||
e.printStackTrace();
|
||||
}
|
||||
}
|
||||
|
||||
// 使用 JKS 文件加载 SSL 上下文,禁用客户端认证并设置为服务器模式
|
||||
private SslContext loadSslContextFromJks(byte[] jks, String keyStorePassword) throws Exception {
|
||||
InputStream jksInputStream = new ByteArrayInputStream(jks);
|
||||
// 初始化 KeyStore
|
||||
KeyStore keyStore = KeyStore.getInstance("JKS");
|
||||
keyStore.load(jksInputStream, keyStorePassword.toCharArray());
|
||||
|
||||
// 初始化 KeyManagerFactory
|
||||
KeyManagerFactory kmf = KeyManagerFactory.getInstance(KeyManagerFactory.getDefaultAlgorithm());
|
||||
kmf.init(keyStore, keyStorePassword.toCharArray());
|
||||
|
||||
// 创建 SslContext 并禁用客户端认证
|
||||
return SslContextBuilder.forServer(kmf)
|
||||
.clientAuth(ClientAuth.NONE) // 禁用客户端认证
|
||||
.build();
|
||||
}
|
||||
|
||||
// 初始化所有域名的 SSL 上下文
|
||||
private void initializeSslContexts() throws Exception {
|
||||
List<DomainNameDO> domainNameDOS = domainMapper.selectList(Wrappers.<DomainNameDO>lambdaQuery()
|
||||
.isNotNull(DomainNameDO::getDomain)
|
||||
.isNotNull(DomainNameDO::getJks)
|
||||
.isNotNull(DomainNameDO::getKeyStorePassword));
|
||||
|
||||
for (DomainNameDO domainNameDO : domainNameDOS) {
|
||||
if (domainNameDO.getDomain() != null && domainNameDO.getKeyStorePassword() != null && domainNameDO.getJks() != null) {
|
||||
String domain = domainNameDO.getDomain();
|
||||
byte[] jks = domainNameDO.getJks();// 假设数据库中存储了每个域名对应的 JKS 路径
|
||||
String keyStorePassword = domainNameDO.getKeyStorePassword();
|
||||
SslContext sslContext = loadSslContextFromJks(jks, keyStorePassword);
|
||||
domainSslContexts.put(domain, sslContext);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 动态添加新的域名和证书
|
||||
public void addDomainAndCert(String domain, byte[] jks, String keyStorePassword) throws Exception {
|
||||
SslContext sslContext = loadSslContextFromJks(jks, keyStorePassword);
|
||||
domainSslContexts.put(domain, sslContext);
|
||||
}
|
||||
|
||||
public SslContext getSslContextByDomain(String domain) {
|
||||
return domainSslContexts.get(domain);
|
||||
}
|
||||
|
||||
public SslContext getSslContextByFullDomain(String fullDomain) {
|
||||
String domain = ProxyUtil.getDomainNameByFullDomain(fullDomain);
|
||||
return domainSslContexts.get(domain);
|
||||
}
|
||||
}
|
||||
+31
-12
@@ -29,11 +29,13 @@ import org.dromara.neutrinoproxy.server.dal.entity.DomainNameDO;
|
||||
import org.dromara.neutrinoproxy.server.dal.entity.DomainPortMappingDO;
|
||||
import org.dromara.neutrinoproxy.server.dal.entity.PortMappingDO;
|
||||
import org.dromara.neutrinoproxy.server.dal.entity.UserDO;
|
||||
import org.dromara.neutrinoproxy.server.proxy.enhance.SslContextManager;
|
||||
import org.dromara.neutrinoproxy.server.service.bo.FullDomainNameBO;
|
||||
import org.dromara.neutrinoproxy.server.util.ParamCheckUtil;
|
||||
import org.dromara.neutrinoproxy.server.util.ProxyUtil;
|
||||
import org.noear.solon.annotation.Component;
|
||||
import org.noear.solon.annotation.Init;
|
||||
import org.noear.solon.annotation.Inject;
|
||||
import org.noear.solon.core.handle.UploadedFile;
|
||||
|
||||
import java.io.ByteArrayOutputStream;
|
||||
@@ -59,6 +61,9 @@ public class DomainService {
|
||||
@Db
|
||||
private DomainPortMappingMapper domainPortMappingMapper;
|
||||
|
||||
@Inject
|
||||
private SslContextManager sslContextManager;
|
||||
|
||||
|
||||
public PageInfo<DomainListRes> page(PageQuery pageQuery, DomainListReq req) {
|
||||
Page<DomainNameDO> page = domainMapper.selectPage(new Page<>(pageQuery.getCurrent(), pageQuery.getSize()), new LambdaQueryWrapper<DomainNameDO>()
|
||||
@@ -107,17 +112,24 @@ public class DomainService {
|
||||
*
|
||||
* @param req
|
||||
*/
|
||||
public void create(DomainCreateReq req, UploadedFile jks) throws IOException {
|
||||
public void create(DomainCreateReq req, UploadedFile jks) {
|
||||
DomainNameDO domainNameCheck = domainMapper.checkRepeat(req.getDomain(), null);
|
||||
ParamCheckUtil.checkMustNull(domainNameCheck, ExceptionConstant.DOMAIN_NAME_CANNOT_REPEAT);
|
||||
DomainNameDO domainNameDO = new DomainNameDO();
|
||||
if (null != jks) {
|
||||
ParamCheckUtil.checkNotEmpty(req.getKeyStorePassword(), "keyStorePassword");
|
||||
try {
|
||||
ParamCheckUtil.checkNotEmpty(req.getKeyStorePassword(), "keyStorePassword");
|
||||
|
||||
InputStream content = jks.getContent();
|
||||
byte[] byteArray = toByteArray(content);
|
||||
domainNameDO.setJks(byteArray);
|
||||
domainNameDO.setKeyStorePassword(req.getKeyStorePassword());
|
||||
InputStream content = jks.getContent();
|
||||
byte[] byteArray = toByteArray(content);
|
||||
domainNameDO.setJks(byteArray);
|
||||
domainNameDO.setKeyStorePassword(req.getKeyStorePassword());
|
||||
//添加证书
|
||||
sslContextManager.addDomainAndCert(req.getDomain(), byteArray, req.getKeyStorePassword());
|
||||
} catch (Exception e) {
|
||||
log.error("证书添加失败", e);
|
||||
e.printStackTrace();
|
||||
}
|
||||
}
|
||||
Integer userId = SystemContextHolder.getUserId();
|
||||
Date now = new Date();
|
||||
@@ -147,7 +159,7 @@ public class DomainService {
|
||||
return buffer.toByteArray();
|
||||
}
|
||||
|
||||
public void update(DomainUpdateReq req, UploadedFile jks) throws IOException {
|
||||
public void update(DomainUpdateReq req, UploadedFile jks) {
|
||||
DomainNameDO domainNameCheck = domainMapper.checkRepeat(req.getDomain(), Sets.newHashSet(req.getId()));
|
||||
ParamCheckUtil.checkMustNull(domainNameCheck, ExceptionConstant.DOMAIN_NAME_CANNOT_REPEAT);
|
||||
|
||||
@@ -156,12 +168,19 @@ public class DomainService {
|
||||
LambdaUpdateWrapper<DomainNameDO> updateWrapper = new LambdaUpdateWrapper<>();
|
||||
updateWrapper.eq(DomainNameDO::getId, req.getId());
|
||||
if (null != jks) {
|
||||
ParamCheckUtil.checkNotEmpty(req.getKeyStorePassword(), "keyStorePassword");
|
||||
try {
|
||||
ParamCheckUtil.checkNotEmpty(req.getKeyStorePassword(), "keyStorePassword");
|
||||
|
||||
InputStream content = jks.getContent();
|
||||
byte[] byteArray = toByteArray(content);
|
||||
updateWrapper.set(DomainNameDO::getKeyStorePassword, req.getKeyStorePassword());
|
||||
updateWrapper.set(DomainNameDO::getJks, byteArray);
|
||||
InputStream content = jks.getContent();
|
||||
byte[] byteArray = toByteArray(content);
|
||||
updateWrapper.set(DomainNameDO::getKeyStorePassword, req.getKeyStorePassword());
|
||||
updateWrapper.set(DomainNameDO::getJks, byteArray);
|
||||
//添加证书
|
||||
sslContextManager.addDomainAndCert(req.getDomain(), byteArray, req.getKeyStorePassword());
|
||||
} catch (Exception e) {
|
||||
log.error("证书添加失败", e);
|
||||
e.printStackTrace();
|
||||
}
|
||||
}
|
||||
updateWrapper.set(DomainNameDO::getDomain, req.getDomain());
|
||||
updateWrapper.set(DomainNameDO::getUpdateTime, new Date());
|
||||
|
||||
+12
@@ -409,6 +409,18 @@ public class ProxyUtil {
|
||||
return getDomainNameIdByDomain(domains.get(0));
|
||||
}
|
||||
|
||||
/**
|
||||
* 通过完整域名获取域名id
|
||||
*/
|
||||
public static String getDomainNameByFullDomain(String fullDomain) {
|
||||
List<String> domains = domainToDomainNameIdMap.keySet().stream().filter(item -> fullDomain.endsWith(item)).collect(Collectors.toList());
|
||||
// 不存在 或者 有多条记录,返回null
|
||||
if (CollectionUtil.isEmpty(domains) || domains.size() > 1) {
|
||||
return null;
|
||||
}
|
||||
return domains.getFirst();
|
||||
}
|
||||
|
||||
/**
|
||||
* 关闭http响应channel
|
||||
* @param channel
|
||||
|
||||
Reference in New Issue
Block a user