域名映射功能重构:支持多域名HTTPS映射和SSL证书自动切换

This commit is contained in:
suxiang
2024-08-27 22:30:05 +08:00
parent 5aac309f47
commit 243605481a
4 changed files with 178 additions and 13 deletions
@@ -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;
}
}
@@ -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);
}
}
@@ -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());
@@ -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