Skip to content

Commit 276f9f8

Browse files
committed
fix: pass command_args/env to decorator and filter, W504 line break
- _safety_wrapper: decorator now passes command_args/environment_variables/ working_directory to SafetyWrapper.check (was dead code path) - _safety_filter: _before now extracts command_args/env_vars/work_dir from req - _scanner: fix W504 line break after binary operator
1 parent 074fd91 commit 276f9f8

3 files changed

Lines changed: 66 additions & 4 deletions

File tree

trpc_agent_sdk/tools/safety/_safety_filter.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -113,10 +113,17 @@ async def _before(self, ctx: AgentContext, req: Any, rsp: FilterResult) -> None:
113113
script_type = _guess_script_type(req, script_content)
114114
tool_name = _extract_tool_name(req)
115115

116+
cmd_args = _extract_list_field(req, "command_args", "cmd_args")
117+
env_vars = _extract_dict_field(req, "environment_variables", "env_vars")
118+
work_dir = _extract_str_field(req, "working_directory", "cwd")
119+
116120
scan_input = SafetyScanInput(
117121
script_content=script_content,
118122
script_type=script_type,
119123
tool_name=tool_name,
124+
command_args=cmd_args,
125+
environment_variables=env_vars,
126+
working_directory=work_dir,
120127
)
121128

122129
report = self._scanner.scan(scan_input)
@@ -239,3 +246,40 @@ def _extract_tool_name(req: Any) -> str:
239246
if hasattr(req, "name"):
240247
return str(getattr(req, "name", "unknown"))
241248
return "unknown"
249+
250+
251+
def _extract_list_field(req: Any, *keys: str) -> Optional[list[str]]:
252+
"""Extract a list-typed field from the request by trying multiple key names."""
253+
if isinstance(req, dict):
254+
for k in keys:
255+
val = req.get(k) or req.get("args", {}).get(k)
256+
if isinstance(val, list):
257+
return val
258+
if hasattr(req, "args"):
259+
args = getattr(req, "args")
260+
if isinstance(args, dict):
261+
for k in keys:
262+
val = args.get(k)
263+
if isinstance(val, list):
264+
return val
265+
return None
266+
267+
268+
def _extract_str_field(req: Any, *keys: str) -> Optional[str]:
269+
"""Extract a string-typed field from the request by trying multiple key names."""
270+
if isinstance(req, dict):
271+
for k in keys:
272+
val = req.get(k) or req.get("args", {}).get(k)
273+
if isinstance(val, str):
274+
return val
275+
return None
276+
277+
278+
def _extract_dict_field(req: Any, *keys: str) -> Optional[dict[str, str]]:
279+
"""Extract a dict-typed field from the request by trying multiple key names."""
280+
if isinstance(req, dict):
281+
for k in keys:
282+
val = req.get(k) or req.get("args", {}).get(k)
283+
if isinstance(val, dict):
284+
return val
285+
return None

trpc_agent_sdk/tools/safety/_safety_wrapper.py

Lines changed: 20 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -223,6 +223,16 @@ def decorator(func: Callable) -> Callable:
223223
raise_on_deny=raise_on_deny,
224224
)
225225

226+
def _extract_extra_fields(call_args: tuple, call_kwargs: dict) -> tuple:
227+
"""Extract extra scan fields from args/kwargs."""
228+
for arg in call_args:
229+
if isinstance(arg, dict):
230+
return (arg.get("command_args") or arg.get("cmd_args"), arg.get("environment_variables")
231+
or arg.get("env_vars"), arg.get("working_directory") or arg.get("cwd"))
232+
return (call_kwargs.get("command_args")
233+
or call_kwargs.get("cmd_args"), call_kwargs.get("environment_variables")
234+
or call_kwargs.get("env_vars"), call_kwargs.get("working_directory") or call_kwargs.get("cwd"))
235+
226236
@functools.wraps(func)
227237
async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
228238
script: Optional[str] = kwargs.get(script_arg_name)
@@ -232,7 +242,11 @@ async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
232242
script = arg[script_arg_name]
233243
break
234244
if script and isinstance(script, str):
235-
wrapper_inst.check(script)
245+
cmd_args, env_vars, work_dir = _extract_extra_fields(args, kwargs)
246+
wrapper_inst.check(script,
247+
command_args=cmd_args,
248+
environment_variables=env_vars,
249+
working_directory=work_dir)
236250
elif require_script:
237251
raise RuntimeError(f"safety_wrapper: '{script_arg_name}' not found in kwargs or positional "
238252
f"dict args for {func.__name__}. Check the decorator configuration "
@@ -255,7 +269,11 @@ def sync_wrapper(*args: Any, **kwargs: Any) -> Any:
255269
script = arg[script_arg_name]
256270
break
257271
if script and isinstance(script, str):
258-
wrapper_inst.check(script)
272+
cmd_args, env_vars, work_dir = _extract_extra_fields(args, kwargs)
273+
wrapper_inst.check(script,
274+
command_args=cmd_args,
275+
environment_variables=env_vars,
276+
working_directory=work_dir)
259277
elif require_script:
260278
raise RuntimeError(f"safety_wrapper: '{script_arg_name}' not found in kwargs or positional "
261279
f"dict args for {func.__name__}. Check the decorator configuration "

trpc_agent_sdk/tools/safety/_scanner.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -281,9 +281,9 @@ def scan(self, scan_input: SafetyScanInput) -> SafetyScanReport:
281281
else:
282282
denied = sum(1 for f in all_findings if f.risk_level in (RiskLevel.CRITICAL, RiskLevel.HIGH))
283283
total = len(all_findings)
284+
auto_suffix = " [auto_allowed by allow_patterns]" if allow_upgraded else ""
284285
summary = (f"Scan of '{scan_input.tool_name or 'unnamed tool'}' found {total} issue(s) "
285-
f"({denied} high/critical). Decision: {decision.value}." +
286-
(" [auto_allowed by allow_patterns]" if allow_upgraded else ""))
286+
f"({denied} high/critical). Decision: {decision.value}.{auto_suffix}")
287287

288288
return SafetyScanReport(
289289
tool_name=scan_input.tool_name,

0 commit comments

Comments
 (0)