diff --git a/strix/agents/factory.py b/strix/agents/factory.py index 25212fb8..b6b15f9c 100644 --- a/strix/agents/factory.py +++ b/strix/agents/factory.py @@ -217,7 +217,12 @@ def _coerce_argument(value: Any, spec: dict[str, Any], *, nullable: bool = False return value -def _coerce_arguments(raw_input: str, schema: dict[str, Any]) -> str: +# Only query tools get nullish coercion: there a literal "null" is a filter that +# matches nothing, while a tool that writes may well be given it as real content. +_QUERY_TOOL_PREFIXES = ("list_", "search_", "view_", "get_") + + +def _coerce_arguments(raw_input: str, schema: dict[str, Any], *, nullish: bool = False) -> str: properties = schema.get("properties") if not isinstance(properties, dict) or not properties: return raw_input @@ -233,7 +238,9 @@ def _coerce_arguments(raw_input: str, schema: dict[str, Any]) -> str: spec = properties.get(key) if not isinstance(spec, dict): continue - coerced = _coerce_argument(value, spec, nullable=_is_nullable(key, spec, schema)) + coerced = _coerce_argument( + value, spec, nullable=nullish and _is_nullable(key, spec, schema) + ) if coerced is not value: payload[key] = coerced changed = True @@ -248,9 +255,10 @@ def _with_coerced_arguments(tool: FunctionTool) -> FunctionTool: return tool invoke_tool = tool.on_invoke_tool schema = tool.params_json_schema + nullish = tool.name.startswith(_QUERY_TOOL_PREFIXES) async def invoke(ctx: Any, raw_input: str) -> Any: - return await invoke_tool(ctx, _coerce_arguments(raw_input, schema)) + return await invoke_tool(ctx, _coerce_arguments(raw_input, schema, nullish=nullish)) tool.on_invoke_tool = invoke tool._strix_coerced = True # type: ignore[attr-defined] diff --git a/strix/tools/notes/tools.py b/strix/tools/notes/tools.py index 0b22e35c..065f8526 100644 --- a/strix/tools/notes/tools.py +++ b/strix/tools/notes/tools.py @@ -14,7 +14,7 @@ from typing import Any from agents import RunContextWrapper, function_tool -from strix.tools.nullish import clean_optional, is_nullish +from strix.tools.nullish import clean_optional logger = logging.getLogger(__name__) @@ -115,7 +115,6 @@ def _filter_notes( ) -> list[dict[str, Any]]: category = clean_optional(category) search_query = clean_optional(search_query) - tags = [tag for tag in tags if not is_nullish(tag)] if tags else None filtered: list[dict[str, Any]] = [] for note_id, note in _notes_storage.items(): diff --git a/tests/test_agent_factory_tool_arguments.py b/tests/test_agent_factory_tool_arguments.py index adbd605c..d7503944 100644 --- a/tests/test_agent_factory_tool_arguments.py +++ b/tests/test_agent_factory_tool_arguments.py @@ -13,22 +13,26 @@ from strix.tools.notes.tools import list_notes from strix.tools.reporting.tool import list_reports -def _capturing_tool(captured: dict[str, str], schema: dict[str, Any]) -> FunctionTool: +def _capturing_tool( + captured: dict[str, str], schema: dict[str, Any], name: str = "probe" +) -> FunctionTool: async def invoke(_ctx: Any, raw_input: str) -> str: captured["raw_input"] = raw_input return "ok" return FunctionTool( - name="probe", + name=name, description="test tool", params_json_schema={"type": "object", "properties": schema}, on_invoke_tool=invoke, ) -async def _roundtrip(schema: dict[str, Any], payload: dict[str, Any]) -> dict[str, Any]: +async def _roundtrip( + schema: dict[str, Any], payload: dict[str, Any], name: str = "probe" +) -> dict[str, Any]: captured: dict[str, str] = {} - wrapped = factory._with_coerced_arguments(_capturing_tool(captured, schema)) + wrapped = factory._with_coerced_arguments(_capturing_tool(captured, schema, name)) assert await wrapped.on_invoke_tool(cast("Any", None), json.dumps(payload)) == "ok" return cast("dict[str, Any]", json.loads(captured["raw_input"])) @@ -149,6 +153,7 @@ async def test_coercion_is_applied_once_per_tool() -> None: _NULLABLE_STRING = {"category": {"anyOf": [{"type": "string"}, {"type": "null"}]}} +_NULLABLE_CONTENT = {"content": {"anyOf": [{"type": "string"}, {"type": "null"}]}} _NULLABLE_STRING_TYPE_LIST = {"category": {"type": ["string", "null"]}} @@ -158,7 +163,7 @@ _NULLABLE_STRING_TYPE_LIST = {"category": {"type": ["string", "null"]}} async def test_nullish_string_on_a_nullable_parameter_becomes_none( schema: dict[str, Any], value: str ) -> None: - parsed = await _roundtrip(schema, {"category": value}) + parsed = await _roundtrip(schema, {"category": value}, "list_probes") assert parsed["category"] is None @@ -167,7 +172,7 @@ async def test_nullish_string_on_a_nullable_parameter_becomes_none( async def test_nullish_string_on_a_required_parameter_is_untouched() -> None: schema = {"content": {"type": "string"}} captured: dict[str, str] = {} - tool = _capturing_tool(captured, schema) + tool = _capturing_tool(captured, schema, "list_probes") tool.params_json_schema["required"] = ["content"] wrapped = factory._with_coerced_arguments(tool) @@ -178,7 +183,7 @@ async def test_nullish_string_on_a_required_parameter_is_untouched() -> None: @pytest.mark.asyncio async def test_a_parameter_absent_from_required_is_treated_as_nullable() -> None: captured: dict[str, str] = {} - tool = _capturing_tool(captured, {"category": {"type": "string"}}) + tool = _capturing_tool(captured, {"category": {"type": "string"}}, "list_probes") tool.params_json_schema["required"] = [] wrapped = factory._with_coerced_arguments(tool) @@ -188,14 +193,25 @@ async def test_a_parameter_absent_from_required_is_treated_as_nullable() -> None @pytest.mark.asyncio async def test_nullish_string_without_a_required_list_is_untouched() -> None: - parsed = await _roundtrip(_STRING, {"todos": "none"}) + parsed = await _roundtrip(_STRING, {"todos": "none"}, "list_probes") assert parsed["todos"] == "none" +@pytest.mark.asyncio +@pytest.mark.parametrize("name", ["update_note", "create_note", "record_coverage"]) +@pytest.mark.parametrize("value", ["null", "none"]) +async def test_a_nullish_value_survives_on_a_tool_that_writes(name: str, value: str) -> None: + parsed = await _roundtrip(_NULLABLE_CONTENT, {"content": value}, name) + + assert parsed["content"] == value + + @pytest.mark.asyncio async def test_nullish_looking_content_is_not_coerced() -> None: - parsed = await _roundtrip(_NULLABLE_STRING, {"category": "none of the endpoints reflect input"}) + parsed = await _roundtrip( + _NULLABLE_STRING, {"category": "none of the endpoints reflect input"}, "list_probes" + ) assert parsed["category"] == "none of the endpoints reflect input" @@ -209,7 +225,7 @@ async def test_empty_string_on_a_nullable_string_parameter_is_untouched() -> Non @pytest.mark.asyncio async def test_nullish_string_on_a_nullable_array_parameter_becomes_none() -> None: - parsed = await _roundtrip(_NULLABLE_ARRAY, {"tags": "null"}) + parsed = await _roundtrip(_NULLABLE_ARRAY, {"tags": "null"}, "list_probes") assert parsed["tags"] is None diff --git a/tests/test_notes.py b/tests/test_notes.py index d048270f..3729dd7d 100644 --- a/tests/test_notes.py +++ b/tests/test_notes.py @@ -113,7 +113,16 @@ def test_list_notes_ignores_nullish_filter_strings(nullish: str) -> None: assert notes_tools._list_notes_impl(category=nullish) == unfiltered assert notes_tools._list_notes_impl(search=nullish) == unfiltered - assert notes_tools._list_notes_impl(tags=[nullish]) == unfiltered + + +@pytest.mark.parametrize("tag", ["null", "none"]) +def test_list_notes_filters_on_a_literal_nullish_tag(tag: str) -> None: + notes_tools._create_note_impl("tagged", "content", tags=[tag]) + notes_tools._create_note_impl("other", "content", tags=["auth"]) + + assert [n["title"] for n in notes_tools._list_notes_impl(tags=[tag])["notes"]] == ["tagged"] + mixed = notes_tools._list_notes_impl(tags=[tag, "auth"]) + assert sorted(n["title"] for n in mixed["notes"]) == ["other", "tagged"] def test_list_notes_still_filters_on_real_values() -> None: