refactor(tools): simplify RPC token branch and kernel stderr/buffer loops
This commit is contained in:
@@ -109,14 +109,11 @@ def _rpc_server_loop(server_sock: socket.socket, task_id: str, tool_call_log: li
|
||||
except (json.JSONDecodeError, UnicodeDecodeError) as exc:
|
||||
resp = tool_error(f"Invalid RPC request: {exc}")
|
||||
else:
|
||||
if not _rpc_token_ok(request, rpc_token):
|
||||
resp = tool_error("Unauthorized RPC request")
|
||||
else:
|
||||
resp = _handle_rpc_request(
|
||||
request, allowed_tools=allowed_tools, tool_call_counter=tool_call_counter,
|
||||
max_tool_calls=max_tool_calls, dispatch=dispatch, tool_call_log=tool_call_log,
|
||||
call_start=call_start, where="sandbox",
|
||||
)
|
||||
resp = _handle_rpc_request(
|
||||
request, allowed_tools=allowed_tools, tool_call_counter=tool_call_counter,
|
||||
max_tool_calls=max_tool_calls, dispatch=dispatch, tool_call_log=tool_call_log,
|
||||
call_start=call_start, where="sandbox",
|
||||
) if _rpc_token_ok(request, rpc_token) else tool_error("Unauthorized RPC request")
|
||||
conn.sendall((resp + "\n").encode())
|
||||
except socket.timeout:
|
||||
logger.debug("RPC listener socket timeout")
|
||||
@@ -146,10 +143,8 @@ def _rpc_poll_loop(env, rpc_dir: str, task_id: str, tool_call_log: list, tool_ca
|
||||
if not output:
|
||||
stop_event.wait(poll_interval)
|
||||
continue
|
||||
req_files = sorted(
|
||||
f for f in (line.strip() for line in output.split("\n"))
|
||||
if f and not f.endswith(".tmp") and "/req_" in f
|
||||
)
|
||||
req_files = sorted(f for f in (line.strip() for line in output.split("\n"))
|
||||
if f and not f.endswith(".tmp") and "/req_" in f)
|
||||
for req_file in req_files:
|
||||
if stop_event.is_set():
|
||||
break
|
||||
|
||||
@@ -202,11 +202,10 @@ class _BoundedBuffer:
|
||||
self.total = 0
|
||||
|
||||
def append(self, data: bytes, cap: int) -> None:
|
||||
if self.total >= cap:
|
||||
return
|
||||
keep = data[: cap - self.total]
|
||||
self.chunks.append(keep)
|
||||
self.total += len(keep)
|
||||
keep = data[: max(0, cap - self.total)]
|
||||
if keep:
|
||||
self.chunks.append(keep)
|
||||
self.total += len(keep)
|
||||
|
||||
def drain(self) -> str:
|
||||
chunks, self.chunks, self.total = self.chunks, [], 0
|
||||
@@ -414,10 +413,7 @@ def _stdout_reader(kernel: SessionKernel) -> None:
|
||||
def _stderr_reader(kernel: SessionKernel) -> None:
|
||||
from tools.code_execution_tool import MAX_STDERR_BYTES
|
||||
assert kernel.proc is not None and kernel.proc.stderr is not None
|
||||
while True:
|
||||
chunk = kernel.proc.stderr.read1(4096)
|
||||
if not chunk:
|
||||
return
|
||||
while chunk := kernel.proc.stderr.read1(4096):
|
||||
kernel.stderr.append(chunk, MAX_STDERR_BYTES)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user