feat(mcp): collapse const-only anyOf/oneOf unions to property enums

MCP servers generated from Rust/TypeScript union types commonly emit
closed value sets as const unions:

    {"anyOf": [{"const": "red"}, {"const": "green"}, {"const": "blue"}]}

Strict tool-calling backends reject or mishandle these; the equivalent
property-level enum form is universally supported. Add
collapse_const_unions() to tools/schema_sanitizer.py and wire it into
the _normalize_mcp_input_schema discovery pipeline after the nullable
strip.

Rules:
- Collapse only when EVERY non-null branch is a pure const of the same
  primitive type (bool never merges with integer).
- Mixed unions, non-uniform const types, and mismatched declared types
  pass through untouched.
- A single {"type": "null"} branch is tolerated: consts -> enum,
  null -> nullable: true hint (matches strip_nullable_unions, which
  leaves null+multi-const unions alone by its one-non-null-branch rule).
- Outer title/description/default/examples carried onto the replacement.
- Deterministic, branch-order-preserving, non-mutating — applied at
  discovery only, so schemas stay byte-stable per conversation.

Ported from: block/goose tool_schema_normalize.rs (Apache-2.0)
This commit is contained in:
Teknium
2026-08-07 07:37:03 -07:00
parent 9fad45fcda
commit 37cc999926
3 changed files with 281 additions and 0 deletions
+171
View File
@@ -346,3 +346,174 @@ def test_dependent_schemas_still_recursively_sanitized():
assert dep_schemas["owner"] == {"type": "object", "properties": {}}, (
f"dependentSchemas['owner'] was not fully sanitized: {dep_schemas['owner']!r}"
)
# ---------------------------------------------------------------------------
# collapse_const_unions — anyOf/oneOf of same-typed const branches -> enum
# Ported from: block/goose tool_schema_normalize.rs (Apache-2.0)
# ---------------------------------------------------------------------------
from tools.schema_sanitizer import collapse_const_unions
def test_pure_const_union_collapses_to_enum():
schema = {
"anyOf": [
{"const": "red"},
{"const": "green"},
{"const": "blue"},
]
}
out = collapse_const_unions(schema)
assert out == {"type": "string", "enum": ["red", "green", "blue"]}
def test_oneof_const_union_collapses_to_enum():
schema = {"oneOf": [{"const": 1}, {"const": 2}, {"const": 3}]}
out = collapse_const_unions(schema)
assert out == {"type": "integer", "enum": [1, 2, 3]}
def test_mixed_union_left_alone():
schema = {
"anyOf": [
{"const": "a"},
{"type": "string", "minLength": 3},
]
}
out = collapse_const_unions(copy.deepcopy(schema))
assert out == schema
def test_non_uniform_const_types_left_alone():
schema = {"anyOf": [{"const": "a"}, {"const": 1}]}
out = collapse_const_unions(copy.deepcopy(schema))
assert out == schema
def test_bool_consts_not_confused_with_integers():
# bool is a subclass of int in Python; True/1 must not merge types.
schema = {"anyOf": [{"const": True}, {"const": 1}]}
out = collapse_const_unions(copy.deepcopy(schema))
assert out == schema
collapsed = collapse_const_unions({"anyOf": [{"const": True}, {"const": False}]})
assert collapsed == {"type": "boolean", "enum": [True, False]}
def test_nested_const_unions_collapse():
schema = {
"type": "object",
"properties": {
"mode": {"anyOf": [{"const": "fast"}, {"const": "slow"}]},
"inner": {
"type": "object",
"properties": {
"level": {"oneOf": [{"const": 1}, {"const": 2}]},
},
},
},
}
out = collapse_const_unions(schema)
assert out["properties"]["mode"] == {"type": "string", "enum": ["fast", "slow"]}
assert out["properties"]["inner"]["properties"]["level"] == {
"type": "integer",
"enum": [1, 2],
}
def test_outer_metadata_carried_onto_collapsed_enum():
schema = {
"title": "Color",
"description": "Pick a color",
"default": "red",
"anyOf": [{"const": "red"}, {"const": "blue"}],
}
out = collapse_const_unions(schema)
assert out == {
"type": "string",
"enum": ["red", "blue"],
"title": "Color",
"description": "Pick a color",
"default": "red",
}
def test_branch_metadata_does_not_block_collapse():
schema = {
"anyOf": [
{"const": "a", "title": "A", "description": "first"},
{"const": "b", "type": "string"},
]
}
out = collapse_const_unions(schema)
assert out == {"type": "string", "enum": ["a", "b"]}
def test_branch_with_mismatched_declared_type_left_alone():
schema = {"anyOf": [{"const": "a", "type": "integer"}, {"const": "b"}]}
out = collapse_const_unions(copy.deepcopy(schema))
assert out == schema
def test_null_plus_const_union_ordering_with_nullable_strip():
"""MCP pipeline: nullable strip runs first, then const collapse.
``anyOf: [{const a}, {const b}, {type: null}]`` has TWO non-null branches
so strip_nullable_unions leaves it; collapse_const_unions must then handle
the remaining null branch by collapsing consts and keeping nullability as
a hint.
"""
from tools.mcp_tool import _normalize_mcp_input_schema
schema = {
"type": "object",
"properties": {
"mode": {
"anyOf": [
{"const": "fast"},
{"const": "slow"},
{"type": "null"},
],
"default": None,
}
},
}
out = _normalize_mcp_input_schema(schema)
mode = out["properties"]["mode"]
assert mode["type"] == "string"
assert mode["enum"] == ["fast", "slow"]
assert mode.get("nullable") is True
def test_normalize_mcp_input_schema_collapses_const_unions():
from tools.mcp_tool import _normalize_mcp_input_schema
schema = {
"type": "object",
"properties": {
"color": {
"description": "Pick one",
"anyOf": [{"const": "red"}, {"const": "green"}],
}
},
}
out = _normalize_mcp_input_schema(schema)
assert out["properties"]["color"] == {
"description": "Pick one",
"type": "string",
"enum": ["red", "green"],
}
def test_collapse_const_unions_does_not_mutate_input():
schema = {"anyOf": [{"const": "x"}, {"const": "y"}]}
snapshot = copy.deepcopy(schema)
collapse_const_unions(schema)
assert schema == snapshot
def test_collapse_is_deterministic():
schema = {"anyOf": [{"const": "b"}, {"const": "a"}]}
first = collapse_const_unions(copy.deepcopy(schema))
second = collapse_const_unions(copy.deepcopy(schema))
assert first == second == {"type": "string", "enum": ["b", "a"]}
+14
View File
@@ -5533,6 +5533,19 @@ def _normalize_mcp_input_schema(schema: dict | None) -> dict:
return strip_nullable_unions(node, keep_nullable_hint=True)
def _collapse_const_unions(node):
"""Collapse anyOf/oneOf unions of same-typed consts to property enums.
Delegates to ``tools.schema_sanitizer.collapse_const_unions``. Runs
AFTER the nullable strip: single-non-null unions are already collapsed
by then, and unions of several const branches plus a null branch are
handled here (consts -> enum, null -> ``nullable: true`` hint).
Ported from block/goose tool_schema_normalize.rs (Apache-2.0).
"""
from tools.schema_sanitizer import collapse_const_unions
return collapse_const_unions(node)
def _repair_object_shape(node):
"""Recursively repair object-shaped nodes: fill type, prune required."""
if isinstance(node, list):
@@ -5573,6 +5586,7 @@ def _normalize_mcp_input_schema(schema: dict | None) -> dict:
normalized = _rewrite_local_refs(schema)
normalized = _strip_nullable_union(normalized)
normalized = _collapse_const_unions(normalized)
normalized = _repair_object_shape(normalized)
# Ensure top-level is a well-formed object schema
+96
View File
@@ -302,6 +302,102 @@ def strip_nullable_unions(
return stripped
_CONST_PRIMITIVE_TYPES: dict[type, str] = {
bool: "boolean",
int: "integer",
float: "number",
str: "string",
}
def _const_branch_type(branch: Any) -> str | None:
"""Return the JSON-Schema primitive type of a pure ``const`` branch.
A branch qualifies when it is a dict carrying ``const`` with a primitive
value, and any declared ``type`` matches the const value's type. Branch
metadata (``title``, ``description``) does not disqualify it, but any
other constraining keyword does. Returns ``None`` for non-qualifying
branches.
"""
if not isinstance(branch, dict) or "const" not in branch:
return None
extra = set(branch) - {"const", "type", "title", "description"}
if extra:
return None
value = branch["const"]
# bool is a subclass of int in Python; check it first so True/False never
# classify as integers.
for py_type, json_type in _CONST_PRIMITIVE_TYPES.items():
if type(value) is py_type:
declared = branch.get("type")
if declared is not None and declared != json_type:
return None
return json_type
return None
def collapse_const_unions(schema: Any) -> Any:
"""Collapse ``anyOf`` / ``oneOf`` unions of same-typed consts to ``enum``.
Ported from block/goose ``tool_schema_normalize.rs`` (Apache-2.0).
MCP servers (particularly ones generated from Rust/TypeScript union types)
commonly emit closed value sets as const unions::
{"anyOf": [{"const": "red"}, {"const": "green"}, {"const": "blue"}]}
Strict tool-calling backends reject or mishandle these, while the
equivalent property-level ``enum`` form is universally supported::
{"type": "string", "enum": ["red", "green", "blue"]}
The collapse applies only when EVERY non-null branch is a pure ``const``
of the same primitive type (bool/int/float/str — ``bool`` never merges
with ``integer``). Mixed unions and non-uniform const types pass through
untouched. A single ``{"type": "null"}`` branch is tolerated: it is
dropped and recorded as ``nullable: true`` (matching the
``strip_nullable_unions`` convention), since strip_nullable_unions only
collapses unions with exactly one non-null branch and therefore leaves
null+multi-const unions for us.
Outer-node metadata (``title``, ``description``, ``default``,
``examples``) is carried onto the replacement. Enum order preserves
branch order, so output is deterministic and byte-stable across
discoveries. Input is never mutated.
"""
if isinstance(schema, list):
return [collapse_const_unions(item) for item in schema]
if not isinstance(schema, dict):
return schema
out = {k: collapse_const_unions(v) for k, v in schema.items()}
for key in ("anyOf", "oneOf"):
variants = out.get(key)
if not isinstance(variants, list) or not variants:
continue
null_branches = [
item for item in variants
if isinstance(item, dict) and item.get("type") == "null" and "const" not in item
]
const_branches = [item for item in variants if item not in null_branches]
if len(null_branches) > 1 or not const_branches:
continue
branch_types = {_const_branch_type(item) for item in const_branches}
if len(branch_types) != 1 or None in branch_types:
continue
replacement: dict = {
"type": branch_types.pop(),
"enum": [item["const"] for item in const_branches],
}
if null_branches:
replacement["nullable"] = True
for meta_key in ("title", "description", "default", "examples"):
if meta_key in out and meta_key not in replacement:
replacement[meta_key] = out[meta_key]
return replacement
return out
def _sanitize_node(node: Any, path: str) -> Any:
"""Recursively sanitize a JSON-Schema fragment.