Limit nullish coercion to query tools and keep literal tags

A literal "null"/"none" is only a mistake where the argument is a filter, so
gate the coercion on read-only query tools; a tool that writes keeps the value,
which stops update_note(content="none") from being read as "leave unchanged".

Stop dropping nullish entries from a notes tag filter too: tags are free-form,
so a literal "none" tag stays filterable and mixed tag queries keep every
branch.
This commit is contained in:
Alex Schapiro 2026-08-25 17:10:52 +00:00
parent dd7ee42336
commit 53a3e56e10
4 changed files with 48 additions and 16 deletions

View file

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

View file

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

View file

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

View file

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