From 08cd582e44e8f724efde577373bdf444b966ab1f Mon Sep 17 00:00:00 2001 From: funnywolf Date: Fri, 26 Dec 2025 21:34:22 +0800 Subject: [PATCH] update the performance of redis --- ASP/settings.py | 8 + Lib/basemodule.py | 1 + Lib/log.py | 2 +- Lib/montior.py | 8 +- ...11-Suspicious-Process-Spawned-by-Office.py | 7 +- .../ES-Rule-21-Phishing_user_report_mail.py | 8 +- .../ES-Rule-22-Phishing_user_report_mail.py | 4 +- .../NDR-Rule-05-Suspect-C2-Communication.py | 8 +- PLUGINS/Redis/CONFIG.example.py | 3 +- PLUGINS/Redis/redis_client.py | 31 ++-- PLUGINS/Redis/redis_stream_api.py | 166 +++++++----------- 11 files changed, 111 insertions(+), 135 deletions(-) diff --git a/ASP/settings.py b/ASP/settings.py index c9026ec..e0f5450 100644 --- a/ASP/settings.py +++ b/ASP/settings.py @@ -96,6 +96,14 @@ LOGGING = { 'level': 'INFO', 'propagate': True }, + 'asp': { + 'handlers': [ + 'console', + 'file', + ], + 'level': 'DEBUG', + 'propagate': True, + }, } } diff --git a/Lib/basemodule.py b/Lib/basemodule.py index 456c75b..0baac2f 100644 --- a/Lib/basemodule.py +++ b/Lib/basemodule.py @@ -46,6 +46,7 @@ class LanggraphModule(BaseModule): self.agent_state = AgentState() for event in self.graph.stream(self.agent_state, config, stream_mode="values"): self.logger.debug(event) + self.logger.debug(f"{self.module_name} finished processing.") def run(self): self.run_graph() diff --git a/Lib/log.py b/Lib/log.py index 075719c..99b1c05 100644 --- a/Lib/log.py +++ b/Lib/log.py @@ -1,3 +1,3 @@ import logging -logger = logging.getLogger("django") +logger = logging.getLogger("asp") diff --git a/Lib/montior.py b/Lib/montior.py index b628f86..6dc747a 100644 --- a/Lib/montior.py +++ b/Lib/montior.py @@ -15,10 +15,9 @@ from Lib.engine import Engine from Lib.log import logger from Lib.xcache import Xcache from PLUGINS.Embeddings.embeddings_qdrant import embedding_api_singleton_qdrant -from PLUGINS.Redis.CONFIG import REDIS_STREAM_STORE_DAYS +from PLUGINS.Mem0.CONFIG import USE as MEM_ZERO_USE from PLUGINS.Redis.redis_stream_api import RedisStreamAPI from PLUGINS.SIRP.sirpapi import Playbook as SIRPPlaybook, Knowledge, KnowledgeAction -from PLUGINS.Mem0.CONFIG import USE as MEM_ZERO_USE if MEM_ZERO_USE: from PLUGINS.Mem0.mem_zero import mem_zero_singleton @@ -101,10 +100,8 @@ class MainMonitor(object): # self.MainScheduler.start() delay_time = 3 - delay_time_clean_thread = 60 * 60 # Start background tasks - self.start_background_task(self.subscribe_clean_thread, "subscribe_clean_thread", delay_time_clean_thread) self.start_background_task(self.subscribe_pending_playbook, "subscribe_pending_playbook", delay_time) self.start_background_task(self.subscribe_knowledge_action, "subscribe_knowledge_action", delay_time) @@ -112,9 +109,6 @@ class MainMonitor(object): self.engine.start() logger.info("Background services started.") - def subscribe_clean_thread(self): - self.redis_stream_api.clean_redis_stream(max_age_days=REDIS_STREAM_STORE_DAYS) - @staticmethod def subscribe_pending_playbook(): records = SIRPPlaybook.get_pending_playbooks() diff --git a/MODULES/EDR-Rule-11-Suspicious-Process-Spawned-by-Office.py b/MODULES/EDR-Rule-11-Suspicious-Process-Spawned-by-Office.py index 932e007..045f92d 100644 --- a/MODULES/EDR-Rule-11-Suspicious-Process-Spawned-by-Office.py +++ b/MODULES/EDR-Rule-11-Suspicious-Process-Spawned-by-Office.py @@ -4,8 +4,10 @@ from datetime import datetime from typing import Optional, Union, Dict, Any from langchain_core.messages import AIMessage, HumanMessage +from langgraph.graph import END from langgraph.graph import StateGraph from langgraph.graph.state import CompiledStateGraph +from langgraph.types import Command from pydantic import BaseModel, Field from Lib.api import string_to_string_time, get_current_time_str @@ -36,7 +38,10 @@ class Module(LanggraphModule): # Get raw alert from stream alert = self.read_message() if alert is None: - return + return Command( + update={"alert": {}}, + goto=END + ) # for splunk webhook # alert = json.loads(alert["_raw"]) diff --git a/MODULES/ES-Rule-21-Phishing_user_report_mail.py b/MODULES/ES-Rule-21-Phishing_user_report_mail.py index d8f1d3e..3b25463 100644 --- a/MODULES/ES-Rule-21-Phishing_user_report_mail.py +++ b/MODULES/ES-Rule-21-Phishing_user_report_mail.py @@ -4,8 +4,9 @@ from datetime import datetime from typing import Optional, Union, Dict, Any from langchain_core.messages import AIMessage, HumanMessage -from langgraph.graph import StateGraph +from langgraph.graph import StateGraph, END from langgraph.graph.state import CompiledStateGraph +from langgraph.types import Command from pydantic import BaseModel, Field from Lib.api import string_to_string_time, get_current_time_str @@ -41,7 +42,10 @@ class Module(LanggraphModule): # Read raw alert from Redis stream alert = self.read_message() if alert is None: - return + return Command( + update={"alert": {}}, + goto=END + ) # Example: For Splunk webhooks, uncomment the following line # alert = json.loads(alert["_raw"]) diff --git a/MODULES/ES-Rule-22-Phishing_user_report_mail.py b/MODULES/ES-Rule-22-Phishing_user_report_mail.py index 3de502c..4d6f563 100644 --- a/MODULES/ES-Rule-22-Phishing_user_report_mail.py +++ b/MODULES/ES-Rule-22-Phishing_user_report_mail.py @@ -28,7 +28,7 @@ class Module(BaseModule): # Get raw alert from stream alert = self.read_message() if alert is None: - return + return None # Parse data, this is an example of handling JSON data sent by Splunk Webhook # alert = json.loads(alert["_raw"]) @@ -40,7 +40,7 @@ class Module(BaseModule): alert["headers"] = headers self.agent_state.alert_raw = alert - return + return None def alert_analyze_node(self): client = Dify() diff --git a/MODULES/NDR-Rule-05-Suspect-C2-Communication.py b/MODULES/NDR-Rule-05-Suspect-C2-Communication.py index 32143ac..67ed77c 100644 --- a/MODULES/NDR-Rule-05-Suspect-C2-Communication.py +++ b/MODULES/NDR-Rule-05-Suspect-C2-Communication.py @@ -3,8 +3,9 @@ from enum import Enum from typing import Optional, Union, Dict, Any from langchain_core.messages import HumanMessage -from langgraph.graph import StateGraph +from langgraph.graph import StateGraph, END from langgraph.graph.state import CompiledStateGraph +from langgraph.types import Command from pydantic import BaseModel, Field, ConfigDict from Lib.api import get_current_time_str @@ -57,7 +58,10 @@ class Module(LanggraphModule): # Get raw alert from stream alert = self.read_message() if alert is None: - return + return Command( + update={"alert": {}}, + goto=END + ) state.alert = alert artifact: list = alert.get("artifact") diff --git a/PLUGINS/Redis/CONFIG.example.py b/PLUGINS/Redis/CONFIG.example.py index 31e7980..80ace75 100644 --- a/PLUGINS/Redis/CONFIG.example.py +++ b/PLUGINS/Redis/CONFIG.example.py @@ -4,4 +4,5 @@ # 根据你的环境修改Redis的连接地址 # 你通过 http://192.168.241.128:8001 和 default/redis-stack-password-for-agentic-soc-platform 登录 Redis Insight REDIS_URL = "redis://:redis-stack-password-for-agentic-soc-platform@192.168.241.128:6379/" -REDIS_STREAM_STORE_DAYS = 7 # 消息在Redis Stream中保存的天数 +REDIS_STREAM_MAX_LENGTH = 10000 # Redis Stream的最大长度 +REDIS_MAX_CONNECTIONS = 10 # Redis连接池大小,建议为模块数量的2倍 diff --git a/PLUGINS/Redis/redis_client.py b/PLUGINS/Redis/redis_client.py index c0ece9a..5a37b6c 100644 --- a/PLUGINS/Redis/redis_client.py +++ b/PLUGINS/Redis/redis_client.py @@ -1,22 +1,25 @@ import redis from Lib.log import logger -from PLUGINS.Redis.CONFIG import REDIS_URL +from PLUGINS.Redis.CONFIG import REDIS_URL, REDIS_MAX_CONNECTIONS class RedisClient(object): + # 类变量,用于存储全局唯一的连接池 + _stream_pool = None - def __init__(self): - pass + @classmethod + def get_stream_connection(cls): + """获取单例连接池中的连接""" + if cls._stream_pool is None: + logger.info("Initializing Global Redis Stream Connection Pool...") + cls._stream_pool = redis.ConnectionPool.from_url( + f"{REDIS_URL}0", + decode_responses=True, + max_connections=REDIS_MAX_CONNECTIONS, # 限制最大连接数 + health_check_interval=30 # 每30秒自动检测连接状态 + ) - @staticmethod - def get_stream_connection(): - """用于订阅类操作,无需使用连接池""" - redis_client = redis.Redis.from_url(f"{REDIS_URL}0", decode_responses=True) - # 测试连接 - try: - redis_client.ping() - return redis_client - except redis.ConnectionError as e: - logger.exception(e) - raise + # 使用同一个池创建 client 实例 + redis_client = redis.Redis(connection_pool=cls._stream_pool) + return redis_client diff --git a/PLUGINS/Redis/redis_stream_api.py b/PLUGINS/Redis/redis_stream_api.py index 22794a8..6090731 100644 --- a/PLUGINS/Redis/redis_stream_api.py +++ b/PLUGINS/Redis/redis_stream_api.py @@ -1,31 +1,47 @@ -import datetime import json +import time from typing import Dict, Any, Optional, List import redis from Lib.configs import REDIS_CONSUMER_GROUP, REDIS_CONSUMER_NAME from Lib.log import logger +from PLUGINS.Redis.CONFIG import REDIS_STREAM_MAX_LENGTH from PLUGINS.Redis.redis_client import RedisClient -class RedisStreamAPI: +class RedisStreamAPI(object): """ Redis Stream API封装类,提供消息发送和读取功能 """ + _instance = None + + def __new__(cls, *args, **kwargs): + """实现单例模式""" + if not cls._instance: + cls._instance = super(RedisStreamAPI, cls).__new__(cls) + cls._instance._initialized = False + return cls._instance def __init__(self): """初始化RedisStreamAPI类""" + if self._initialized: + return + self.redis_client = RedisClient.get_stream_connection() - def send_message(self, stream_key: str, message: Dict[str, Any]) -> Optional[str]: + self._checked_groups = set() + self._initialized = True + logger.info("RedisStreamAPI init finished.") + + def send_message(self, stream_key: str, message: Dict[str, Any], maxlen: int = REDIS_STREAM_MAX_LENGTH) -> Optional[str]: """ 发送消息到指定stream Args: stream_key (str): Redis stream的key名称 message (Dict[str, Any]): 要发送的消息内容 - + maxlen (int): stream最大长度,超过则删除最旧的消息,默认10000 Returns: Optional[str]: 发送成功返回消息ID,失败返回None """ @@ -34,7 +50,9 @@ class RedisStreamAPI: # 发送消息到stream message_id = self.redis_client.xadd( stream_key, - {"data": data} + {"data": data}, + maxlen=maxlen, + approximate=True ) return message_id @@ -43,62 +61,50 @@ class RedisStreamAPI: return None def read_message(self, stream_key: str, consumer_group: str = None, - consumer_name: str = None, timeout: int = 0, noack: bool = False) -> Optional[Dict[str, Any]]: + consumer_name: str = None, timeout: int = 5000) -> Optional[Dict[str, Any]]: """ 从指定stream读取一条消息 - + Args: - noack: stream_key (str): Redis stream的key名称 consumer_group (str): 消费者组名称,如果为None则使用默认配置 consumer_name (str): 消费者名称,如果为None则使用默认配置 - timeout (int): 读取超时时间(毫秒),默认5000毫秒 - - Returns: - Optional[Dict[str, Any]]: 读取到的消息,如果没有消息或出错则返回None + timeout (int): 阻塞等待时间,单位毫秒,默认5000ms """ - try: - if consumer_group is None: - consumer_group = REDIS_CONSUMER_GROUP - if consumer_name is None: - consumer_name = REDIS_CONSUMER_NAME + if consumer_group is None: + consumer_group = REDIS_CONSUMER_GROUP + if consumer_name is None: + consumer_name = REDIS_CONSUMER_NAME - # 确保消费者组存在 - flag = self._ensure_consumer_group(stream_key, consumer_group) - if not flag: - logger.error(f"无法确保消费者组 {consumer_group} 存在.") - return None - # 从消费者组读取消息 - messages = self.redis_client.xreadgroup( - consumer_group, - consumer_name, - {stream_key: '>'}, # '>' 表示只读取新消息 - count=1, - block=timeout, - noack=noack, - ) + # 确保消费者组存在(只在第一次调用时执行) + if stream_key not in self._checked_groups: + if self._ensure_consumer_group(stream_key, consumer_group): + self._checked_groups.add(stream_key) - if not messages or not messages[0][1]: - return None + while True: + try: + messages = self.redis_client.xreadgroup( + consumer_group, + consumer_name, + {stream_key: '>'}, + count=1, + block=timeout, + noack=True, + ) - # 解析消息 - stream_name, stream_messages = messages[0] - if not stream_messages: - return None + if not messages or not messages[0][1]: + continue - message_id, fields = stream_messages[0] + _, stream_messages = messages[0] + message_id, fields = stream_messages[0] + data = json.loads(fields["data"]) - value = fields["data"] - data = json.loads(value) + logger.info(f"Received: {stream_key} -> {message_id}") + return data - # 确认消息 - flag = self.redis_client.xack(stream_key, consumer_group, message_id) - - logger.info(f"{consumer_group} : {consumer_name} : {message_id}") - return data - except Exception as e: - logger.exception(e) - return None + except Exception as e: + logger.error(f"Error reading from stream {stream_key}: {e}") + time.sleep(1) # 发生异常(如网络闪断)时稍作停顿,防止死循环刷屏 def read_stream_from_start(self, stream_key, start_id='0-0'): """ @@ -199,43 +205,19 @@ class RedisStreamAPI: logger.exception(e) return [] - def _ensure_consumer_group(self, stream_key: str, consumer_group: str): + def _ensure_consumer_group(self, stream_key: str, consumer_group: str) -> bool: """ - 确保消费者组存在,如果不存在则创建 - - Args: - stream_key (str): Redis stream的key名称 - consumer_group (str): 消费者组名称 + 静默确保消费者组存在 """ try: - # 检查消费者组是否存在 - groups = self.redis_client.xinfo_groups(stream_key) - group_names = [group['name'] for group in groups] - - if consumer_group not in group_names: - # 创建消费者组 - self.redis_client.xgroup_create(stream_key, consumer_group, '$', mkstream=True) + # 直接创建,利用 mkstream=True。如果流不存在会自动创建流,如果组已存在会报错 + self.redis_client.xgroup_create(stream_key, consumer_group, id='$', mkstream=True) return True except redis.ResponseError as e: - # 如果流不存在,xinfo_groups会报错.捕获此错误并创建流和组. - if "no such key" in str(e).lower(): - try: - self.redis_client.xgroup_create(stream_key, consumer_group, '$', mkstream=True) - return True - except Exception as create_e: - if "BUSYGROUP Consumer Group name already exists" in str(e): - return True - else: - logger.exception(f"创建流 {stream_key} 和组 {consumer_group} 失败: {create_e}") - return False - elif "BUSYGROUP" in str(e): - # 消费者组已存在,这是正常情况 + # 如果报错信息包含 BUSYGROUP,说明组已经存在,属于正常情况 + if "BUSYGROUP" in str(e): return True - else: - logger.exception(f"检查或创建消费者组时发生未知Redis响应错误: {e}") - return False - except Exception as e: - logger.exception(e) + # logger.error(f"创建消费者组失败: {e}") return False def get_stream_info(self, stream_key: str) -> Optional[Dict[str, Any]]: @@ -281,29 +263,3 @@ class RedisStreamAPI: self.redis_client.close() except Exception as e: logger.exception(e) - - def clean_redis_stream(self, max_age_days=30): - """ - 清理Redis Stream中超过指定天数的键值对. - """ - # 计算最老允许的时间戳,单位为毫秒 - # Unix时间戳(秒)* 1000 - logger.info(f"Clean Redis Stream older than {max_age_days} days.") - cutoff_timestamp_ms = int((datetime.datetime.now() - datetime.timedelta(days=max_age_days)).timestamp() * 1000) - - try: - for key in self.redis_client.scan_iter(match='*'): - # 检查键的类型是否为 stream - if self.redis_client.type(key) == 'stream': - # 使用 XTRIM 命令删除早于给定ID的条目 - # ID的格式是 `unix_time_ms-sequence_number` - # 我们可以使用 `unix_time_ms-0` 作为删除的上限ID - trim_id = f'{cutoff_timestamp_ms}-0' - - # `XTRIM` 带有 `MINID` 选项,用于删除ID小于指定ID的所有条目 - trimmed_count = self.redis_client.xtrim(key, minid=trim_id) - logger.info(f"Delete {trimmed_count} items from '{key}'") - except Exception as e: - logger.exception(e) - finally: - logger.info("clean task finished.")