refactor(tools): simplify RPC token branch and kernel stderr/buffer loops

This commit is contained in:
Teknium
2026-09-02 23:08:22 -07:00
parent 9a9b719533
commit cf8d53ba6b
2 changed files with 12 additions and 21 deletions

View File

@@ -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

View File

@@ -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)