fixed NIM_INTERCEPT for chunked tool calls

This commit is contained in:
Alishahryar1
2026-02-14 03:25:32 -08:00
parent d2e6e52742
commit 4b95429c32
7 changed files with 260 additions and 17 deletions
+76
View File
@@ -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)
+8
View File
@@ -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}")
+95 -13
View File
@@ -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)
@@ -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:
+6 -1
View File
@@ -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()
+34
View File
@@ -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()
+35 -3
View File
@@ -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()