fix DoH热加载

This commit is contained in:
景大侠
2025-06-13 17:43:45 +08:00
parent db5d81d7f0
commit e162bd1168
3 changed files with 42 additions and 10 deletions
-1
View File
@@ -1,2 +1 @@
from .doh import doh_query_json
from .cloudflare import under_challenge from .cloudflare import under_challenge
+39 -9
View File
@@ -10,10 +10,14 @@ import socket
import struct import struct
import urllib import urllib
import urllib.request import urllib.request
from threading import Lock
from typing import Dict, Optional from typing import Dict, Optional
from app.core.config import settings from app.core.config import settings
from app.core.event import Event, eventmanager
from app.log import logger from app.log import logger
from app.schemas import ConfigChangeEventData
from app.schemas.types import EventType
# 定义一个全局线程池执行器 # 定义一个全局线程池执行器
_executor = concurrent.futures.ThreadPoolExecutor() _executor = concurrent.futures.ThreadPoolExecutor()
@@ -21,11 +25,15 @@ _executor = concurrent.futures.ThreadPoolExecutor()
# 定义默认的DoH配置 # 定义默认的DoH配置
_doh_timeout = 5 _doh_timeout = 5
_doh_cache: Dict[str, str] = {} _doh_cache: Dict[str, str] = {}
_doh_lock = Lock()
# 保存原始的 socket.getaddrinfo 方法
_orig_getaddrinfo = socket.getaddrinfo
# 对 socket.getaddrinfo 进行补丁
if settings.DOH_ENABLE: def enable_doh(enable: bool):
# 保存原始的 socket.getaddrinfo 方法 """
_orig_getaddrinfo = socket.getaddrinfo 对 socket.getaddrinfo 进行补丁
"""
def _patched_getaddrinfo(host, *args, **kwargs): def _patched_getaddrinfo(host, *args, **kwargs):
""" """
@@ -34,8 +42,9 @@ if settings.DOH_ENABLE:
if host not in settings.DOH_DOMAINS.split(","): if host not in settings.DOH_DOMAINS.split(","):
return _orig_getaddrinfo(host, *args, **kwargs) return _orig_getaddrinfo(host, *args, **kwargs)
# 检查主机是否已解析 # 检查主机是否已解析
if host in _doh_cache: with _doh_lock:
ip = _doh_cache[host] ip = _doh_cache.get("host", None)
if ip is not None:
logger.info("已解析 [%s] 为 [%s] (缓存)", host, ip) logger.info("已解析 [%s] 为 [%s] (缓存)", host, ip)
return _orig_getaddrinfo(ip, *args, **kwargs) return _orig_getaddrinfo(ip, *args, **kwargs)
# 使用DoH解析主机 # 使用DoH解析主机
@@ -46,13 +55,34 @@ if settings.DOH_ENABLE:
ip = future.result() ip = future.result()
if ip is not None: if ip is not None:
logger.info("已解析 [%s] 为 [%s]", host, ip) logger.info("已解析 [%s] 为 [%s]", host, ip)
_doh_cache[host] = ip with _doh_lock:
_doh_cache[host] = ip
host = ip host = ip
break break
return _orig_getaddrinfo(host, *args, **kwargs) return _orig_getaddrinfo(host, *args, **kwargs)
# 替换 socket.getaddrinfo 方法 if enable:
socket.getaddrinfo = _patched_getaddrinfo # 替换 socket.getaddrinfo 方法
socket.getaddrinfo = _patched_getaddrinfo
else:
socket.getaddrinfo = _orig_getaddrinfo
class DohHelper:
def __init__(self):
enable_doh(settings.DOH_ENABLE)
@eventmanager.register(EventType.ConfigChanged)
@staticmethod
def handle_config_changed(event: Event):
if not event:
return
event_data: ConfigChangeEventData = event.event_data
if event_data.key not in ["DOH_ENABLE", "DOH_DOMAINS", "DOH_RESOLVERS"]:
return
with _doh_lock:
# DOH配置有变动的情况下,清空缓存
_doh_cache.clear()
enable_doh(settings.DOH_ENABLE)
def _doh_query(resolver: str, host: str) -> Optional[str]: def _doh_query(resolver: str, host: str) -> Optional[str]:
+3
View File
@@ -19,6 +19,7 @@ except ImportError as e:
from app.core.event import EventManager from app.core.event import EventManager
from app.helper.thread import ThreadHelper from app.helper.thread import ThreadHelper
from app.helper.display import DisplayHelper from app.helper.display import DisplayHelper
from app.helper.doh import DohHelper
from app.helper.resource import ResourceHelper from app.helper.resource import ResourceHelper
from app.helper.message import MessageHelper from app.helper.message import MessageHelper
from app.schemas import Notification, NotificationType from app.schemas import Notification, NotificationType
@@ -132,6 +133,8 @@ def init_modules():
""" """
# 虚拟显示 # 虚拟显示
DisplayHelper() DisplayHelper()
# DoH
DohHelper()
# 站点管理 # 站点管理
SitesHelper() SitesHelper()
# 资源包检测 # 资源包检测