From 132d2a9f71addf471386477ac8b4c5f9b565e58b Mon Sep 17 00:00:00 2001 From: John Smith Date: Sat, 2 Nov 2024 23:22:54 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0=E6=9C=8D=E5=8A=A1=E5=8F=91?= =?UTF-8?q?=E7=8E=B0=E6=8E=A5=E5=8F=A3=E3=80=81=E6=94=AF=E6=8C=81=E8=B7=A8?= =?UTF-8?q?=E5=9F=9F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/base.py | 21 +++++++++++++++- api/chat.py | 9 ++++--- api/main.py | 9 +++++++ api/open_live.py | 3 +++ api/plugin.py | 3 ++- config.py | 55 +++++++++++++++++++++++++++++++++++------ data/config.example.ini | 10 ++++++++ 7 files changed, 96 insertions(+), 14 deletions(-) diff --git a/api/base.py b/api/base.py index fc3bdbf..14c5b6e 100644 --- a/api/base.py +++ b/api/base.py @@ -4,18 +4,37 @@ from typing import * import tornado.web +import config + class ApiHandler(tornado.web.RequestHandler): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.json_args: Optional[dict] = None - def prepare(self): + def set_default_headers(self): self.set_header('Cache-Control', 'no-cache') + self.add_header('Vary', 'Origin') + origin = self.request.headers.get('Origin', None) + if origin is None: + return + cfg = config.get_config() + if not cfg.is_allowed_cors_origin(origin): + return + + self.set_header('Access-Control-Allow-Origin', origin) + self.set_header('Access-Control-Allow-Methods', '*') + self.set_header('Access-Control-Allow-Headers', '*') + self.set_header('Access-Control-Max-Age', '3600') + + def prepare(self): if not self.request.headers.get('Content-Type', '').startswith('application/json'): return try: self.json_args = json.loads(self.request.body) except json.JSONDecodeError: pass + + async def options(self, *_args, **_kwargs): + self.set_status(204) diff --git a/api/chat.py b/api/chat.py index 93c628d..9bd979e 100644 --- a/api/chat.py +++ b/api/chat.py @@ -223,12 +223,13 @@ class ChatHandler(tornado.websocket.WebSocketHandler): self._refresh_receive_timeout_timer() - # 跨域测试用 def check_origin(self, origin): cfg = config.get_config() - if cfg.debug: - return True - return super().check_origin(origin) + return ( + cfg.debug # 开发时前端localhost直连 + or cfg.is_allowed_cors_origin(origin) + or super().check_origin(origin) # 和Host相同 + ) @property def has_joined_room(self): diff --git a/api/main.py b/api/main.py index be9d7a6..97f9229 100644 --- a/api/main.py +++ b/api/main.py @@ -52,6 +52,14 @@ class ServerInfoHandler(api.base.ApiHandler): }) +class ServiceDiscoveryHandler(api.base.ApiHandler): + async def get(self): + cfg = config.get_config() + self.write({ + 'endpoints': cfg.registered_endpoints, + }) + + class UploadEmoticonHandler(api.base.ApiHandler): async def post(self): cfg = config.get_config() @@ -94,6 +102,7 @@ class NoCacheStaticFileHandler(tornado.web.StaticFileHandler): ROUTES = [ (r'/api/server_info', ServerInfoHandler), + (r'/api/endpoints', ServiceDiscoveryHandler), (r'/api/emoticon', UploadEmoticonHandler), ] # 通配的放在最后 diff --git a/api/open_live.py b/api/open_live.py index 9ecbbc4..0d2e9a0 100644 --- a/api/open_live.py +++ b/api/open_live.py @@ -169,6 +169,9 @@ class _OpenLiveHandlerBase(api.base.ApiHandler): def prepare(self): super().prepare() + if self.request.method == 'OPTIONS': + return + if not isinstance(self.json_args, dict): raise tornado.web.MissingArgumentError('body') diff --git a/api/plugin.py b/api/plugin.py index 6dfff9b..d4f6787 100644 --- a/api/plugin.py +++ b/api/plugin.py @@ -24,7 +24,8 @@ class _AdminHandlerBase(api.base.ApiHandler): if not cfg.enable_admin_plugins: raise tornado.web.HTTPError(403) - logger.info('client=%s requesting admin plugin, cls=%s', self.request.remote_ip, type(self).__name__) + if self.request.method != 'OPTIONS': + logger.info('client=%s requesting admin plugin, cls=%s', self.request.remote_ip, type(self).__name__) super().prepare() diff --git a/config.py b/config.py index 3fe8a78..19b40d0 100644 --- a/config.py +++ b/config.py @@ -2,6 +2,7 @@ import configparser import logging import os +import re from typing import * logger = logging.getLogger(__name__) @@ -30,7 +31,7 @@ def init(cmd_args): _config = config -def reload(cmd_args): +def reload(cmd_args=None): config_path = '' for path in CONFIG_PATH_LIST: if os.path.exists(path): @@ -73,12 +74,15 @@ class AppConfig: self.open_live_app_id = 0 self.enable_translate = True - self.allow_translate_rooms = set() + self.allow_translate_rooms: Set[int] = set() self.translate_max_queue_size = 10 self.translation_cache_size = 50000 - self.translator_configs = [] + self.translator_configs: List[dict] = [] - self.text_emoticons = [] + self.text_emoticons: List[dict] = [] + + self.registered_endpoints: List[str] = [] + self.cors_origins: List[re.Pattern[str]] = [] @property def is_open_live_configured(self): @@ -86,7 +90,9 @@ class AppConfig: self.open_live_access_key_id != '' and self.open_live_access_key_secret != '' and self.open_live_app_id != 0 ) - def load_cmd_args(self, args): + def load_cmd_args(self, args=None): + if args is None: + return if args.host is not None: self.host = args.host if args.port is not None: @@ -101,6 +107,8 @@ class AppConfig: self._load_app_config(config) self._load_translator_configs(config) self._load_text_emoticons(config) + self._load_registered_endpoints(config) + self._load_cors_origins(config) except Exception: # noqa logger.exception('Failed to load config:') return False @@ -145,7 +153,10 @@ class AppConfig: return url def _load_translator_configs(self, config: configparser.ConfigParser): - app_section = config['app'] + try: + app_section = config['app'] + except KeyError: + return section_names = _str_to_list(app_section.get('translator_configs', '')) translator_configs = [] for section_name in section_names: @@ -191,13 +202,41 @@ class AppConfig: self.translator_configs = translator_configs def _load_text_emoticons(self, config: configparser.ConfigParser): - mappings_section = config['text_emoticon_mappings'] + try: + mappings_section = config['text_emoticon_mappings'] + except KeyError: + return text_emoticons = [] for value in mappings_section.values(): keyword, _, url = value.partition(',') text_emoticons.append({'keyword': keyword, 'url': url}) self.text_emoticons = text_emoticons + def _load_registered_endpoints(self, config: configparser.ConfigParser): + try: + registered_endpoints_section = config['registered_endpoints'] + except KeyError: + return + registered_endpoints = list(registered_endpoints_section.values()) + self.registered_endpoints = registered_endpoints + + def _load_cors_origins(self, config: configparser.ConfigParser): + try: + cors_origins_section = config['cors_origins'] + except KeyError: + return + cors_origins = [ + re.compile(origin, re.IGNORECASE) + for origin in cors_origins_section.values() + ] + self.cors_origins = cors_origins + + def is_allowed_cors_origin(self, origin): + return any( + pattern.fullmatch(origin) is not None + for pattern in self.cors_origins + ) + def _str_to_list(value, item_type: Type = str, container_type: Type = list): value = value.strip() @@ -206,5 +245,5 @@ def _str_to_list(value, item_type: Type = str, container_type: Type = list): items = value.split(',') items = map(lambda item: item.strip(), items) if item_type is not str: - items = map(lambda item: item_type(item), items) + items = map(item_type, items) return container_type(items) diff --git a/data/config.example.ini b/data/config.example.ini index 2ce2320..461a811 100644 --- a/data/config.example.ini +++ b/data/config.example.ini @@ -266,3 +266,13 @@ temperature = 0.4 80 = [抱拳],http://i0.hdslb.com/bfs/live/3f170894dd08827ee293afcb5a3d2b60aecdb5b1.png 81 = [给力],http://i0.hdslb.com/bfs/live/d1ba5f4c54332a21ed2ca0dcecaedd2add587839.png 82 = [耶],http://i0.hdslb.com/bfs/live/eb2d84ba623e2335a48f73fb5bef87bcf53c1239.png + + +# 用于服务发现返回的后端端点 +[registered_endpoints] +# 1 = https://api1.blive.chat + + +# 允许跨域的源,正则表达式 +[cors_origins] +# 1 = https://(?:|.+\.)blive\.chat