diff --git a/cli/process_registry.py b/cli/process_registry.py new file mode 100644 index 00000000..5a7a1bdb --- /dev/null +++ b/cli/process_registry.py @@ -0,0 +1,76 @@ +"""Track and clean up spawned CLI subprocesses. + +This is a safety net for cases where the server is interrupted (Ctrl+C) and the +FastAPI lifespan cleanup doesn't run to completion. We only track processes we +spawn so we don't accidentally kill unrelated system processes. +""" + +from __future__ import annotations + +import atexit +import logging +import os +import subprocess +import threading +from typing import Set + +logger = logging.getLogger(__name__) + +_lock = threading.Lock() +_pids: Set[int] = set() +_atexit_registered = False + + +def ensure_atexit_registered() -> None: + global _atexit_registered + with _lock: + if _atexit_registered: + return + atexit.register(kill_all_best_effort) + _atexit_registered = True + + +def register_pid(pid: int) -> None: + if not pid: + return + ensure_atexit_registered() + with _lock: + _pids.add(int(pid)) + + +def unregister_pid(pid: int) -> None: + if not pid: + return + with _lock: + _pids.discard(int(pid)) + + +def kill_all_best_effort() -> None: + """Kill any still-running registered pids (best-effort).""" + with _lock: + pids = list(_pids) + _pids.clear() + + if not pids: + return + + if os.name == "nt": + for pid in pids: + try: + # /T kills child processes, /F forces termination. + subprocess.run( + ["taskkill", "/PID", str(pid), "/T", "/F"], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + check=False, + ) + except Exception as e: + logger.debug("process_registry: taskkill failed pid=%s: %s", pid, e) + return + + # Best-effort fallback for non-Windows. + for pid in pids: + try: + os.kill(pid, 9) + except Exception as e: + logger.debug("process_registry: kill failed pid=%s: %s", pid, e) diff --git a/cli/session.py b/cli/session.py index 1a28552b..c28626d6 100644 --- a/cli/session.py +++ b/cli/session.py @@ -6,6 +6,8 @@ import json import logging from typing import AsyncGenerator, Optional, Dict, List, Any +from .process_registry import register_pid, unregister_pid + logger = logging.getLogger(__name__) @@ -102,6 +104,8 @@ class CLISession: cwd=self.workspace, env=env, ) + if self.process and self.process.pid: + register_pid(self.process.pid) if not self.process or not self.process.stdout: yield {"type": "exit", "code": 1} @@ -181,6 +185,8 @@ class CLISession: } finally: self._is_busy = False + if self.process and self.process.pid: + unregister_pid(self.process.pid) async def _handle_line_gen( self, line_str: str, session_id_extracted: bool @@ -236,6 +242,8 @@ class CLISession: except asyncio.TimeoutError: self.process.kill() await self.process.wait() + if self.process and self.process.pid: + unregister_pid(self.process.pid) return True except Exception as e: logger.error(f"Error stopping process: {e}") diff --git a/providers/nvidia_nim/client.py b/providers/nvidia_nim/client.py index 9269b5b4..199c1a50 100644 --- a/providers/nvidia_nim/client.py +++ b/providers/nvidia_nim/client.py @@ -118,6 +118,10 @@ class NvidiaNimProvider(BaseProvider): yield event block_idx = sse.blocks.allocate_index() + if tool_use.get("name") == "Task" and isinstance( + tool_use.get("input"), dict + ): + tool_use["input"]["run_in_background"] = False yield sse.content_block_start( block_idx, "tool_use", @@ -183,6 +187,10 @@ class NvidiaNimProvider(BaseProvider): id=tool_use["id"], name=tool_use["name"], ) + if tool_use.get("name") == "Task" and isinstance( + tool_use.get("input"), dict + ): + tool_use["input"]["run_in_background"] = False yield sse.content_block_delta( block_idx, "input_json_delta", @@ -199,6 +207,10 @@ class NvidiaNimProvider(BaseProvider): yield event yield sse.emit_text_delta(" ") + # Flush buffered Task args before closing tool blocks. + for event in self._flush_task_arg_buffers(sse): + yield event + for event in sse.close_all_blocks(): yield event @@ -245,10 +257,24 @@ class NvidiaNimProvider(BaseProvider): tc_index = len(sse.blocks.tool_indices) fn_delta = tc.get("function", {}) - if fn_delta.get("name") is not None: - sse.blocks.tool_names[tc_index] = ( - sse.blocks.tool_names.get(tc_index, "") + fn_delta["name"] - ) + incoming_name = fn_delta.get("name") + if incoming_name is not None: + # Some providers stream tool names as fragments; others resend the full name. + # Avoid "TaskTask" while still supporting fragment streams. + prev = sse.blocks.tool_names.get(tc_index, "") + if not prev: + sse.blocks.tool_names[tc_index] = incoming_name + elif prev == incoming_name: + pass + elif isinstance(prev, str) and isinstance(incoming_name, str): + if incoming_name.startswith(prev): + sse.blocks.tool_names[tc_index] = incoming_name + elif prev.startswith(incoming_name): + pass + else: + sse.blocks.tool_names[tc_index] = prev + incoming_name + else: + sse.blocks.tool_names[tc_index] = str(prev) + str(incoming_name) if tc_index not in sse.blocks.tool_indices: name = sse.blocks.tool_names.get(tc_index, "") @@ -273,20 +299,76 @@ class NvidiaNimProvider(BaseProvider): yield sse.start_tool_block(tc_index, tool_id, name) sse.blocks.tool_started[tc_index] = True - # INTERCEPTION: If this is a Task tool, force background=False current_name = sse.blocks.tool_names.get(tc_index, "") + # INTERCEPTION: Task args can stream in many partial chunks. Buffer until we + # have valid JSON, then emit a single delta with run_in_background forced off. if current_name == "Task": - try: - args_json = json.loads(args) + # Allow older tests to pass with MagicMock'd builders by lazily creating fields. + if not isinstance(getattr(sse.blocks, "task_arg_buffer", None), dict): + sse.blocks.task_arg_buffer = {} + if not isinstance(getattr(sse.blocks, "task_args_emitted", None), dict): + sse.blocks.task_args_emitted = {} + if not isinstance(getattr(sse.blocks, "tool_ids", None), dict): + sse.blocks.tool_ids = {} + + if not sse.blocks.task_args_emitted.get(tc_index, False): + buf = sse.blocks.task_arg_buffer.get(tc_index, "") + args + sse.blocks.task_arg_buffer[tc_index] = buf + try: + args_json = json.loads(buf) + except Exception: + return if args_json.get("run_in_background") is not False: logger.info( - f"NIM_INTERCEPT: Forcing run_in_background=False for Task {tc.get('id', 'unknown')}" + "NIM_INTERCEPT: Forcing run_in_background=False for Task %s", + ( + tc.get("id") + or sse.blocks.tool_ids.get(tc_index, "unknown") + ), ) args_json["run_in_background"] = False - args = json.dumps(args_json) - except Exception as e: - logger.warning( - f"NIM_INTERCEPT: Failed to parse/modify Task args: {e}" - ) + sse.blocks.task_args_emitted[tc_index] = True + sse.blocks.task_arg_buffer.pop(tc_index, None) + yield sse.emit_tool_delta(tc_index, json.dumps(args_json)) + return yield sse.emit_tool_delta(tc_index, args) + + def _flush_task_arg_buffers(self, sse: Any): + """Emit buffered Task args as a single JSON delta (best-effort).""" + if not isinstance(getattr(sse.blocks, "task_arg_buffer", None), dict): + return + if not isinstance(getattr(sse.blocks, "task_args_emitted", None), dict): + sse.blocks.task_args_emitted = {} + if not isinstance(getattr(sse.blocks, "tool_ids", None), dict): + sse.blocks.tool_ids = {} + # Iterate over a copy; we will mutate dicts. + for tool_index, buf in list(getattr(sse.blocks, "task_arg_buffer", {}).items()): + if sse.blocks.task_args_emitted.get(tool_index, False): + sse.blocks.task_arg_buffer.pop(tool_index, None) + continue + + tool_id = sse.blocks.tool_ids.get(tool_index, "unknown") + out = "{}" + try: + args_json = json.loads(buf) + if args_json.get("run_in_background") is not False: + logger.info( + "NIM_INTERCEPT: Forcing run_in_background=False for Task %s", + tool_id, + ) + args_json["run_in_background"] = False + out = json.dumps(args_json) + except Exception as e: + prefix = buf[:120] + logger.warning( + "NIM_INTERCEPT: Task args invalid JSON (id=%s len=%d prefix=%r): %s", + tool_id, + len(buf), + prefix, + e, + ) + + sse.blocks.task_args_emitted[tool_index] = True + sse.blocks.task_arg_buffer.pop(tool_index, None) + yield sse.emit_tool_delta(tool_index, out) diff --git a/providers/nvidia_nim/utils/sse_builder.py b/providers/nvidia_nim/utils/sse_builder.py index 845ed707..e0def469 100644 --- a/providers/nvidia_nim/utils/sse_builder.py +++ b/providers/nvidia_nim/utils/sse_builder.py @@ -42,7 +42,11 @@ class ContentBlockManager: tool_indices: Dict[int, int] = field(default_factory=dict) tool_contents: Dict[int, str] = field(default_factory=dict) tool_names: Dict[int, str] = field(default_factory=dict) + tool_ids: Dict[int, str] = field(default_factory=dict) tool_started: Dict[int, bool] = field(default_factory=dict) + # Buffer streaming args for tools where we don't want to emit partial deltas. + task_arg_buffer: Dict[int, str] = field(default_factory=dict) + task_args_emitted: Dict[int, bool] = field(default_factory=dict) def allocate_index(self) -> int: """Allocate and return the next block index.""" @@ -200,6 +204,8 @@ class SSEBuilder: block_idx = self.blocks.allocate_index() self.blocks.tool_indices[tool_index] = block_idx self.blocks.tool_contents[tool_index] = "" + self.blocks.tool_ids[tool_index] = tool_id + self.blocks.task_args_emitted.setdefault(tool_index, False) return self.content_block_start(block_idx, "tool_use", id=tool_id, name=name) def emit_tool_delta(self, tool_index: int, partial_json: str) -> str: diff --git a/server.py b/server.py index 0e72e2d3..244e703c 100644 --- a/server.py +++ b/server.py @@ -12,6 +12,11 @@ __all__ = ["app", "create_app"] if __name__ == "__main__": import uvicorn from config.settings import get_settings + from cli.process_registry import kill_all_best_effort settings = get_settings() - uvicorn.run(app, host=settings.host, port=settings.port, log_level="debug") + try: + uvicorn.run(app, host=settings.host, port=settings.port, log_level="debug") + finally: + # Safety net for Ctrl+C cases where lifespan shutdown doesn't fully run. + kill_all_best_effort() diff --git a/tests/test_process_registry.py b/tests/test_process_registry.py new file mode 100644 index 00000000..b97b6366 --- /dev/null +++ b/tests/test_process_registry.py @@ -0,0 +1,34 @@ +import os + + +def test_process_registry_register_unregister_does_not_crash(): + from cli import process_registry as pr + + pr.register_pid(12345) + pr.unregister_pid(12345) + + +def test_process_registry_kill_all_best_effort_empty_is_noop(): + from cli import process_registry as pr + + # Ensure no exception on empty set + pr.kill_all_best_effort() + + +def test_process_registry_kill_all_best_effort_windows_noop_when_taskkill_missing( + monkeypatch, +): + from cli import process_registry as pr + + # Simulate windows path in a stable way. + monkeypatch.setattr(pr, "_pids", {12345}) + monkeypatch.setattr(os, "name", "nt", raising=False) + + # If taskkill isn't callable, we still should not crash. + import subprocess + + def _boom(*args, **kwargs): + raise FileNotFoundError("taskkill missing") + + monkeypatch.setattr(subprocess, "run", _boom) + pr.kill_all_best_effort() diff --git a/tests/test_streaming_errors.py b/tests/test_streaming_errors.py index 298d59fc..e3920a6e 100644 --- a/tests/test_streaming_errors.py +++ b/tests/test_streaming_errors.py @@ -316,8 +316,33 @@ class TestProcessToolCall: # The intercepted args should have run_in_background=false assert "false" in event_text.lower() - def test_task_tool_invalid_json_logs_warning(self): - """Invalid JSON args for Task tool doesn't crash.""" + def test_task_tool_chunked_args_forces_background_false(self): + """Chunked Task args are buffered until valid JSON, then forced to false.""" + provider = _make_provider() + from providers.nvidia_nim.utils import SSEBuilder + + sse = SSEBuilder("msg_test", "test-model") + tc1 = { + "index": 0, + "id": "call_task_chunked", + "function": {"name": "Task", "arguments": '{"run_in_background": true,'}, + } + tc2 = { + "index": 0, + "id": "call_task_chunked", + "function": {"name": None, "arguments": ' "prompt": "test"}'}, + } + + events1 = list(provider._process_tool_call(tc1, sse)) + assert len(events1) > 0 + assert "false" not in "".join(events1).lower() + + events2 = list(provider._process_tool_call(tc2, sse)) + event_text = "".join(events1 + events2) + assert "false" in event_text.lower() + + def test_task_tool_invalid_json_logs_warning_on_flush(self, caplog): + """Invalid JSON args for Task tool emits {} on flush and logs a warning.""" provider = _make_provider() from providers.nvidia_nim.utils import SSEBuilder @@ -327,10 +352,17 @@ class TestProcessToolCall: "id": "call_task2", "function": {"name": "Task", "arguments": "not json"}, } - # Should not raise events = list(provider._process_tool_call(tc, sse)) assert len(events) > 0 + with caplog.at_level("WARNING"): + flushed = list(provider._flush_task_arg_buffers(sse)) + assert len(flushed) > 0 + assert "{}" in "".join(flushed) + assert any( + "NIM_INTERCEPT: Task args invalid JSON" in r.message for r in caplog.records + ) + def test_negative_tool_index_fallback(self): """tc_index < 0 uses len(tool_indices) as fallback.""" provider = _make_provider()