diff --git a/think/cortex.py b/think/cortex.py index 5171c1fc7..d9205cd05 100644 --- a/think/cortex.py +++ b/think/cortex.py @@ -20,6 +20,7 @@ from __future__ import annotations import json import logging import os +import signal import subprocess import sys import threading @@ -58,16 +59,34 @@ class TalentProcess: if self.process.poll() is None: # First try SIGTERM for graceful shutdown - self.process.terminate() + try: + self.process.terminate() + except ProcessLookupError: + pass + self._signal_process_group(signal.SIGTERM) try: self.process.wait(timeout=10) # Give more time for graceful shutdown except subprocess.TimeoutExpired: logging.getLogger(__name__).warning( f"Talent {self.use_id} didn't stop gracefully, killing" ) - self.process.kill() + self._signal_process_group(signal.SIGKILL) + try: + self.process.kill() + except ProcessLookupError: + pass self.process.wait() # Ensure zombie is reaped + def _signal_process_group(self, sig: int) -> None: + try: + pgid = os.getpgid(self.process.pid) + except ProcessLookupError: + return + try: + os.killpg(pgid, sig) + except ProcessLookupError: + return + class CortexService: """Callosum-based talent process manager.""" @@ -321,6 +340,7 @@ class CortexService: env=env, bufsize=1, cwd=subprocess_cwd, + start_new_session=True, ) # Send input and close stdin diff --git a/think/talents.py b/think/talents.py index 5696820ad..1e3b597cc 100644 --- a/think/talents.py +++ b/think/talents.py @@ -18,6 +18,7 @@ import asyncio import json import logging import os +import signal import sys import traceback from datetime import datetime @@ -1239,6 +1240,16 @@ async def main_async() -> None: app_logger = setup_logging(args.verbose) event_writer = JSONEventWriter(None) + loop = asyncio.get_running_loop() + main_task = asyncio.current_task() + registered_signals: list[signal.Signals] = [] + if main_task: + for sig in (signal.SIGTERM, signal.SIGINT): + try: + loop.add_signal_handler(sig, main_task.cancel) + registered_signals.append(sig) + except (NotImplementedError, RuntimeError): + LOG.debug("Signal handler registration unavailable for %s", sig) def emit_event(data: Event) -> None: if "ts" not in data: @@ -1301,6 +1312,8 @@ async def main_async() -> None: emit_event(err) raise finally: + for sig in registered_signals: + loop.remove_signal_handler(sig) event_writer.close()