Files
geMoldInsight/src/moldinsight/services/occ_process_pool.py
T

145 lines
5.8 KiB
Python
Raw Normal View History

# services/occ_process_pool.py
"""常驻 OCC 工作进程池(方案 B,见 docs/topics/performance/OCC_THROUGHPUT.md)。
替代原进程内 `ThreadPoolExecutor(max_workers=1)`:
- 每个工作进程是一个独立 OCC 通道(OCC 非线程安全,通道内串行),常驻不随任务拉起
(spawn 下 import OCC 秒级,按任务拉起会把开销摊到每个任务上)
- 超时/崩溃 = terminate() 换新进程补位——进程边界干净回收(线程级无法击杀 C++ 栈,
旧方案每次超时滞留 1 个线程)
- 输入输出走文件路径 + 普通字典,杜绝 pickle OCC 对象(见 core/occ_worker.py)
"""
import asyncio
import multiprocessing
from typing import Any, Dict, Optional
from moldinsight.core.occ_worker import worker_main
from shared.utils.logger import get_logger
logger = get_logger(__name__)
class _OccWorker:
"""单个 OCC 工作进程的父进程侧封装。"""
def __init__(self, process, conn):
self.process = process
self.conn = conn
self.lock = asyncio.Lock() # 通道串行:同一进程同时只有一个操作在途
async def run(self, op_name: str, payload: Dict[str, Any], timeout: float):
# 阻塞式管道收发放 asyncio.to_thread,不卡事件循环;
# 超时后父进程 terminate() 子进程 → 管道 EOF → 该线程 recv 立即返回,无泄漏
result = await asyncio.wait_for(
asyncio.to_thread(self._run_blocking, op_name, payload),
timeout=timeout,
)
return result
def _run_blocking(self, op_name: str, payload: Dict[str, Any]):
self.conn.send((op_name, payload))
status, data = self.conn.recv()
if status == "error":
raise RuntimeError(data.get("error") or "OCC 子进程操作失败")
return data
class OccProcessPool:
"""OCC 工作进程池(默认 1 进程 = 1 串行通道,与旧单线程语义一致)。
池大小与 celery 并发解耦(celery 并发走多 worker 子进程,每个持自己的池)。
"""
def __init__(self, size: int = 1):
if size < 1:
raise ValueError("size 必须 >= 1")
self._size = size
self._ctx = multiprocessing.get_context("spawn")
self._workers: list[_OccWorker] = []
self._rr = 0
# 保护 workers 列表与轮转指针;操作执行期不持锁
self._pool_lock = asyncio.Lock()
# 任务级整体超时(process_file_with_storage 外层 wait_for)时的在途 worker 追踪
self._busy: Optional[_OccWorker] = None
async def run(self, op_name: str, payload: Dict[str, Any], timeout: float = 600):
for attempt in range(3):
async with self._pool_lock:
await self._ensure_started()
worker = self._workers[self._rr]
self._rr = (self._rr + 1) % len(self._workers)
self._busy = worker
try:
async with worker.lock:
return await worker.run(op_name, payload, timeout)
except asyncio.TimeoutError:
await self._replace(worker)
raise
except Exception as exc:
if not worker.process.is_alive():
# OCC segfault 等进程死亡:换新补位后重试该操作
logger.warning(f"OCC 工作进程异常退出,重试操作 {op_name}: {exc}")
await self._replace(worker)
continue
raise
finally:
if self._busy is worker:
self._busy = None
raise RuntimeError(f"OCC 工作进程连续异常,操作 {op_name} 未能完成")
async def recover(self):
"""任务级整体超时恢复:重建整个池,丢弃可能正卡在挂死 OCC 操作上的进程。
单通道池重建代价可忽略;重建后下次 run 自动按需补拉。
"""
async with self._pool_lock:
for worker in self._workers:
await self._terminate(worker)
self._workers = []
self._busy = None
logger.warning("OCC 进程池已整体重建(任务级超时恢复)")
async def shutdown(self):
async with self._pool_lock:
for worker in self._workers:
await self._terminate(worker)
self._workers = []
self._busy = None
async def _ensure_started(self):
if not self._workers:
for _ in range(self._size):
self._workers.append(self._spawn_one())
logger.info(f"OCC 进程池已启动: {self._size} 个常驻工作进程")
def _spawn_one(self) -> _OccWorker:
parent_conn, child_conn = self._ctx.Pipe(duplex=True)
proc = self._ctx.Process(target=worker_main, args=(child_conn,), daemon=True)
proc.start()
child_conn.close() # 父进程侧关闭子端,只留 parent_conn
return _OccWorker(proc, parent_conn)
async def _replace(self, worker: _OccWorker):
async with self._pool_lock:
await self._terminate(worker)
new_worker = self._spawn_one()
try:
idx = self._workers.index(worker)
except ValueError:
# 已被并发重建移除,新进程追加补位
self._workers.append(new_worker)
else:
self._workers[idx] = new_worker
logger.warning("OCC 工作进程已重建补位")
@staticmethod
async def _terminate(worker: _OccWorker):
try:
worker.process.terminate()
worker.process.join(timeout=5)
except Exception as exc:
logger.warning(f"终止 OCC 工作进程异常: {exc}")
try:
worker.conn.close()
except Exception:
pass