Compare commits

...
7 changed files with 451 additions and 51 deletions
+66 -13
View File
@@ -55,7 +55,7 @@ from strix.tools.web_search.tool import web_search
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from collections.abc import Awaitable, Callable, Sequence
from agents import RunContextWrapper
from agents.tool import FunctionToolResult
@@ -349,6 +349,48 @@ _BASE_TOOLS: tuple[Tool, ...] = (
)
# Extra tools registered for scan agents. Mirrors
# ``strix.runtime.backends.register_backend``: register before the first
# ``build_strix_agent`` call and every agent (root + children) gets them.
_EXTRA_TOOLS: list[Tool] = []
def _ensure_unique_tool_names(tools: Sequence[Tool]) -> None:
seen: set[str] = set()
duplicates: set[str] = set()
for tool in tools:
if tool.name in seen:
duplicates.add(tool.name)
seen.add(tool.name)
if duplicates:
msg = f"Agent tools must have unique names: {sorted(duplicates)}"
raise ValueError(msg)
def register_agent_tools(*tools: Tool) -> None:
"""Register tools for every scan agent built afterwards.
Tools are added to both root and child agents, after the base set and
before the lifecycle tool (``finish_scan`` / ``agent_finish``). Duplicate
tool objects are ignored so repeated imports don't double-register.
"""
new_tools: list[Tool] = []
for tool in tools:
if tool not in _EXTRA_TOOLS and tool not in new_tools:
new_tools.append(tool)
_ensure_unique_tool_names([*_BASE_TOOLS, *_EXTRA_TOOLS, *new_tools, finish_scan, agent_finish])
for tool in new_tools:
_EXTRA_TOOLS.append(tool)
logger.info("Registered extra agent tool: %s", getattr(tool, "name", tool))
def registered_agent_tools() -> tuple[Tool, ...]:
"""Return the currently registered scan-agent tools."""
return tuple(_EXTRA_TOOLS)
def build_strix_agent(
*,
name: str = "strix",
@@ -359,26 +401,37 @@ def build_strix_agent(
interactive: bool = False,
chat_completions_tools: bool = False,
system_prompt_context: dict[str, Any] | None = None,
extra_tools: Sequence[Tool] | None = None,
instructions_override: str | None = None,
) -> SandboxAgent[Any]:
"""Build a SandboxAgent for either root or child use.
Args:
chat_completions_tools: Wrap SDK custom tools as function tools
when the selected backend cannot accept Responses custom tools.
extra_tools: Additional tools for this scan agent only, on top of any
registered via ``register_agent_tools``.
instructions_override: Use this verbatim as the system prompt instead
of rendering the built-in scan prompt.
"""
instructions = render_system_prompt(
skills=skills,
scan_mode=scan_mode,
is_whitebox=is_whitebox,
is_root=is_root,
interactive=interactive,
system_prompt_context=system_prompt_context,
)
if is_root:
tools: list[Tool] = [*_BASE_TOOLS, finish_scan]
if instructions_override is not None:
instructions = instructions_override
else:
tools = [*_BASE_TOOLS, agent_finish]
instructions = render_system_prompt(
skills=skills,
scan_mode=scan_mode,
is_whitebox=is_whitebox,
is_root=is_root,
interactive=interactive,
system_prompt_context=system_prompt_context,
)
agent_tools = [*_EXTRA_TOOLS, *(extra_tools or [])]
if is_root:
tools: list[Tool] = [*_BASE_TOOLS, *agent_tools, finish_scan]
else:
tools = [*_BASE_TOOLS, *agent_tools, agent_finish]
_ensure_unique_tool_names(tools)
logger.info(
"Built %s agent '%s' (skills=%d, tools=%d, scan_mode=%s, whitebox=%s)",
+3 -3
View File
@@ -7,7 +7,7 @@ from typing import Any
from jinja2 import Environment, FileSystemLoader, select_autoescape
from strix.skills import get_available_skills, load_skills
from strix.skills import get_available_skills, load_skills, skill_search_dirs
from strix.utils.resource_paths import get_strix_resource_path
@@ -69,9 +69,9 @@ def render_system_prompt(
"""Render the system prompt. Returns empty string on template failure."""
try:
prompt_dir = get_strix_resource_path("agents", _PROMPT_DIRNAME)
skills_dir = get_strix_resource_path("skills")
loader_dirs = [prompt_dir, *skill_search_dirs()]
env = Environment(
loader=FileSystemLoader([prompt_dir, skills_dir]),
loader=FileSystemLoader(loader_dirs),
autoescape=select_autoescape(
enabled_extensions=(),
default_for_string=False,
+4
View File
@@ -24,6 +24,10 @@ DEFAULT_MAX_TURNS = 500
def _accepts_required_tool_choice(model_name: str | None) -> bool:
name = (model_name or "").strip().lower()
for prefix in ("litellm/", "any-llm/"):
if name.startswith(prefix):
name = name[len(prefix) :]
break
return name.startswith("openai/") or is_known_openai_bare_model(name)
+150 -35
View File
@@ -1,6 +1,8 @@
import logging
import re
from collections import Counter
from collections.abc import Iterator
from pathlib import Path
from strix.utils.resource_paths import get_strix_resource_path
@@ -10,20 +12,82 @@ logger = logging.getLogger(__name__)
_FRONTMATTER_PATTERN = re.compile(r"^---\s*\n.*?\n---\s*\n", re.DOTALL)
_INTERNAL_SKILL_CATEGORIES: frozenset[str] = frozenset({"scan_modes", "coordination"})
_ROOT_SKILL_CATEGORY = "root"
_EXTRA_SKILL_DIRS: list[Path] = []
def register_skill_dir(path: str | Path) -> None:
"""Add a directory searched for skills ahead of the built-in set.
The directory uses the same layout as the packaged skills
(``<root>/<category>/<name>.md``). Skills found in a registered
directory shadow packaged skills with the same relative path, so
callers can both add new skills and override existing ones without
editing the package. The most recently registered directory has the
highest precedence.
"""
resolved = Path(path)
if resolved not in _EXTRA_SKILL_DIRS:
_EXTRA_SKILL_DIRS.append(resolved)
logger.info("Registered extra skill dir: %s", resolved)
def registered_skill_dirs() -> tuple[Path, ...]:
"""Return registered extra skill directories, highest precedence first."""
return tuple(reversed(_EXTRA_SKILL_DIRS))
def skill_search_dirs() -> tuple[Path, ...]:
"""All existing skill roots, highest precedence first (built-in last)."""
roots = [d for d in registered_skill_dirs() if d.is_dir()]
builtin = get_strix_resource_path("skills")
if builtin.is_dir():
roots.append(builtin)
return tuple(roots)
def _iter_user_skill_files() -> Iterator[tuple[str, str]]:
"""Yield ``(category_name, skill_name)`` for every user-selectable skill."""
skills_dir = get_strix_resource_path("skills")
if not skills_dir.exists():
return
for category_dir in sorted(skills_dir.iterdir()):
if not category_dir.is_dir() or category_dir.name.startswith("__"):
continue
if category_dir.name in _INTERNAL_SKILL_CATEGORIES:
continue
for file_path in sorted(category_dir.glob("*.md")):
yield category_dir.name, file_path.stem
seen: set[tuple[str, str]] = set()
for skills_dir in skill_search_dirs():
for file_path in sorted(skills_dir.glob("*.md")):
if file_path.name.startswith("__") or file_path.name == "README.md":
continue
key = (_ROOT_SKILL_CATEGORY, file_path.stem)
if key in seen:
continue
seen.add(key)
yield key
for category_dir in sorted(skills_dir.iterdir()):
if not category_dir.is_dir() or category_dir.name.startswith("__"):
continue
if category_dir.name in _INTERNAL_SKILL_CATEGORIES:
continue
for file_path in sorted(category_dir.glob("*.md")):
key = (category_dir.name, file_path.stem)
if key in seen:
continue
seen.add(key)
yield key
def _is_selectable_root_skill_file(file_path: Path) -> bool:
return file_path.suffix == ".md" and not (
file_path.name.startswith("__") or file_path.name == "README.md"
)
def _qualified_skill_file(skills_dir: Path, category: str, name: str) -> Path | None:
if category == _ROOT_SKILL_CATEGORY:
candidate = skills_dir / f"{name}.md"
if candidate.exists() and _is_selectable_root_skill_file(candidate):
return candidate
return None
candidate = skills_dir / category / f"{name}.md"
return candidate if candidate.exists() else None
def get_all_skill_names() -> set[str]:
@@ -31,6 +95,54 @@ def get_all_skill_names() -> set[str]:
return {name for _, name in _iter_user_skill_files()}
def _get_all_skill_keys() -> set[str]:
keys: set[str] = set()
for category, name in _iter_user_skill_files():
keys.add(f"{category}/{name}")
return keys
def _get_ambiguous_skill_names() -> set[str]:
counts = Counter(name for _, name in _iter_user_skill_files())
return {name for name, count in counts.items() if count > 1}
def _qualified_skill_files(skill_name: str) -> list[Path]:
category, _, name = skill_name.partition("/")
for skills_dir in skill_search_dirs():
candidate = _qualified_skill_file(skills_dir, category, name)
if candidate is not None:
return [candidate]
return []
def _bare_skill_files(skill_name: str) -> list[Path]:
seen: set[tuple[str, str]] = set()
candidates: list[Path] = []
for skills_dir in skill_search_dirs():
for category_dir in sorted(skills_dir.iterdir()):
if not category_dir.is_dir() or category_dir.name.startswith("__"):
continue
if category_dir.name in _INTERNAL_SKILL_CATEGORIES:
continue
key = (category_dir.name, skill_name)
if key in seen:
continue
candidate = category_dir / f"{skill_name}.md"
if candidate.exists():
seen.add(key)
candidates.append(candidate)
key = (_ROOT_SKILL_CATEGORY, skill_name)
if key in seen:
continue
candidate = _qualified_skill_file(skills_dir, _ROOT_SKILL_CATEGORY, skill_name)
if candidate is not None:
seen.add(key)
candidates.append(candidate)
return candidates
def get_available_skills() -> dict[str, list[str]]:
grouped: dict[str, list[str]] = {}
for category, name in _iter_user_skill_files():
@@ -52,48 +164,51 @@ def validate_requested_skills(skill_list: list[str], max_skills: int = 5) -> str
if not skill_list:
return None
available = get_all_skill_names()
invalid = sorted({s for s in skill_list if s not in available})
available_keys = _get_all_skill_keys()
invalid = sorted({s for s in skill_list if s not in available and s not in available_keys})
if invalid:
return f"Invalid skill name(s): {invalid}. Available skills: {sorted(available)}"
ambiguous = sorted({s for s in skill_list if "/" not in s} & _get_ambiguous_skill_names())
if ambiguous:
return (
f"Ambiguous skill name(s): {ambiguous}. Use category-qualified names from: "
f"{sorted(available_keys)}"
)
return None
def _candidate_skill_files(skill_name: str) -> list[Path]:
"""Resolve *skill_name* to effective matching files."""
if "/" in skill_name:
return _qualified_skill_files(skill_name)
return _bare_skill_files(skill_name)
def load_skills(skill_names: list[str]) -> dict[str, str]:
"""Load skill markdown bodies (frontmatter stripped) by name.
Skill files live at ``strix/skills/<category>/<name>.md``. Names
can be ``"name"`` (any category), ``"category/name"``, or a bare
file at the skills root. Missing skills are logged and skipped.
Skill files live at ``strix/skills/<category>/<name>.md`` (or any
directory added via :func:`register_skill_dir`, searched first).
Names can be ``"name"`` (any category), ``"category/name"``, or a
bare file at the skills root. Missing skills are logged and skipped.
"""
skills_dir = get_strix_resource_path("skills")
if not skills_dir.exists():
search_dirs = skill_search_dirs()
if not search_dirs:
return {}
by_category: dict[str, str] = {}
for category_dir in skills_dir.iterdir():
if not category_dir.is_dir() or category_dir.name.startswith("__"):
continue
for file_path in category_dir.glob("*.md"):
by_category[file_path.stem] = f"{category_dir.name}/{file_path.stem}.md"
skill_content: dict[str, str] = {}
for skill_name in skill_names:
rel_path: str | None
if "/" in skill_name:
rel_path = f"{skill_name}.md"
elif skill_name in by_category:
rel_path = by_category[skill_name]
elif (skills_dir / f"{skill_name}.md").exists():
rel_path = f"{skill_name}.md"
else:
rel_path = None
if rel_path is None or not (skills_dir / rel_path).exists():
candidates = _candidate_skill_files(skill_name)
if not candidates:
logger.warning("Skill not found: %s", skill_name)
continue
if len(candidates) > 1:
logger.warning("Ambiguous skill name %s; use a category-qualified name", skill_name)
continue
file_path = candidates[0]
try:
content = (skills_dir / rel_path).read_text(encoding="utf-8")
content = file_path.read_text(encoding="utf-8")
except (OSError, ValueError) as e:
logger.warning("Failed to load skill %s: %s", skill_name, e)
continue
+88
View File
@@ -0,0 +1,88 @@
"""Tests for scan-agent tool registration in factory."""
from __future__ import annotations
import pytest
from agents.tool import FunctionTool
from strix.agents import factory
def _tool(name: str) -> FunctionTool:
return FunctionTool(
name=name,
description="test tool",
params_json_schema={"type": "object", "properties": {}, "additionalProperties": False},
on_invoke_tool=lambda _ctx, _inp: "ok",
)
@pytest.fixture(autouse=True)
def _reset_registry() -> object:
saved = list(factory._EXTRA_TOOLS)
factory._EXTRA_TOOLS.clear()
try:
yield
finally:
factory._EXTRA_TOOLS[:] = saved
def test_register_agent_tools_is_deduped() -> None:
tool = _tool("dup")
factory.register_agent_tools(tool)
factory.register_agent_tools(tool)
assert factory.registered_agent_tools() == (tool,)
def test_registered_tools_appear_before_lifecycle_tool() -> None:
tool = _tool("extra")
factory.register_agent_tools(tool)
root = factory.build_strix_agent(is_root=True)
child = factory.build_strix_agent(is_root=False)
root_names = [t.name for t in root.tools]
child_names = [t.name for t in child.tools]
assert root_names[-2:] == ["extra", "finish_scan"]
assert child_names[-2:] == ["extra", "agent_finish"]
def test_per_call_extra_tools_stack_with_registry() -> None:
factory.register_agent_tools(_tool("registered"))
agent = factory.build_strix_agent(is_root=True, extra_tools=[_tool("per_call")])
names = [t.name for t in agent.tools]
assert "registered" in names
assert "per_call" in names
assert names[-1] == "finish_scan"
def test_register_agent_tools_rejects_duplicate_names() -> None:
factory.register_agent_tools(_tool("same_name"))
with pytest.raises(ValueError, match="same_name"):
factory.register_agent_tools(_tool("same_name"))
def test_per_call_extra_tools_reject_duplicate_registered_names() -> None:
factory.register_agent_tools(_tool("same_name"))
with pytest.raises(ValueError, match="same_name"):
factory.build_strix_agent(is_root=True, extra_tools=[_tool("same_name")])
def test_instructions_override_is_used_verbatim() -> None:
custom = "You are a scan agent. Follow the provided scope."
agent = factory.build_strix_agent(is_root=True, instructions_override=custom)
assert agent.instructions == custom
def test_no_override_renders_builtin_prompt() -> None:
agent = factory.build_strix_agent(is_root=True)
assert isinstance(agent.instructions, str)
assert agent.instructions != ""
+20
View File
@@ -135,3 +135,23 @@ def test_make_model_settings_skips_required_tool_choice_for_non_openai_models()
)
assert settings.tool_choice is None
def test_make_model_settings_forces_required_for_routed_openai_model() -> None:
settings = make_model_settings(
None,
model_name="litellm/openai/gpt-4o",
force_required_tool_choice=True,
)
assert settings.tool_choice == "required"
def test_make_model_settings_forces_required_for_anyllm_routed_openai_model() -> None:
settings = make_model_settings(
None,
model_name="any-llm/openai/gpt-4o",
force_required_tool_choice=True,
)
assert settings.tool_choice == "required"
+120
View File
@@ -0,0 +1,120 @@
from pathlib import Path
import pytest
import strix.skills as skills_mod
from strix.skills import (
get_all_skill_names,
get_available_skills,
load_skills,
register_skill_dir,
registered_skill_dirs,
skill_search_dirs,
validate_requested_skills,
)
@pytest.fixture(autouse=True)
def _clear_extra_dirs() -> None:
original = list(skills_mod._EXTRA_SKILL_DIRS)
skills_mod._EXTRA_SKILL_DIRS.clear()
try:
yield
finally:
skills_mod._EXTRA_SKILL_DIRS[:] = original
def _write_skill(root: Path, category: str, name: str, body: str) -> None:
category_dir = root / category
category_dir.mkdir(parents=True, exist_ok=True)
(category_dir / f"{name}.md").write_text(body, encoding="utf-8")
def _write_root_skill(root: Path, name: str, body: str) -> None:
root.mkdir(parents=True, exist_ok=True)
(root / f"{name}.md").write_text(body, encoding="utf-8")
def test_no_registration_leaves_builtin_only() -> None:
assert registered_skill_dirs() == ()
builtin = skills_mod.get_strix_resource_path("skills")
assert skill_search_dirs() == (builtin,)
assert {"nmap", "subfinder"}.issubset(get_available_skills()["tooling"])
def test_register_is_idempotent_and_ordered(tmp_path: Path) -> None:
a = tmp_path / "a"
b = tmp_path / "b"
a.mkdir()
b.mkdir()
register_skill_dir(a)
register_skill_dir(b)
register_skill_dir(a)
# Most recently registered wins → highest precedence first.
assert registered_skill_dirs() == (b, a)
def test_registered_dir_adds_new_skill(tmp_path: Path) -> None:
_write_skill(tmp_path, "extra", "widget", "widget body")
register_skill_dir(tmp_path)
assert "widget" in get_all_skill_names()
assert get_available_skills()["extra"] == ["widget"]
assert load_skills(["widget"]) == {"widget": "widget body"}
def test_registered_root_skill_is_discoverable_and_valid(tmp_path: Path) -> None:
_write_root_skill(tmp_path, "widget", "widget body")
register_skill_dir(tmp_path)
assert "widget" in get_all_skill_names()
assert get_available_skills()["root"] == ["widget"]
assert validate_requested_skills(["widget"]) is None
assert validate_requested_skills(["root/widget"]) is None
assert load_skills(["widget"]) == {"widget": "widget body"}
assert load_skills(["root/widget"]) == {"widget": "widget body"}
def test_ambiguous_bare_skill_requires_qualified_name(tmp_path: Path) -> None:
_write_skill(tmp_path, "alpha", "widget", "alpha body")
_write_skill(tmp_path, "beta", "widget", "beta body")
register_skill_dir(tmp_path)
assert "widget" in get_all_skill_names()
assert get_available_skills()["alpha"] == ["widget"]
assert get_available_skills()["beta"] == ["widget"]
assert validate_requested_skills(["alpha/widget"]) is None
assert validate_requested_skills(["beta/widget"]) is None
error = validate_requested_skills(["widget"])
assert error is not None
assert "Ambiguous skill name" in error
assert "alpha/widget" in error
assert "beta/widget" in error
assert load_skills(["widget"]) == {}
assert load_skills(["alpha/widget"]) == {"widget": "alpha body"}
assert load_skills(["beta/widget"]) == {"widget": "beta body"}
def test_registered_dir_overrides_builtin_skill(tmp_path: Path) -> None:
_write_skill(tmp_path, "coordination", "root_agent", "overridden root agent")
register_skill_dir(tmp_path)
loaded = load_skills(["coordination/root_agent"])
assert loaded["root_agent"] == "overridden root agent"
def test_builtin_skill_still_loads_when_not_overridden(tmp_path: Path) -> None:
_write_skill(tmp_path, "extra", "widget", "widget body")
register_skill_dir(tmp_path)
# A packaged skill the registered dir does not shadow still resolves.
assert load_skills(["scan_modes/deep"]).get("deep")
def test_missing_skill_is_skipped(tmp_path: Path) -> None:
register_skill_dir(tmp_path)
assert load_skills(["does_not_exist"]) == {}