-
Notifications
You must be signed in to change notification settings - Fork 64
Expand file tree
/
Copy pathmain.py
More file actions
465 lines (395 loc) · 17.2 KB
/
Copy pathmain.py
File metadata and controls
465 lines (395 loc) · 17.2 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
import os
import sys
import argparse
import logging
import threading # Added for MQ Consumer
import atexit # Added for graceful shutdown
import traceback # Added for error logging
from sqlalchemy import text
from flask import Flask, jsonify, render_template, redirect, url_for, request, make_response
from flask_socketio import SocketIO
from flask_migrate import Migrate
from flask_jwt_extended import JWTManager, jwt_required, get_jwt_identity, verify_jwt_in_request, get_jwt, decode_token
from functools import wraps
from dotenv import load_dotenv
from app.models import db
from app.utils.logging_config import configure_logging
from app.models.models import Prompt, User
from app.prompts.default_prompts import DEFAULT_PROMPTS
from app.utils.mq_consumer import RabbitMQConsumer # Added MQ Consumer
from app import __version__, get_version, get_version_info
import sys
# 配置日志
logger = configure_logging()
# 显示版本信息
def print_version_info():
"""打印版本信息"""
version_info = get_version_info()
version_info['python_version'] = f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}"
print("=" * 60)
print(f"🚀 DeepSOC - AI-Powered Security Operations Center")
print("=" * 60)
print(f"版本: {version_info['version']}")
print(f"发布名称: {version_info['release_name']}")
print(f"构建日期: {version_info['build_date']}")
print(f"Python 版本: {version_info['python_version']}")
print(f"描述: {version_info['description']}")
print("=" * 60)
logger.info(f"DeepSOC v{version_info['version']} 启动")
# 加载环境变量
load_dotenv(override=True)
# 创建Flask应用
app = Flask(__name__,
static_folder='app/static',
template_folder='app/templates')
app.config['SECRET_KEY'] = os.getenv('SECRET_KEY', 'deepsoc_secret_key')
app.config['SQLALCHEMY_DATABASE_URI'] = os.getenv('DATABASE_URL', 'mysql+pymysql://root:123456@localhost:3306/deepsoc')
app.config['SQLALCHEMY_TRACK_MODIFICATIONS'] = False
# JWT配置
app.config['JWT_SECRET_KEY'] = os.getenv('JWT_SECRET_KEY', 'deepsoc_jwt_secret_key')
app.config['JWT_ACCESS_TOKEN_EXPIRES'] = int(os.getenv('JWT_ACCESS_TOKEN_EXPIRES', 86400))
app.config['JWT_TOKEN_LOCATION'] = ['headers', 'cookies']
app.config['JWT_COOKIE_CSRF_PROTECT'] = False
# 初始化JWT
jwt = JWTManager(app)
# 初始化数据库
db.init_app(app)
# 初始化迁移
migrate = Migrate(app, db)
# 初始化SocketIO
socketio = SocketIO(
app,
cors_allowed_origins="*",
async_mode='threading',
ping_timeout=60,
ping_interval=25,
engineio_logger=os.getenv('ENGINEIO_LOGGER', 'False').lower() == 'true', # Control via env var
logger=os.getenv('SOCKETIO_LOGGER', 'False').lower() == 'true', # Control via env var
manage_session=False,
always_connect=True,
max_http_buffer_size=int(os.getenv('SOCKETIO_MAX_HTTP_BUFFER_SIZE', 100000000)) # Control via env var (1e8)
)
# 导入路由
from app.controllers.event_controller import event_bp
app.register_blueprint(event_bp, url_prefix='/api/event')
from app.controllers.auth_controller import auth_bp
app.register_blueprint(auth_bp, url_prefix='/api/auth')
from app.controllers.prompt_controller import prompt_bp
app.register_blueprint(prompt_bp, url_prefix='/api/prompt')
from app.controllers.state_controller import state_bp
app.register_blueprint(state_bp, url_prefix='/api/state')
from app.controllers.user_controller import user_bp
app.register_blueprint(user_bp, url_prefix='/api/user')
from app.controllers.engineer_chat_api import engineer_chat_bp
app.register_blueprint(engineer_chat_bp, url_prefix='/api/engineer-chat')
from app.controllers.socket_controller import register_socket_events
register_socket_events(socketio)
# --- RabbitMQ Consumer Setup ---
mq_consumer_thread = None
mq_consumer = None
def handle_mq_message_to_socketio(message_data):
"""Callback function to process messages from RabbitMQ and emit them via SocketIO."""
event_id = message_data.get('event_id')
message_type = message_data.get('message_type', 'generic_notification') # Default type if not present
message_id = message_data.get('message_id', 'N/A')
if not event_id:
logger.warning(f"MQ Consumer: Received message (ID: {message_id}) without event_id. Cannot route to SocketIO room. Message: {message_data}")
return
# The message_data is already a dict from create_standard_message().to_dict()
# It should be directly usable by the frontend if it expects the Message model structure.
# Determine the SocketIO event name.
# For now, using 'new_message' for all, as previously used by broadcast_message.
# This could be made more specific based on message_type if frontend handles different events.
socketio_event_name = 'new_message'
logger.debug(f"MQ Consumer: Relaying message (ID: {message_id}, Type: {message_type}) to SocketIO room '{event_id}' for event '{socketio_event_name}'")
try:
# Emit with app.app_context() to ensure context for operations like url_for if used by SocketIO internals
# although socketio.emit itself is generally thread-safe and handles context for its own operations.
with app.app_context():
socketio.emit(socketio_event_name, message_data, room=event_id)
logger.debug(f"MQ Consumer: Successfully emitted message ID {message_id} to room {event_id}.")
except Exception as e:
logger.error(f"MQ Consumer: Error emitting message ID {message_id} to SocketIO room '{event_id}': {e}")
logger.error(traceback.format_exc())
def start_rabbitmq_consumer():
global mq_consumer, mq_consumer_thread
logger.debug("Initializing RabbitMQ consumer...")
# 明确传递环境变量参数,确保使用正确的配置
rabbitmq_host = os.getenv('RABBITMQ_HOST', '127.0.0.1')
rabbitmq_port = int(os.getenv('RABBITMQ_PORT', 5672))
rabbitmq_user = os.getenv('RABBITMQ_USER', 'guest')
rabbitmq_password = os.getenv('RABBITMQ_PASSWORD', 'guest')
rabbitmq_vhost = os.getenv('RABBITMQ_VHOST', '/')
mq_consumer = RabbitMQConsumer(
host=rabbitmq_host,
port=rabbitmq_port,
username=rabbitmq_user,
password=rabbitmq_password,
virtual_host=rabbitmq_vhost
)
# Start consuming in a separate thread
# The start_consuming method is blocking, so it needs its own thread.
mq_consumer_thread = threading.Thread(
target=mq_consumer.start_consuming,
args=(handle_mq_message_to_socketio,),
name="RabbitMQConsumerThread",
daemon=True # Daemon thread will exit when the main program exits
)
mq_consumer_thread.start()
logger.debug("RabbitMQ consumer thread started.")
def stop_rabbitmq_consumer():
global mq_consumer, mq_consumer_thread
if mq_consumer:
logger.debug("Stopping RabbitMQ consumer...")
mq_consumer.stop_consuming()
if mq_consumer_thread and mq_consumer_thread.is_alive():
logger.debug("Waiting for RabbitMQ consumer thread to join...")
mq_consumer_thread.join(timeout=10) # Wait for up to 10 seconds
if mq_consumer_thread.is_alive():
logger.warning("RabbitMQ consumer thread did not join in time.")
else:
logger.debug("RabbitMQ consumer thread joined successfully.")
logger.debug("RabbitMQ consumer stopped.")
# Register the cleanup function to be called on exit
atexit.register(stop_rabbitmq_consumer)
# --- End RabbitMQ Consumer Setup ---
# 定义登录检查装饰器
def login_required(f):
@wraps(f)
def decorated_function(*args, **kwargs):
# 从多个来源获取Token
token = None
# 1. 尝试从Authorization头部获取
auth_header = request.headers.get('Authorization')
if auth_header and auth_header.startswith('Bearer '):
token = auth_header.split(' ')[1]
# 2. 尝试从Cookie获取
if not token:
token = request.cookies.get('access_token')
# 3. 尝试从查询参数获取(不推荐,但为了兼容)
if not token:
token = request.args.get('access_token')
if token:
try:
# 验证令牌有效性
app.config['JWT_TOKEN_LOCATION'] = ['headers'] # 临时设置为仅从头部获取
app.config['JWT_HEADER_NAME'] = 'Authorization'
app.config['JWT_HEADER_TYPE'] = 'Bearer'
# 手动解码Token
jwt_data = decode_token(token)
identity = jwt_data['sub']
# 如果到这里没有异常,则用户已登录
return f(*args, **kwargs)
except Exception as e:
logger.error(f"Token验证失败: {str(e)}")
pass # Token无效,继续执行重定向逻辑
# Token不存在或无效
if request.path.startswith('/api/'):
return jsonify({
'status': 'error',
'message': '请先登录',
'authenticated': False
}), 401
# 如果是页面请求,重定向到登录页
return redirect(url_for('login'))
return decorated_function
@app.route('/')
def index():
return render_template('index.html')
@app.route('/login')
def login():
return render_template('login.html')
@app.route('/warroom/<event_id>')
@login_required
def warroom(event_id):
return render_template('warroom.html', event_id=event_id)
@app.route('/settings/prompts')
@login_required
def prompt_settings():
return render_template('prompt_management.html')
@app.route('/settings/background-security')
@login_required
def background_security():
return render_template('background_security.html')
@app.route('/settings/soar-playbooks')
@login_required
def soar_playbooks():
return render_template('soar_playbooks.html')
@app.route('/settings/mcp-tools')
@login_required
def mcp_tools():
return render_template('mcp_tools.html')
@app.route('/user-management')
@login_required
def user_management():
return render_template('user_management.html')
@app.route('/change-password')
@login_required
def change_password_page():
return render_template('change_password.html')
@app.route('/health')
def health():
return jsonify({
'status': 'success',
'message': 'DeepSOC API is healthy'
})
@app.route('/api/version')
def api_version():
"""获取系统版本信息API"""
version_info = get_version_info()
version_info['python_version'] = f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}"
return jsonify({
'status': 'success',
'data': version_info
})
def create_tables():
"""重新创建数据库表,确保结构最新"""
with app.app_context():
db.drop_all()
db.create_all()
logger.info("数据库表重新创建成功")
def import_sql_file(sql_path: str = "initial_data.sql"):
"""Import initial data from a SQL file if it exists."""
if not os.path.exists(sql_path):
logger.warning(f"SQL file {sql_path} not found, skipping import")
return
with app.app_context():
engine = db.engine
with engine.begin() as connection:
sql_content = open(sql_path, "r", encoding="utf-8").read()
for statement in sql_content.split(";"):
stmt = statement.strip()
if stmt:
# Use exec_driver_sql to avoid SQLAlchemy interpreting
# colon-prefixed values inside JSON as bind parameters.
connection.exec_driver_sql(stmt)
logger.info(f"初始数据 {sql_path} 导入完成")
def create_default_prompts():
"""Load built-in prompt content into the database if not already present."""
with app.app_context():
for name, content in DEFAULT_PROMPTS.items():
prompt = Prompt.query.filter_by(name=name).first()
if not prompt:
prompt = Prompt(name=name)
db.session.add(prompt)
if not prompt.content:
prompt.content = content
db.session.commit()
logger.info("默认提示词导入完成")
def create_admin_user():
"""创建默认管理员用户"""
with app.app_context():
# 检查admin用户是否已存在
existing_admin = User.query.filter_by(username='admin').first()
if existing_admin:
logger.info("管理员用户已存在,跳过创建")
return
# 创建admin用户
admin_user = User(
username='admin',
nickname='管理员',
email='admin@deepsoc.local',
phone='18999990000',
role='admin',
is_active=True
)
admin_user.set_password('admin123')
db.session.add(admin_user)
db.session.commit()
logger.info("管理员用户创建成功")
logger.info("用户名: admin")
logger.info("密码: admin123")
logger.info("邮箱: admin@deepsoc.local")
def start_agent(role):
"""启动特定角色的Agent"""
logger.info(f"启动 {role} Agent")
if role == '_captain':
from app.services.captain_service import run_captain
run_captain()
elif role == '_manager':
from app.services.manager_service import run_manager
run_manager()
elif role == '_operator':
from app.services.operator_service import run_operator
run_operator()
elif role == '_executor':
from app.services.executor_service import run_executor
run_executor()
elif role == '_expert':
from app.services.expert_service import run_expert
run_expert()
else:
logger.error(f"未知角色: {role}")
sys.exit(1)
if __name__ == '__main__':
# 显示版本信息
print_version_info()
parser = argparse.ArgumentParser(description='DeepSOC - AI驱动的安全运营中心')
parser.add_argument('-role', type=str, help='Agent角色: _captain, _manager, _operator, _executor, _expert')
parser.add_argument('-init', action='store_true', help='系统初始化:数据库表 + 提示词 + 管理员用户')
parser.add_argument('-init-with-demo', action='store_true', help='完整初始化:数据库 + 演示数据(推荐)')
parser.add_argument('-load_demo', action='store_true', help='仅加载演示数据(需要已存在的数据库)')
parser.add_argument('-version', action='store_true', help='显示版本信息')
args = parser.parse_args()
if args.version:
sys.exit(0)
if getattr(args, 'init_with_demo', False):
logger.info("开始完整初始化(包含演示数据)...")
create_tables()
create_default_prompts()
create_admin_user()
logger.info("系统初始化完成 - 数据库表、提示词、管理员用户已创建")
logger.info("开始加载演示数据...")
import_sql_file("sql_data/initial_data.sql")
logger.info("演示数据加载完成")
logger.info("完整初始化全部完成")
sys.exit(0)
if args.init:
logger.info("开始系统初始化...")
create_tables()
create_default_prompts()
create_admin_user()
logger.info("系统初始化完成 - 数据库表、提示词、管理员用户已创建")
sys.exit(0)
if args.load_demo:
logger.info("开始加载演示数据...")
import_sql_file("sql_data/initial_data.sql")
logger.info("演示数据加载完成")
sys.exit(0)
if args.role:
# When running as an agent, do not start the MQ consumer or web server.
start_agent(args.role)
else:
# This is the main web server process
logger.info("Starting DeepSOC Web Server and services...")
# 先检查端口可用性,避免启动后台线程后才发现端口冲突
host = os.getenv('LISTEN_HOST', '0.0.0.0')
port = int(os.getenv('LISTEN_PORT', 5007))
try:
import socket
# 测试端口是否可用
test_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
test_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
test_socket.bind((host, port))
test_socket.close()
logger.info(f"Port {port} is available, starting services...")
except OSError as e:
logger.error(f"Port {port} is not available: {e}")
logger.error("Please stop the existing process or use a different port.")
sys.exit(1)
# 端口可用后再启动RabbitMQ consumer
start_rabbitmq_consumer()
# 启动Web服务器
try:
socketio.run(
app,
host=host,
port=port,
debug=(os.getenv('FLASK_DEBUG', 'False').lower() == 'true'), # Control via env var
use_reloader=False, # Important: reloader can cause issues with threads and SocketIO
allow_unsafe_werkzeug=True # Allow for development and testing
)
except Exception as e:
logger.error(f"Failed to start web server: {e}")
stop_rabbitmq_consumer()
sys.exit(1)