mirror of
https://github.com/Alishahryar1/free-claude-code.git
synced 2026-07-03 14:05:26 +02:00
fixed NIM_INTERCEPT for chunked tool calls
This commit is contained in:
@@ -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)
|
||||
@@ -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}")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user