diff --git a/Makefile b/Makefile index a23e82e..7b097c5 100644 --- a/Makefile +++ b/Makefile @@ -11,6 +11,11 @@ gen-all: ## Generate all code from schema @uv run ruff check --fix @uv run ruff format . +.PHONY: gen-check +gen-check: ## Verify generated schema bindings without changing the worktree + @echo "🚀 Checking generated schema bindings" + @uv run --frozen python -m scripts.gen_schema --check + .PHONY: check check: ## Run code quality tools. @echo "🚀 Checking lock file consistency with 'pyproject.toml'" diff --git a/pyproject.toml b/pyproject.toml index 2602140..5934086 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -22,6 +22,7 @@ classifiers = [ ] dependencies = [ "pydantic>=2.7", + "pydantic-core>=2.18.1", ] @@ -32,7 +33,7 @@ Documentation = "https://agentclientprotocol.github.io/python-sdk/" [dependency-groups] dev = [ - "datamodel-code-generator>=0.71.0", + "datamodel-code-generator==0.71.0", "pytest>=7.2.0", "pytest-asyncio>=0.21.0", "tox-uv>=1.11.3", diff --git a/scripts/gen_schema.py b/scripts/gen_schema.py index 638bbe7..0e01c6e 100644 --- a/scripts/gen_schema.py +++ b/scripts/gen_schema.py @@ -1,1115 +1,437 @@ #!/usr/bin/env python3 from __future__ import annotations -import ast +import argparse import copy +import difflib import json -import re import subprocess import sys -import tempfile import textwrap -from collections.abc import Callable -from dataclasses import dataclass from pathlib import Path from typing import Any +from datamodel_code_generator import ( + DataModelType, + Formatter, + InputFileType, + LiteralType, + PythonVersion, + generate, +) +from datamodel_code_generator.enums import NamingStrategy, VersionMode +from datamodel_code_generator.validators import ModelValidators, ValidatorDefinition, ValidatorMode +from pydantic.alias_generators import to_snake + ROOT = Path(__file__).resolve().parents[1] -SCHEMA_DIR = ROOT / "schema" -SCHEMA_JSON = SCHEMA_DIR / "schema.json" -VERSION_FILE = SCHEMA_DIR / "VERSION" +SCHEMA_JSON = ROOT / "schema" / "schema.json" +VERSION_FILE = ROOT / "schema" / "VERSION" SCHEMA_OUT = ROOT / "src" / "acp" / "schema.py" -STDIO_TYPE_LITERAL = 'Literal["2#-datamodel-code-generator-#-object-#-special-#"]' -MODELS_TO_REMOVE = [ - "AgentClientProtocol", - "AgentClientProtocol1", - "AgentClientProtocol2", - "AgentClientProtocol3", - "AgentClientProtocol4", - "AgentClientProtocol5", - "AgentClientProtocol6", - "AgentClientProtocol7", -] - -# Map of numbered classes produced by datamodel-code-generator to descriptive names. -# Keep this in sync with the Rust/TypeScript SDK nomenclature. -RENAME_MAP: dict[str, str] = { - "AgentResponse1": "AgentResponseMessage", - "AgentResponse2": "AgentErrorMessage", - "ClientResponse1": "ClientResponseMessage", - "ClientResponse2": "ClientErrorMessage", - "ContentBlock1": "TextContentBlock", - "ContentBlock2": "ImageContentBlock", - "ContentBlock3": "AudioContentBlock", - "ContentBlock4": "ResourceContentBlock", - "ContentBlock5": "EmbeddedResourceContentBlock", - "McpServer1": "HttpMcpServer", - "McpServer2": "SseMcpServer", - "McpServer3": "AcpMcpServer", - "RequestPermissionOutcome1": "DeniedOutcome", - "RequestPermissionOutcome2": "AllowedOutcome", - "AuthMethod1": "EnvVarAuthMethod", - "AuthMethod2": "TerminalAuthMethod", - "SessionConfigOption1": "SessionConfigOptionSelect", - "SessionConfigOption2": "SessionConfigOptionBoolean", - "SetSessionConfigOptionRequest1": "SetSessionConfigOptionBooleanRequest", - "SetSessionConfigOptionRequest2": "SetSessionConfigOptionSelectRequest", - "SessionUpdate1": "UserMessageChunk", - "SessionUpdate2": "AgentMessageChunk", - "SessionUpdate3": "AgentThoughtChunk", - "SessionUpdate4": "ToolCallStart", - "SessionUpdate5": "ToolCallProgress", - "SessionUpdate6": "AgentPlanUpdate", - "SessionUpdate7": "AgentPlanContentUpdate", - "SessionUpdate8": "AgentPlanRemovedUpdate", - "SessionUpdate9": "AvailableCommandsUpdate", - "SessionUpdate10": "CurrentModeUpdate", - "SessionUpdate11": "ConfigOptionUpdate", - "SessionUpdate12": "SessionInfoUpdate", - "SessionUpdate13": "UsageUpdate", - "PlanUpdateContent1": "PlanUpdateItems", - "PlanUpdateContent2": "PlanUpdateFile", - "PlanUpdateContent3": "PlanUpdateMarkdown", - "ToolCallContent1": "ContentToolCallContent", - "ToolCallContent2": "FileEditToolCallContent", - "ToolCallContent3": "TerminalToolCallContent", - "CreateElicitationRequest1": "CreateFormSessionElicitationRequest", - "CreateElicitationRequest2": "CreateFormRequestElicitationRequest", - "CreateElicitationRequest3": "CreateUrlSessionElicitationRequest", - "CreateElicitationRequest4": "CreateUrlRequestElicitationRequest", - "CreateElicitationRequest5": "CreateOtherElicitationRequest", - "CreateElicitationResponse1": "AcceptElicitationResponse", - "CreateElicitationResponse2": "DeclineElicitationResponse", - "CreateElicitationResponse3": "CancelElicitationResponse", - "CreateElicitationResponse4": "OtherElicitationResponse", - "ElicitationFormMode1": "ElicitationFormSessionMode", - "ElicitationFormMode2": "ElicitationFormRequestMode", - "ElicitationPropertySchema1": "ElicitationStringPropertySchema", - "ElicitationPropertySchema2": "ElicitationNumberPropertySchema", - "ElicitationPropertySchema3": "ElicitationIntegerPropertySchema", - "ElicitationPropertySchema4": "ElicitationBooleanPropertySchema", - "ElicitationPropertySchema5": "ElicitationMultiSelectPropertySchema", - "ElicitationPropertySchema6": "ElicitationOtherPropertySchema", - "MultiSelectItems1": "StringMultiSelectItems", - "MultiSelectItems2": "OtherMultiSelectItems", - "ElicitationUrlMode1": "ElicitationUrlSessionMode", - "ElicitationUrlMode2": "ElicitationUrlRequestMode", - "NesSuggestion1": "NesEditSuggestionVariant", - "NesSuggestion2": "NesJumpSuggestionVariant", - "NesSuggestion3": "NesRenameSuggestionVariant", - "NesSuggestion4": "NesSearchAndReplaceSuggestionVariant", -} - -# Extensible ("custom or future") unions: known const-tagged variants plus a -# catch-all member tagged `"title": "other"`. _normalize_catchall_unions strips the -# discriminator and the catch-all's `not` clause so datamodel-codegen produces a plain -# union; the exclusion is restored at runtime by a field_validator injected into the -# catch-all class, so a malformed known variant fails instead of silently parsing as -# custom (mirrors the TypeScript SDK's excludeKnownTags). Maps union def name -> -# catch-all class name; the set is asserted against the schema in -# _validate_schema_alignment. -EXTENSIBLE_UNIONS: dict[str, str] = { - "CreateElicitationRequest": "CreateOtherElicitationRequest", - "CreateElicitationResponse": "OtherElicitationResponse", - "ElicitationPropertySchema": "ElicitationOtherPropertySchema", - "MultiSelectItems": "OtherMultiSelectItems", -} - -ENUM_LITERAL_MAP: dict[str, tuple[str, ...]] = { - "PermissionOptionKind": ( - "allow_once", - "allow_always", - "reject_once", - "reject_always", - ), - "PlanEntryPriority": ("high", "medium", "low"), - "PlanEntryStatus": ("pending", "in_progress", "completed"), - "StopReason": ("end_turn", "max_tokens", "max_turn_requests", "refusal", "cancelled"), - "ToolCallStatus": ("pending", "in_progress", "completed", "failed"), - "ToolKind": ("read", "edit", "delete", "move", "search", "execute", "think", "fetch", "switch_mode", "other"), -} - -# datamodel-code-generator 0.64 promotes referenced string enums to Enum classes. -# Keep the existing Python API, where these schema types are plain strings; the -# selected public fields below are narrowed back to the named Literal aliases. -STRING_ENUM_TYPES = ( - *ENUM_LITERAL_MAP, - "ElicitationSchemaType", - "NesDiagnosticSeverity", - "NesRejectReason", - "NesTriggerKind", - "PositionEncodingKind", - "Role", - "StringFormat", - "TextDocumentSyncKind", +UNSIGNED_TYPE_MAPPINGS = ( + "integer+uint16=integer", + "integer+uint32=integer", + "integer+uint64=integer", ) -# Preserve RootModel classes that existed in the generated public surface before -# 0.64; other unreferenced RootModels are intermediates left after collapsing. -PUBLIC_ROOT_MODELS = { - "AgentResponse", - "ClientResponse", - "ElicitationContentValue", - "ElicitationFormMode", - "ElicitationUrlMode", +OPEN_UNIONS = { + "CreateElicitationRequest": 2, + "CreateElicitationResponse": 3, + "ElicitationPropertySchema": 5, + "MultiSelectItems": 1, } -FIELD_TYPE_OVERRIDES: tuple[tuple[str, str, str, bool], ...] = ( - ("PermissionOption", "kind", "PermissionOptionKind", False), - ("PlanEntry", "priority", "PlanEntryPriority", False), - ("PlanEntry", "status", "PlanEntryStatus", False), - ("PromptResponse", "stop_reason", "StopReason", False), - ("ToolCall", "kind", "ToolKind", True), - ("ToolCall", "status", "ToolCallStatus", True), - ("ToolCallUpdate", "kind", "ToolKind", True), - ("ToolCallUpdate", "status", "ToolCallStatus", True), -) +def _inline_model_ref(definition: str, *steps: tuple[str, int | None]) -> str: + ref = f"#/$defs/{definition}" + for keyword, index in steps: + ref += f"#-datamodel-code-generator-#-{keyword}-#-special-#" + if index is not None: + ref += f"/{index}" + return ref -@dataclass(frozen=True) -class FieldValidatorInjection: - """A generated field validator that should be appended to one schema class.""" - class_name: str - field_name: str - method_name: str - argument_name: str - return_type: str - comment_lines: tuple[str, ...] - body_lines: tuple[str, ...] +def _variant_model_map( + definition: str, + keyword: str, + branch: str, + names: tuple[str, ...], +) -> dict[str, str]: + return {_inline_model_ref(definition, (keyword, index), (branch, None)): name for index, name in enumerate(names)} - def render(self) -> str: - lines = [ - f'@field_validator("{self.field_name}", mode="before")', - "@classmethod", - f"def {self.method_name}(cls, {self.argument_name}: Any) -> {self.return_type}:", - ] - lines.extend(f" # {line}" for line in self.comment_lines) - lines.extend(f" {line}" for line in self.body_lines) - return "\n".join(lines) - -DEFAULT_VALUE_OVERRIDES: tuple[tuple[str, str, str], ...] = ( - ("AgentCapabilities", "mcp_capabilities", "McpCapabilities()"), - ("AgentCapabilities", "session_capabilities", "SessionCapabilities()"), - ( - "AgentCapabilities", - "prompt_capabilities", - "PromptCapabilities()", +MODEL_NAME_MAP = { + "#/$defs/AvailableCommandsUpdate": "AvailableCommandsUpdateBase", + "#/$defs/ConfigOptionUpdate": "ConfigOptionUpdateBase", + "#/$defs/CurrentModeUpdate": "CurrentModeUpdateBase", + "#/$defs/SessionInfoUpdate": "SessionInfoUpdateBase", + "#/$defs/StringMultiSelectItems": "StringMultiSelectItemsBase", + "#/$defs/UsageUpdate": "UsageUpdateBase", +} +for variant_map in ( + _variant_model_map("AgentResponse", "anyOf", "object", ("AgentResponseMessage", "AgentErrorMessage")), + _variant_model_map("ClientResponse", "anyOf", "object", ("ClientResponseMessage", "ClientErrorMessage")), + _variant_model_map("AuthMethod", "anyOf", "allOf", ("EnvVarAuthMethod", "TerminalAuthMethod")), + _variant_model_map("McpServer", "anyOf", "allOf", ("HttpMcpServer", "SseMcpServer", "AcpMcpServer")), + _variant_model_map( + "SetSessionConfigOptionRequest", + "anyOf", + "object", + ("SetSessionConfigOptionBooleanRequest", "SetSessionConfigOptionSelectRequest"), + ), + _variant_model_map( + "ContentBlock", + "oneOf", + "allOf", + ( + "TextContentBlock", + "ImageContentBlock", + "AudioContentBlock", + "ResourceContentBlock", + "EmbeddedResourceContentBlock", + ), ), - ("ClientCapabilities", "fs", "FileSystemCapabilities()"), - ("ClientCapabilities", "terminal", "False"), - ( - "InitializeRequest", - "client_capabilities", - "ClientCapabilities()", + _variant_model_map( + "ToolCallContent", + "oneOf", + "allOf", + ("ContentToolCallContent", "FileEditToolCallContent", "TerminalToolCallContent"), ), - ( - "InitializeResponse", - "agent_capabilities", - "AgentCapabilities()", + _variant_model_map( + "PlanUpdateContent", + "oneOf", + "allOf", + ("PlanUpdateItems", "PlanUpdateFile", "PlanUpdateMarkdown"), ), -) - -# Classes that need a field_validator injected after generation. -CLASS_VALIDATOR_INJECTIONS: tuple[FieldValidatorInjection, ...] = ( - FieldValidatorInjection( - class_name="InitializeRequest", - field_name="protocol_version", - method_name="_coerce_protocol_version", - argument_name="value", - return_type="int", - comment_lines=( - 'Some clients (e.g. Zed) send a date string like "2024-11-05" instead', - "of an integer. The Rust SDK treats legacy strings as version 0; this", - "SDK maps unparsable values to 1 so the connection is not rejected.", - "See: https://github.com/agentclientprotocol/rust-sdk/blob/main/crates/agent-client-protocol-schema/src/version.rs", + _variant_model_map( + "NesSuggestion", + "oneOf", + "allOf", + ( + "NesEditSuggestionVariant", + "NesJumpSuggestionVariant", + "NesRenameSuggestionVariant", + "NesSearchAndReplaceSuggestionVariant", ), - body_lines=( - "if isinstance(value, int):", - " return value", - "try:", - " return int(value)", - "except (TypeError, ValueError):", - " return 1", + ), + _variant_model_map( + "SessionUpdate", + "oneOf", + "allOf", + ( + "UserMessageChunk", + "AgentMessageChunk", + "AgentThoughtChunk", + "ToolCallStart", + "ToolCallProgress", + "AgentPlanUpdate", + "AgentPlanContentUpdate", + "AgentPlanRemovedUpdate", + "AvailableCommandsUpdate", + "CurrentModeUpdate", + "ConfigOptionUpdate", + "SessionInfoUpdate", + "UsageUpdate", ), ), -) - - -@dataclass(frozen=True) -class _ProcessingStep: - """A named transformation applied to the generated schema content.""" - - name: str - apply: Callable[[str], str] - - -def main() -> None: - generate_schema() - - -def generate_schema() -> None: - if not SCHEMA_JSON.exists(): - print( - "Schema file missing. Ensure schema/schema.json exists (run gen_all.py --version to download).", - file=sys.stderr, - ) - sys.exit(1) - - with tempfile.TemporaryDirectory() as tmp_dir: - codegen_input = Path(tmp_dir) / "schema.codegen.json" - codegen_input.write_text(json.dumps(_preprocess_schema_for_codegen(_load_schema()), indent=2), encoding="utf-8") - - cmd = [ - sys.executable, - "-m", - "datamodel_code_generator", - "--input", - str(codegen_input), - "--input-file-type", - "jsonschema", - "--output", - str(SCHEMA_OUT), - "--target-python-version", - "3.12", - "--collapse-root-models", - "--output-model-type", - "pydantic_v2.BaseModel", - "--no-use-specialized-enum", - "--no-use-standard-collections", - "--no-use-union-operator", - "--type-overrides", - json.dumps(dict.fromkeys(STRING_ENUM_TYPES, "builtins.str")), - "--formatters", - "black", - "isort", - "--use-annotated", - "--use-field-description", - "--snake-case-field", - ] - - subprocess.check_call(cmd) # noqa: S603 - warnings = postprocess_generated_schema(SCHEMA_OUT) - for warning in warnings: - print(f"Warning: {warning}", file=sys.stderr) - - -def _load_schema() -> dict[str, Any]: - return json.loads(SCHEMA_JSON.read_text(encoding="utf-8")) - - -COMBINATOR_KEYS = ("oneOf", "anyOf") - - -def _preprocess_schema_for_codegen(schema: dict[str, Any]) -> dict[str, Any]: - schema = _normalize_catchall_unions(schema) - defs = schema.get("$defs", {}) - return _distribute_composed_object_schemas(schema, defs) - - -def _normalize_catchall_unions(node: Any) -> Any: - # ACP "custom or future" unions include a member tagged `"title": "other"` whose - # discriminator (type/mode/action) is a free-form string. datamodel-codegen cannot - # put that in a discriminated union, so it emits `#-special-#` placeholder literals. - # Drop the discriminator (the union is then validated structurally) and collapse the - # catch-all to a permissive object so unknown variants round-trip their raw payload. - if isinstance(node, list): - return [_normalize_catchall_unions(item) for item in node] - if not isinstance(node, dict): - return node - - transformed = {key: _normalize_catchall_unions(value) for key, value in node.items()} - for combinator in COMBINATOR_KEYS: - members = transformed.get(combinator) - if not isinstance(members, list): - continue - if not any(isinstance(member, dict) and member.get("title") == "other" for member in members): - continue - transformed.pop("discriminator", None) - transformed[combinator] = [ - _collapse_catchall_member(member) if isinstance(member, dict) and member.get("title") == "other" else member - for member in members - ] - return transformed - - -def _collapse_catchall_member(member: dict[str, Any]) -> dict[str, Any]: - collapsed: dict[str, Any] = {"type": "object", "additionalProperties": True} - for key in ("title", "description", "properties", "required"): - if key in member: - collapsed[key] = member[key] - return collapsed - - -def _distribute_composed_object_schemas(node: Any, defs: dict[str, Any]) -> Any: - if isinstance(node, list): - return [_distribute_composed_object_schemas(item, defs) for item in node] - if not isinstance(node, dict): - return node - - transformed = {key: _distribute_composed_object_schemas(value, defs) for key, value in node.items()} - for combinator in COMBINATOR_KEYS: - if combinator not in transformed or "properties" not in transformed: - continue - result = {combinator: _expand_composed_object_variants(transformed, defs)} - for key in ("title", "description", "discriminator"): - if key in transformed: - result[key] = transformed[key] - return result - return transformed - - -def _expand_composed_object_variants(node: dict[str, Any], defs: dict[str, Any]) -> list[Any]: - for combinator in COMBINATOR_KEYS: - if combinator not in node or "properties" not in node: - continue - - common_schema = _without_combinators(node) - expanded: list[Any] = [] - for option in node[combinator]: - for variant in _expand_allof_union_refs(option, defs): - expanded.append(_merge_object_schema(common_schema, variant) if isinstance(variant, dict) else variant) - return expanded - - return _expand_allof_union_refs(node, defs) - - -def _expand_allof_union_refs(node: Any, defs: dict[str, Any]) -> list[Any]: - if not isinstance(node, dict): - return [node] - - variants = [{key: copy.deepcopy(value) for key, value in node.items() if key != "allOf"}] - for item in node.get("allOf", []): - ref_name = _local_def_ref_name(item.get("$ref")) if isinstance(item, dict) else None - ref_schema = defs.get(ref_name) if ref_name else None - if isinstance(ref_schema, dict) and any(key in ref_schema for key in COMBINATOR_KEYS): - ref_variants = _expand_composed_object_variants(ref_schema, defs) - else: - ref_variants = [item] - - variants = [ - _merge_object_schema(variant, ref_variant) if isinstance(ref_variant, dict) else variant - for variant in variants - for ref_variant in ref_variants - ] - return variants - - -def _without_combinators(node: dict[str, Any]) -> dict[str, Any]: - return { - key: copy.deepcopy(value) - for key, value in node.items() - if key not in COMBINATOR_KEYS and key != "discriminator" - } - - -def _local_def_ref_name(ref: Any) -> str | None: - if isinstance(ref, str) and ref.startswith("#/$defs/"): - return ref.rsplit("/", 1)[-1] - return None - - -def _pop_ref_as_allof(schema: dict[str, Any]) -> tuple[dict[str, Any], list[dict[str, Any]]]: - schema = copy.deepcopy(schema) - if "$ref" not in schema: - return schema, [] - return schema, [{"$ref": schema.pop("$ref")}] - - -def _merge_object_schema(left: dict[str, Any], right: dict[str, Any]) -> dict[str, Any]: - left, left_refs = _pop_ref_as_allof(left) - right, right_refs = _pop_ref_as_allof(right) - merged: dict[str, Any] = {} + _variant_model_map( + "ElicitationFormMode", + "anyOf", + "allOf", + ("ElicitationFormSessionMode", "ElicitationFormRequestMode"), + ), + _variant_model_map( + "ElicitationUrlMode", + "anyOf", + "allOf", + ("ElicitationUrlSessionMode", "ElicitationUrlRequestMode"), + ), + _variant_model_map( + "ElicitationPropertySchema", + "anyOf", + "allOf", + ( + "ElicitationStringPropertySchema", + "ElicitationNumberPropertySchema", + "ElicitationIntegerPropertySchema", + "ElicitationBooleanPropertySchema", + "ElicitationMultiSelectPropertySchema", + ), + ), +): + MODEL_NAME_MAP.update(variant_map) + +MODEL_NAME_MAP.update({ + _inline_model_ref("RequestPermissionOutcome", ("oneOf", 0), ("object", None)): "DeniedOutcome", + _inline_model_ref("RequestPermissionOutcome", ("oneOf", 1), ("allOf", None)): "AllowedOutcome", + _inline_model_ref("CreateElicitationResponse", ("anyOf", 0), ("allOf", None)): "AcceptElicitationResponse", + _inline_model_ref("CreateElicitationResponse", ("anyOf", 1), ("object", None)): "DeclineElicitationResponse", + _inline_model_ref("CreateElicitationResponse", ("anyOf", 2), ("object", None)): "CancelElicitationResponse", + _inline_model_ref("CreateElicitationResponse", ("anyOf", 3), ("object", None)): "OtherElicitationResponse", + _inline_model_ref("ElicitationPropertySchema", ("anyOf", 5), ("object", None)): ("ElicitationOtherPropertySchema"), + _inline_model_ref("MultiSelectItems", ("anyOf", 0), ("allOf", None)): "StringMultiSelectItems", + _inline_model_ref("MultiSelectItems", ("anyOf", 1), ("object", None)): "OtherMultiSelectItems", + _inline_model_ref("CreateElicitationRequest", ("anyOf", 0), ("allOf", None), ("allOf", None)): ( + "CreateFormElicitationRequestBase" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 0), ("allOf", 0), ("allOf", None)): ( + "CreateFormSessionElicitationRequestBase" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 0), ("allOf", 1), ("allOf", None)): ( + "CreateFormRequestElicitationRequestBase" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 0), ("allOf", None), ("union_model-0", None)): ( + "CreateFormSessionElicitationRequest" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 0), ("allOf", None), ("union_model-1", None)): ( + "CreateFormRequestElicitationRequest" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 1), ("allOf", None), ("allOf", None)): ( + "CreateUrlElicitationRequestBase" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 1), ("allOf", 0), ("allOf", None)): ( + "CreateUrlSessionElicitationRequestBase" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 1), ("allOf", 1), ("allOf", None)): ( + "CreateUrlRequestElicitationRequestBase" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 1), ("allOf", None), ("union_model-0", None)): ( + "CreateUrlSessionElicitationRequest" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 1), ("allOf", None), ("union_model-1", None)): ( + "CreateUrlRequestElicitationRequest" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 2), ("anyOf", 0), ("allOf", None)): ( + "CreateOtherSessionElicitationRequest" + ), + _inline_model_ref("CreateElicitationRequest", ("anyOf", 2), ("anyOf", 1), ("allOf", None)): ( + "CreateOtherRequestElicitationRequest" + ), +}) + +# datamodel-code-generator owns schema interpretation and its internal model names. +# This block only preserves the Python names already published by the SDK. +COMPATIBILITY_ALIASES = textwrap.dedent(""" + PermissionOptionKind = Literal["allow_once", "allow_always", "reject_once", "reject_always"] + PlanEntryPriority = Literal["high", "medium", "low"] + PlanEntryStatus = Literal["pending", "in_progress", "completed"] + StopReason = Literal["end_turn", "max_tokens", "max_turn_requests", "refusal", "cancelled"] + ToolCallStatus = Literal["pending", "in_progress", "completed", "failed"] + ToolKind = Literal[ + "read", + "edit", + "delete", + "move", + "search", + "execute", + "think", + "fetch", + "switch_mode", + "other", + ] - for key in set(left) | set(right): - if key in COMBINATOR_KEYS or key in {"allOf", "discriminator"}: - continue - if key == "properties": - merged[key] = {**left.get(key, {}), **right.get(key, {})} - elif key == "required": - required = [] - for item in left.get(key, []) + right.get(key, []): - if item not in required: - required.append(item) - if required: - merged[key] = required - elif key in right: - merged[key] = right[key] - else: - merged[key] = left[key] + CreateOtherElicitationRequest = Union[ + CreateOtherSessionElicitationRequest, + CreateOtherRequestElicitationRequest, + ] + CreateFormElicitationRequest = Union[ + CreateFormSessionElicitationRequest, + CreateFormRequestElicitationRequest, + ] + CreateUrlElicitationRequest = Union[ + CreateUrlSessionElicitationRequest, + CreateUrlRequestElicitationRequest, + ] + CreateElicitationRequest = Union[ + CreateFormElicitationRequest, + CreateUrlElicitationRequest, + CreateOtherElicitationRequest, + ] - all_of = left_refs + left.get("allOf", []) + right_refs + right.get("allOf", []) - if all_of: - merged["allOf"] = all_of - return merged + CreateElicitationResponse = Union[ + AcceptElicitationResponse, + DeclineElicitationResponse, + CancelElicitationResponse, + OtherElicitationResponse, + ] + ElicitationMode = Union[ + ElicitationFormSessionMode, + ElicitationFormRequestMode, + ElicitationUrlSessionMode, + ElicitationUrlRequestMode, + ] + _AvailableCommandsUpdate = AvailableCommandsUpdateBase + _CurrentModeUpdate = CurrentModeUpdateBase + _ConfigOptionUpdate = ConfigOptionUpdateBase + _SessionInfoUpdate = SessionInfoUpdateBase + _UsageUpdate = UsageUpdateBase + _StringMultiSelectItems = StringMultiSelectItemsBase -def _required_nullable_fields(schema: dict[str, Any]) -> dict[str, list[str]]: - defs = schema.get("$defs", {}) - fields: dict[str, list[str]] = {} - for class_name, definition in defs.items(): - if not isinstance(definition, dict): - continue + class Jsonrpc(Enum): + field_2_0 = "2.0" + """).strip() - required = set(definition.get("required", [])) - if not required: - continue - properties = definition.get("properties", {}) - nullable_fields = [ - _schema_field_name(property_name) - for property_name in sorted(required) - if _schema_allows_null(properties.get(property_name), defs) - ] - if nullable_fields: - fields[class_name] = nullable_fields - return fields +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Generate ACP v1 schema bindings.") + parser.add_argument("--check", action="store_true", help="Fail if the committed bindings are stale.") + return parser.parse_args() -def _schema_allows_null(node: Any, defs: dict[str, Any]) -> bool: - if not isinstance(node, dict): - return False +def main() -> None: + args = parse_args() + if not generate_schema(check=args.check): + raise SystemExit(1) - schema_type = node.get("type") - if schema_type == "null" or (isinstance(schema_type, list) and "null" in schema_type): - return True - for combinator in COMBINATOR_KEYS: - if any(_schema_allows_null(option, defs) for option in node.get(combinator, [])): +def generate_schema(*, check: bool = False) -> bool: + candidate = render_schema() + current = SCHEMA_OUT.read_text(encoding="utf-8") if SCHEMA_OUT.exists() else "" + if check: + if current == candidate: return True - - ref_name = _local_def_ref_name(node.get("$ref")) - if ref_name is not None: - return _schema_allows_null(defs.get(ref_name), defs) - - return any(_schema_allows_null(option, defs) for option in node.get("allOf", [])) - - -def _schema_field_name(name: str) -> str: - if name.startswith("_"): - return "field" + name - return re.sub(r"(? list[str]: - if not output_path.exists(): - raise RuntimeError(f"Generated schema not found at {output_path}") - - raw_content = output_path.read_text(encoding="utf-8") - header_block = _build_header_block() - - content = _strip_existing_header(raw_content) - # Type overrides for builtins are rendered as imports in 0.64, but the - # annotations should continue to use Python's builtin `str` directly. - content = content.replace("from builtins import str\n", "") - content = _remove_unused_models(content) - content, leftover_classes = _rename_numbered_models(content) - - processing_steps: tuple[_ProcessingStep, ...] = ( - _ProcessingStep("apply field overrides", _apply_field_overrides), - _ProcessingStep("apply default overrides", _apply_default_overrides), - _ProcessingStep("restore required nullable fields", _restore_required_nullable_fields), - _ProcessingStep("ensure custom BaseModel", _ensure_custom_base_model), - _ProcessingStep("enable RootModel attribute docstrings", _enable_root_model_attribute_docstrings), - _ProcessingStep("inject field validators", _inject_field_validators), - _ProcessingStep("inject deserialize defaults", _inject_deserialize_defaults), - _ProcessingStep("inject schema aliases", _inject_schema_aliases), - ) - - for step in processing_steps: - content = step.apply(content) - - missing_targets = _find_missing_targets(content) - - content = _inject_enum_aliases(content) - content = _remove_unreferenced_root_models(content) - final_content = header_block + content.rstrip() + "\n" - if not final_content.endswith("\n"): - final_content += "\n" - output_path.write_text(final_content, encoding="utf-8") - - warnings: list[str] = [] - if leftover_classes: - warnings.append( - "Unrenamed schema models detected: " - + ", ".join(leftover_classes) - + ". Update RENAME_MAP in scripts/gen_schema.py." - ) - if missing_targets: - warnings.append( - "Renamed schema targets not found after generation: " - + ", ".join(sorted(missing_targets)) - + ". Check RENAME_MAP or upstream schema changes." + print( + "".join( + difflib.unified_diff( + current.splitlines(keepends=True), + candidate.splitlines(keepends=True), + fromfile=str(SCHEMA_OUT.relative_to(ROOT)), + tofile=f"{SCHEMA_OUT.relative_to(ROOT)} (generated)", + ) + ), + end="", ) - warnings.extend(_validate_schema_alignment()) - - return warnings - - -def _build_header_block() -> str: - header_lines = ["# Generated from schema/schema.json. Do not edit by hand."] - if VERSION_FILE.exists(): - ref = VERSION_FILE.read_text(encoding="utf-8").strip() - if ref: - header_lines.append(f"# Schema ref: {ref}") - return "\n".join(header_lines) + "\n\n" - - -def _strip_existing_header(content: str) -> str: - existing_header = re.match(r"(#.*\n)+", content) - if existing_header: - return content[existing_header.end() :].lstrip("\n") - return content.lstrip("\n") - - -def _rename_numbered_models(content: str) -> tuple[str, list[str]]: - renamed = content - for old, new in sorted(RENAME_MAP.items(), key=lambda item: len(item[0]), reverse=True): - if re.search(rf"\b{re.escape(new)}\b", renamed) is not None: - renamed = re.sub(rf"\b{re.escape(new)}\b", f"_{new}", renamed) - pattern = re.compile(rf"\b{re.escape(old)}\b") - renamed = pattern.sub(new, renamed) - - leftover_class_pattern = re.compile(r"^class (\w+\d+)\(", re.MULTILINE) - leftover_classes = sorted(set(leftover_class_pattern.findall(renamed))) - return renamed, leftover_classes - - -def _find_missing_targets(content: str) -> list[str]: - missing: list[str] = [] - for new_name in RENAME_MAP.values(): - pattern = re.compile(rf"^class {re.escape(new_name)}\(", re.MULTILINE) - if not pattern.search(content): - missing.append(new_name) - return missing + return False + SCHEMA_OUT.write_text(candidate, encoding="utf-8") + return True -def _validate_schema_alignment() -> list[str]: - warnings: list[str] = [] +def render_schema() -> str: + """Generate v1 directly from the currently pinned JSON Schema.""" if not SCHEMA_JSON.exists(): - warnings.append("schema/schema.json missing; unable to validate enum aliases.") - return warnings - - try: - schema_enums = _load_schema_enum_literals() - except json.JSONDecodeError as exc: - warnings.append(f"Failed to parse schema/schema.json: {exc}") - return warnings - - for enum_name, expected_values in ENUM_LITERAL_MAP.items(): - schema_values = schema_enums.get(enum_name) - if schema_values is None: - warnings.append( - f"Enum '{enum_name}' not found in schema.json; update ENUM_LITERAL_MAP or investigate schema changes." - ) - continue - if tuple(schema_values) != expected_values: - warnings.append( - f"Enum mismatch for '{enum_name}': schema.json -> {schema_values}, generated aliases -> {expected_values}" - ) - - detected_unions = _detect_extensible_unions() - if detected_unions != set(EXTENSIBLE_UNIONS): - warnings.append( - f"Extensible union drift: schema defines {sorted(detected_unions)}, " - f"EXTENSIBLE_UNIONS lists {sorted(EXTENSIBLE_UNIONS)}. Update EXTENSIBLE_UNIONS, the " - "RENAME_MAP catch-all names, and the alias template together." - ) - return warnings - - -def _detect_extensible_unions() -> set[str]: - defs = _load_schema().get("$defs", {}) - detected: set[str] = set() - for name, definition in defs.items(): - if not isinstance(definition, dict) or "discriminator" not in definition: - continue - members = definition.get("anyOf") or definition.get("oneOf") or [] - if any(isinstance(member, dict) and member.get("title") == "other" for member in members): - detected.add(name) - return detected - - -def _load_schema_enum_literals() -> dict[str, tuple[str, ...]]: - schema_data = json.loads(SCHEMA_JSON.read_text(encoding="utf-8")) - defs = schema_data.get("$defs", {}) - enum_literals: dict[str, tuple[str, ...]] = {} - - for name, definition in defs.items(): - values: list[str] = [] - if "enum" in definition: - values = [str(item) for item in definition["enum"]] - elif "oneOf" in definition: - values = [ - str(option["const"]) - for option in definition.get("oneOf", []) - if isinstance(option, dict) and "const" in option - ] - if values: - enum_literals[name] = tuple(values) - - return enum_literals - - -def _ensure_custom_base_model(content: str) -> str: - if "class BaseModel(_BaseModel):" in content: - return content - lines = content.splitlines() - for idx, line in enumerate(lines): - if not line.startswith("from pydantic import "): - continue - imports = [part.strip() for part in line[len("from pydantic import ") :].split(",")] - has_alias = any(part == "BaseModel as _BaseModel" for part in imports) - has_config = any(part == "ConfigDict" for part in imports) - new_imports = [] - for part in imports: - if part == "BaseModel": - new_imports.append("BaseModel as _BaseModel") - has_alias = True - else: - new_imports.append(part) - if not has_alias: - new_imports.append("BaseModel as _BaseModel") - if not has_config: - new_imports.append("ConfigDict") - lines[idx] = "from pydantic import " + ", ".join(new_imports) - to_insert = textwrap.dedent("""\ - class BaseModel(_BaseModel): - model_config = ConfigDict(populate_by_name=True, use_attribute_docstrings=True) - - def __getattr__(self, item: str) -> Any: - if item.lower() != item: - snake_cased = "".join("_" + c.lower() if c.isupper() and i > 0 else c.lower() for i, c in enumerate(item)) - return getattr(self, snake_cased) - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{item}'") - """) - insert_idx = idx + 1 - lines.insert(insert_idx, "") - for offset, line in enumerate(to_insert.splitlines(), 1): - lines.insert(insert_idx + offset, line) - break - return "\n".join(lines) + "\n" - - -def _enable_root_model_attribute_docstrings(content: str) -> str: - lines = content.splitlines(keepends=True) - tree = ast.parse(content) - insertion_points = [ - node.body[0].lineno - 1 - for node in tree.body - if isinstance(node, ast.ClassDef) - and node.body - and any( - isinstance(base, ast.Subscript) and isinstance(base.value, ast.Name) and base.value.id == "RootModel" - for base in node.bases - ) - ] - for line_index in reversed(insertion_points): - lines.insert(line_index, " model_config = ConfigDict(use_attribute_docstrings=True)\n\n") - return "".join(lines) - - -def _ensure_pydantic_import(content: str, name: str) -> str: - """Add *name* to the ``from pydantic import ...`` line if not already present.""" - lines = content.splitlines() - for idx, line in enumerate(lines): - if not line.startswith("from pydantic import "): - continue - imports = [part.strip() for part in line[len("from pydantic import ") :].split(",")] - if name not in imports: - imports.append(name) - lines[idx] = "from pydantic import " + ", ".join(imports) - return "\n".join(lines) + "\n" - return content - - -def _extensible_union_excluded_tags(union_def: dict[str, Any], discriminator: str) -> tuple[str, ...]: - members = union_def.get("anyOf") or union_def.get("oneOf") or [] - other = next((member for member in members if isinstance(member, dict) and member.get("title") == "other"), None) - if other is None: - return () - tags: list[str] = [] - for excluded in other.get("not", {}).get("anyOf", []): - const = excluded.get("properties", {}).get(discriminator, {}).get("const") - if isinstance(const, str) and const not in tags: - tags.append(const) - return tuple(tags) - - -def _catchall_exclusion_injections() -> list[FieldValidatorInjection]: - defs = _load_schema().get("$defs", {}) - injections: list[FieldValidatorInjection] = [] - for union_name, catchall_class in EXTENSIBLE_UNIONS.items(): - union_def = defs.get(union_name) - if not isinstance(union_def, dict): - continue - discriminator = union_def.get("discriminator", {}).get("propertyName") - if not discriminator: - continue - tags = _extensible_union_excluded_tags(union_def, discriminator) - if not tags: - continue - field = _schema_field_name(discriminator) - injections.append( - FieldValidatorInjection( - class_name=catchall_class, - field_name=field, - method_name=f"_reject_known_{field}", - argument_name="value", - return_type="Any", - comment_lines=( - "Restore the schema's `not` clause dropped for codegen: reject the known", - "variants' discriminator values so a malformed known variant fails instead", - "of silently parsing as this catch-all.", - ), - body_lines=( - f"if value in {tags!r}:", - f' raise ValueError("{field} value is reserved by a known variant")', - "return value", - ), - ) - ) - return injections - - -def _inject_field_validators(content: str) -> str: - """Inject field_validator methods for CLASS_VALIDATOR_INJECTIONS and catch-all exclusions.""" - for injection in (*CLASS_VALIDATOR_INJECTIONS, *_catchall_exclusion_injections()): - content = _ensure_pydantic_import(content, "field_validator") - - class_pattern = re.compile( - rf"(class {injection.class_name}\(BaseModel\):)(.*?)(?=\nclass |\Z)", - re.DOTALL, - ) - - def _append_validator( - match: re.Match[str], - _injection: FieldValidatorInjection = injection, - ) -> str: - header, block = match.group(1), match.group(2) - indented = "\n" + textwrap.indent(_injection.render(), " ") - return header + block + indented + "\n" - - content, count = class_pattern.subn(_append_validator, content, count=1) - if count == 0: - print( - f"Warning: class {injection.class_name} not found for validator injection", - file=sys.stderr, - ) - return content - - -def _inject_deserialize_defaults(content: str) -> str: - defs = _load_schema().get("$defs", {}) - - # `_meta` carries x-deserialize-default-on-error on almost every model; handle it once - # on the shared BaseModel with check_fields=False so every subclass inherits the salvage. - meta_validator = ( - '@field_validator("field_meta", mode="wrap", check_fields=False)\n' - "@classmethod\n" - "def _salvage_meta_on_error(cls, value: Any, handler: Any) -> Any:\n" - " return salvage_on_error(value, handler, lambda: None)\n" + raise FileNotFoundError("schema/schema.json is missing; fetch a pinned schema release first") + + schema = _schema_for_codegen(json.loads(SCHEMA_JSON.read_text(encoding="utf-8"))) + generated = generate( + schema, + input_file_type=InputFileType.JsonSchema, + custom_file_header=_build_header(), + target_python_version=PythonVersion.PY_310, + collapse_root_models=True, + skip_root_model=True, + output_model_type=DataModelType.PydanticV2BaseModel, + base_class="acp._schema_base.BaseModel", + use_specialized_enum=False, + use_standard_collections=False, + use_union_operator=False, + additional_imports=["enum.Enum"], + enum_field_as_literal=LiteralType.All, + use_one_literal_as_default=True, + validators=_build_validators_config(schema), + formatters=[Formatter.BUILTIN], + infer_union_variant_names=True, + naming_strategy=NamingStrategy.PrimaryFirst, + model_name_map=MODEL_NAME_MAP, + strict_refs=True, + schema_version="2020-12", + schema_version_mode=VersionMode.Strict, + type_mappings=list(UNSIGNED_TYPE_MAPPINGS), + generate_schema_validators=True, + use_annotated=True, + field_constraints=True, + use_field_description=True, + snake_case_field=True, ) - content, count = _append_class_method(content, r"class BaseModel\(_BaseModel\):", meta_validator) - if count == 0: - print("Warning: custom BaseModel not found for _meta salvage injection", file=sys.stderr) - - for class_name, definition in defs.items(): + if not isinstance(generated, str): + raise TypeError("Schema generation did not produce a single Python module") + return _format_python(f"{generated.rstrip()}\n\n\n{COMPATIBILITY_ALIASES}\n") + + +def _schema_for_codegen(schema: dict[str, Any]) -> dict[str, Any]: + """Drop open-union constraints that Pydantic cannot represent statically.""" + patched = copy.deepcopy(schema) + for name, catchall_index in OPEN_UNIONS.items(): + try: + del patched["$defs"][name]["discriminator"] + del patched["$defs"][name]["anyOf"][catchall_index]["not"] + except KeyError: + raise ValueError(f"{name} no longer has the expected open-union shape") from None + return patched + + +def _build_validators_config(schema: dict[str, Any]) -> dict[str, ModelValidators]: + validators: dict[str, list[ValidatorDefinition]] = { + "InitializeRequest": [ + ValidatorDefinition( + field="protocol_version", + function="acp._deserialize.coerce_protocol_version", + mode=ValidatorMode.BEFORE, + ) + ] + } + for class_name, definition in schema.get("$defs", {}).items(): if not isinstance(definition, dict): continue - salvage_groups, skip_fields = _deserialize_field_specs(definition) - methods: list[str] = [] - for index, (fallback, fields) in enumerate(sorted(salvage_groups.items())): - arguments = ", ".join(f'"{field}"' for field in sorted(fields)) - methods.append( - f'@field_validator({arguments}, mode="wrap")\n' - "@classmethod\n" - f"def _salvage_on_error_{index}(cls, value: Any, handler: Any) -> Any:\n" - f" return salvage_on_error(value, handler, {fallback})\n" - ) - for index, field in enumerate(sorted(skip_fields)): - methods.append( - f'@field_validator("{field}", mode="wrap")\n' - "@classmethod\n" - f"def _skip_invalid_items_{index}(cls, value: Any, handler: Any) -> Any:\n" - " return skip_invalid_items(value, handler)\n" + default_fields, skip_fields = _deserialize_field_specs(definition) + definitions = validators.setdefault(class_name, []) + definitions.extend( + ValidatorDefinition(fields=sorted(fields), function=function, mode=ValidatorMode.WRAP) + for fields, function in ( + (default_fields, "acp._deserialize.use_default_on_error"), + (skip_fields, "acp._deserialize.skip_invalid_items"), ) - # A plain object $def renders as `class Name(BaseModel)` (or `_Name` after a - # collision rename). A union $def has no class of its own; its common properties - # distribute to the member variant classes, so target those instead. - targets = [rf"class _?{re.escape(class_name)}\(BaseModel\):"] - members = _union_member_classes(class_name) - if members: - targets = [rf"class {re.escape(member)}\(\w+\):" for member in members] - for method in methods: - for target in targets: - content, count = _append_class_method(content, target, method) - if count == 0: - print(f"Warning: no class matched {target!r} for deserialize injection", file=sys.stderr) - - content = _ensure_pydantic_import(content, "field_validator") - return _ensure_deserialize_import(content) - - -def _union_member_classes(union_name: str) -> list[str]: - return [new for old, new in RENAME_MAP.items() if re.fullmatch(rf"{re.escape(union_name)}\d+", old)] + if fields + ) + return { + class_name: ModelValidators(validators=definitions) + for class_name, definitions in validators.items() + if definitions + } -def _deserialize_field_specs(definition: dict[str, Any]) -> tuple[dict[str, list[str]], list[str]]: - """Return ({fallback_expr: [field, ...]}, [skip_field, ...]) for a $def. `_meta` is handled - on the shared BaseModel and excluded here.""" +def _deserialize_field_specs(definition: dict[str, Any]) -> tuple[list[str], list[str]]: required = set(definition.get("required", [])) - salvage: dict[str, list[str]] = {} + use_default: list[str] = [] skip: list[str] = [] - for prop_name, prop in definition.get("properties", {}).items(): - if not isinstance(prop, dict) or prop_name == "_meta": + for property_name, property_schema in definition.get("properties", {}).items(): + if not isinstance(property_schema, dict) or property_name == "_meta": continue - field = _schema_field_name(prop_name) - if prop.get("x-deserialize-skip-invalid-items"): - skip.append(field) - elif prop.get("x-deserialize-default-on-error"): - salvage.setdefault(_fallback_expression(prop, prop_name in required), []).append(field) - return salvage, skip - - -def _fallback_expression(prop: dict[str, Any], is_required: bool) -> str: - if "default" in prop: - return f"lambda: {prop['default']!r}" - if _is_array_schema(prop) and (is_required or not _schema_allows_null(prop, {})): - return "lambda: []" - return "lambda: None" - - -def _is_array_schema(prop: dict[str, Any]) -> bool: - prop_type = prop.get("type") - if prop_type == "array" or (isinstance(prop_type, list) and "array" in prop_type): - return True - return "items" in prop - - -def _append_class_method(content: str, header_pattern: str, method_text: str) -> tuple[str, int]: - pattern = re.compile(rf"({header_pattern})(.*?)(?=\nclass |\Z)", re.DOTALL) - - def _append(match: re.Match[str]) -> str: - indented = "\n" + textwrap.indent(method_text, " ") - return match.group(1) + match.group(2) + indented + "\n" - - return pattern.subn(_append, content, count=1) - - -def _ensure_deserialize_import(content: str) -> str: - # Absolute import (not relative): gen_signature.py loads schema.py as a standalone - # module with no package context, where `from ._deserialize` cannot resolve. - statement = "from acp._deserialize import salvage_on_error, skip_invalid_items" - if statement in content: - return content - lines = content.splitlines() - for idx, line in enumerate(lines): - if line.startswith("from pydantic import "): - lines.insert(idx + 1, statement) - return "\n".join(lines) + "\n" - return content - - -def _inject_schema_aliases(content: str) -> str: - if "CreateElicitationRequest = Union[" in content: - return content - - aliases = textwrap.dedent("""\ - ElicitationMode = Union[ - ElicitationFormSessionMode, - ElicitationFormRequestMode, - ElicitationUrlSessionMode, - ElicitationUrlRequestMode, - ] - CreateFormElicitationRequest = Union[ - CreateFormSessionElicitationRequest, - CreateFormRequestElicitationRequest, - ] - CreateUrlElicitationRequest = Union[ - CreateUrlSessionElicitationRequest, - CreateUrlRequestElicitationRequest, - ] - CreateElicitationRequest = Union[ - CreateFormElicitationRequest, - CreateUrlElicitationRequest, - CreateOtherElicitationRequest, - ] - CreateElicitationResponse = Union[ - AcceptElicitationResponse, - DeclineElicitationResponse, - CancelElicitationResponse, - OtherElicitationResponse, - ] - """) - pattern = re.compile( - r"^(class CreateFormRequestElicitationRequest\([\s\S]*?\):[\s\S]*?)(?=^class \w+\(|\Z)", - re.MULTILINE, + field_name = to_snake(property_name) + if property_schema.get("x-deserialize-skip-invalid-items"): + skip.append(field_name) + elif property_schema.get("x-deserialize-default-on-error"): + if property_name in required: + raise ValueError(f"{property_name!r} requests default-on-error but is required") + use_default.append(field_name) + return use_default, skip + + +def _build_header() -> str: + lines = ["# Generated from schema/schema.json. Do not edit by hand."] + if VERSION_FILE.exists() and (ref := VERSION_FILE.read_text(encoding="utf-8").strip()): + lines.append(f"# Schema ref: {ref}") + return "\n".join(lines) + + +def _format_python(source: str) -> str: + commands = ( + ("check", "--fix"), + ("format",), ) - content, count = pattern.subn(lambda match: match.group(1).rstrip() + "\n\n" + aliases + "\n", content, count=1) - if count == 0: - print("Warning: failed to insert schema aliases", file=sys.stderr) - return content - - -def _restore_required_nullable_fields(content: str, schema: dict[str, Any] | None = None) -> str: - schema = _load_schema() if schema is None else schema - for class_name, field_names in _required_nullable_fields(schema).items(): - class_pattern = re.compile( - rf"(class {re.escape(class_name)}\([^)]*\):)(.*?)(?=\nclass |\Z)", - re.DOTALL, - ) - - def restore_block(match: re.Match[str], _field_names: list[str] = field_names) -> str: - header, block = match.group(1), match.group(2) - for field_name in _field_names: - field_patterns = ( - re.compile(rf"(\n\s+{re.escape(field_name)}:[^\n]*?)\s*=\s*None(?=\n)"), - re.compile(rf"(\n\s+{re.escape(field_name)}:[^\n]*\[\s*\n[\s\S]*?\n\s+\]\s*)=\s*None"), - ) - for field_pattern in field_patterns: - block, count = field_pattern.subn(r"\1", block, count=1) - if count: - break - return header + block - - content = class_pattern.sub(restore_block, content, count=1) - return content - - -def _apply_field_overrides(content: str) -> str: - for class_name, field_name, new_type, optional in FIELD_TYPE_OVERRIDES: - old_type = "Optional[str]" if optional else "str" - replacement_type = f"Optional[{new_type}]" if optional else new_type - pattern = re.compile( - rf"(class {re.escape(class_name)}\(BaseModel\):.*?\n\s+{re.escape(field_name)}:\s+" - rf"(?:Annotated\[\s*)?){re.escape(old_type)}(?=\s*(?:,|=|\n))", - re.DOTALL, - ) - content, count = pattern.subn(rf"\g<1>{replacement_type}", content, count=1) - if count == 0: - print( - f"Warning: failed to apply type override for {class_name}.{field_name} -> {new_type}", - file=sys.stderr, - ) - return content - - -def _apply_default_overrides(content: str) -> str: - for class_name, field_name, replacement in DEFAULT_VALUE_OVERRIDES: - class_pattern = re.compile( - rf"(class {class_name}\(BaseModel\):)(.*?)(?=\nclass |\Z)", - re.DOTALL, - ) - - def replace_block( - match: re.Match[str], - _field_name: str = field_name, - _replacement: str = replacement, - _class_name: str = class_name, - ) -> str: - header, block = match.group(1), match.group(2) - field_patterns: tuple[tuple[re.Pattern[str], Callable[[re.Match[str]], str]], ...] = ( - ( - re.compile( - rf"(\n\s+{_field_name}:.*?\]\s*=\s*)([\s\S]*?)" - rf"(?=\n\s{{4}}(?:[A-Za-z_][A-Za-z0-9_]*\s*:|[rRuUbBfF]*(?:'''|\"\"\"))|$)", - re.DOTALL, - ), - lambda m, _rep=_replacement: m.group(1) + _rep, - ), - ( - re.compile( - rf"(\n\s+{_field_name}:[^\n]*=)\s*([^\n]+)", - re.MULTILINE, - ), - lambda m, _rep=_replacement: m.group(1) + " " + _rep, - ), - ) - for pattern, replacer in field_patterns: - new_block, count = pattern.subn(replacer, block, count=1) - if count: - return header + new_block - print( - f"Warning: failed to override default for {_class_name}.{_field_name}", - file=sys.stderr, - ) - return match.group(0) - - content, count = class_pattern.subn(replace_block, content, count=1) - if count == 0: - print( - f"Warning: class {class_name} not found for default override on {field_name}", - file=sys.stderr, - ) - return content - - -def _inject_enum_aliases(content: str) -> str: - enum_lines = [ - f"{name} = Literal[{', '.join(repr(value) for value in values)}]" for name, values in ENUM_LITERAL_MAP.items() - ] - if not enum_lines: - return content - block = "\n".join(enum_lines) + "\n\n" - class_index = content.find("\nclass ") - if class_index == -1: - return content - insertion_point = class_index + 1 # include leading newline - return content[:insertion_point] + block + content[insertion_point:] - - -def _remove_unreferenced_root_models(content: str) -> str: - tree = ast.parse(content) - root_models = { - node.name: node - for node in tree.body - if isinstance(node, ast.ClassDef) - and any( - isinstance(base, ast.Subscript) and isinstance(base.value, ast.Name) and base.value.id == "RootModel" - for base in node.bases - ) - } - - referenced_roots = set(PUBLIC_ROOT_MODELS) - root_dependencies: dict[str, set[str]] = {} - for statement in tree.body: - loaded_names = { - node.id - for node in ast.walk(statement) - if isinstance(node, ast.Name) and isinstance(node.ctx, ast.Load) and node.id in root_models - } - if isinstance(statement, ast.ClassDef) and statement.name in root_models: - root_dependencies[statement.name] = loaded_names - else: - referenced_roots.update(loaded_names) - - pending = list(referenced_roots) - while pending: - root_name = pending.pop() - for dependency in root_dependencies.get(root_name, set()) - referenced_roots: - referenced_roots.add(dependency) - pending.append(dependency) - - unused_models = [model for name, model in root_models.items() if name not in referenced_roots] - lines = content.splitlines(keepends=True) - for model in sorted(unused_models, key=lambda item: item.lineno, reverse=True): - del lines[model.lineno - 1 : model.end_lineno] - return re.sub(r"\n{4,}", "\n\n\n", "".join(lines)) - - -def _remove_unused_models(content: str) -> str: - for model_name in MODELS_TO_REMOVE: - pattern = re.compile( - rf"^(class {model_name}\([\s\S]*?\):)([\s\S]*?)(?=^\S|\Z)", - re.MULTILINE, + for arguments in commands: + result = subprocess.run( # noqa: S603 + [sys.executable, "-m", "ruff", *arguments, "--stdin-filename", str(SCHEMA_OUT), "-"], + input=source, + text=True, + capture_output=True, + check=False, + cwd=ROOT, ) - content, count = pattern.subn("", content) - if count > 0: - print(f"Removed unused model: {model_name}", file=sys.stderr) - return content + if result.returncode: + raise RuntimeError(f"ruff {' '.join(arguments)} failed:\n{result.stderr}") + source = result.stdout + return source if __name__ == "__main__": diff --git a/src/acp/_deserialize.py b/src/acp/_deserialize.py index 1cc2ea9..d847d8e 100644 --- a/src/acp/_deserialize.py +++ b/src/acp/_deserialize.py @@ -12,6 +12,25 @@ from typing import Any from pydantic import ValidationError +from pydantic_core import PydanticUseDefault + + +def coerce_protocol_version(value: Any, _info: Any = None) -> int: + """Coerce legacy string protocol versions without rejecting the connection.""" + if isinstance(value, int): + return value + try: + return int(value) + except (TypeError, ValueError): + return 1 + + +def use_default_on_error(value: Any, handler: Callable[[Any], Any], _info: Any = None) -> Any: + """Ask Pydantic to use the declared field default after validation fails.""" + try: + return handler(value) + except ValidationError: + raise PydanticUseDefault() from None def salvage_on_error(value: Any, handler: Callable[[Any], Any], fallback: Callable[[], Any]) -> Any: @@ -26,7 +45,7 @@ def salvage_on_error(value: Any, handler: Callable[[Any], Any], fallback: Callab return fallback() -def skip_invalid_items(value: Any, handler: Callable[[Any], Any]) -> Any: +def skip_invalid_items(value: Any, handler: Callable[[Any], Any], _info: Any = None) -> Any: """Drop array items that fail validation instead of failing the whole array. Restores ``x-deserialize-skip-invalid-items``. Each item is validated through the field's diff --git a/src/acp/_schema_base.py b/src/acp/_schema_base.py new file mode 100644 index 0000000..ad979b3 --- /dev/null +++ b/src/acp/_schema_base.py @@ -0,0 +1,77 @@ +from __future__ import annotations + +from typing import Any, ClassVar, Literal, get_args, get_origin + +import pydantic +from pydantic import ( + ConfigDict, + SerializationInfo, + SerializerFunctionWrapHandler, + field_validator, + model_serializer, + model_validator, +) +from pydantic.alias_generators import to_snake + +from ._deserialize import use_default_on_error + + +class BaseModel(pydantic.BaseModel): + """Runtime behavior shared by generated ACP schema models.""" + + _reserved_tags: ClassVar[dict[str, tuple[str, frozenset[str]]]] = { + "CreateOtherSessionElicitationRequest": ("mode", frozenset({"form", "url"})), + "CreateOtherRequestElicitationRequest": ("mode", frozenset({"form", "url"})), + "OtherElicitationResponse": ("action", frozenset({"accept", "cancel", "decline"})), + "ElicitationOtherPropertySchema": ( + "type", + frozenset({"array", "boolean", "integer", "number", "string"}), + ), + "OtherMultiSelectItems": ("type", frozenset({"string"})), + } + + model_config = ConfigDict( + populate_by_name=True, + use_attribute_docstrings=True, + ) + + def __getattr__(self, item: str) -> Any: + if item.lower() != item: + return getattr(self, to_snake(item)) + raise AttributeError(f"{type(self).__name__!r} object has no attribute {item!r}") + + @model_serializer(mode="wrap") + def _include_literal_defaults( + self, + handler: SerializerFunctionWrapHandler, + info: SerializationInfo, + ) -> Any: + data = handler(self) + for name, field in type(self).model_fields.items(): + annotation = field.annotation + if info.include is not None and name not in info.include: + continue + if info.exclude is not None and name in info.exclude: + continue + if get_origin(annotation) is not Literal or len(get_args(annotation)) != 1: + continue + key = (field.serialization_alias or field.alias or name) if info.by_alias else name + data[key] = getattr(self, name) + return data + + @field_validator("field_meta", mode="wrap", check_fields=False) + @classmethod + def _use_meta_default_on_error(cls, value: Any, handler: Any) -> Any: + return use_default_on_error(value, handler) + + @model_validator(mode="before") + @classmethod + def _reject_malformed_known_variant(cls, value: Any) -> Any: + rule = cls._reserved_tags.get(cls.__name__) + if rule is None or not isinstance(value, dict): + return value + wire_field, reserved = rule + tag = value.get(wire_field, value.get(to_snake(wire_field))) + if tag in reserved: + raise ValueError(f"{wire_field} value is reserved by a known variant") + return value diff --git a/src/acp/schema.py b/src/acp/schema.py index 5e1be37..f525578 100644 --- a/src/acp/schema.py +++ b/src/acp/schema.py @@ -6,34 +6,9 @@ from enum import Enum from typing import Annotated, Any, Dict, List, Literal, Optional, Union -from pydantic import AnyUrl, BaseModel as _BaseModel, ConfigDict, Field, RootModel, field_validator -from acp._deserialize import salvage_on_error, skip_invalid_items - -PermissionOptionKind = Literal["allow_once", "allow_always", "reject_once", "reject_always"] -PlanEntryPriority = Literal["high", "medium", "low"] -PlanEntryStatus = Literal["pending", "in_progress", "completed"] -StopReason = Literal["end_turn", "max_tokens", "max_turn_requests", "refusal", "cancelled"] -ToolCallStatus = Literal["pending", "in_progress", "completed", "failed"] -ToolKind = Literal["read", "edit", "delete", "move", "search", "execute", "think", "fetch", "switch_mode", "other"] - - -class BaseModel(_BaseModel): - model_config = ConfigDict(populate_by_name=True, use_attribute_docstrings=True) - - def __getattr__(self, item: str) -> Any: - if item.lower() != item: - snake_cased = "".join("_" + c.lower() if c.isupper() and i > 0 else c.lower() for i, c in enumerate(item)) - return getattr(self, snake_cased) - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{item}'") - - @field_validator("field_meta", mode="wrap", check_fields=False) - @classmethod - def _salvage_meta_on_error(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) - - -class Jsonrpc(Enum): - field_2_0 = "2.0" +from acp._deserialize import coerce_protocol_version, skip_invalid_items, use_default_on_error +from acp._schema_base import BaseModel +from pydantic import AnyUrl, ConfigDict, Field, RootModel, ValidationInfo, ValidatorFunctionWrapHandler, field_validator class ReadTextFileRequest(BaseModel): @@ -64,8 +39,8 @@ class ReadTextFileRequest(BaseModel): @field_validator("limit", "line", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class TextResourceContents(BaseModel): @@ -92,8 +67,8 @@ class TextResourceContents(BaseModel): @field_validator("mime_type", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class BlobResourceContents(BaseModel): @@ -120,8 +95,8 @@ class BlobResourceContents(BaseModel): @field_validator("mime_type", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class Diff(BaseModel): @@ -148,8 +123,8 @@ class Diff(BaseModel): @field_validator("old_text", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class Terminal(BaseModel): @@ -187,8 +162,8 @@ class ToolCallLocation(BaseModel): @field_validator("line", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class EnvVariable(BaseModel): @@ -286,10 +261,7 @@ class KillTerminalRequest(BaseModel): """ -class CreateOtherElicitationRequest(BaseModel): - model_config = ConfigDict( - extra="allow", - ) +class CreateFormElicitationRequestBase(BaseModel): message: str """ A human-readable message describing what input is needed. @@ -302,24 +274,23 @@ class CreateOtherElicitationRequest(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - mode: str - """ - Custom or future elicitation mode. + mode: Literal["form"] = "form" - Values beginning with `_` are reserved for implementation-specific - extensions. Unknown values that do not begin with `_` are reserved for - future ACP variants. + +class CreateUrlElicitationRequestBase(BaseModel): + message: str + """ + A human-readable message describing what input is needed. + """ + field_meta: Annotated[Optional[Dict[str, Any]], Field(alias="_meta")] = None """ + The _meta property is reserved by ACP to allow clients and agents to attach additional + metadata to their interactions. Implementations MUST NOT make assumptions about values at + these keys. - @field_validator("mode", mode="before") - @classmethod - def _reject_known_mode(cls, value: Any) -> Any: - # Restore the schema's `not` clause dropped for codegen: reject the known - # variants' discriminator values so a malformed known variant fails instead - # of silently parsing as this catch-all. - if value in ("form", "url"): - raise ValueError("mode value is reserved by a known variant") - return value + See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) + """ + mode: Literal["url"] = "url" class ElicitationSessionScope(BaseModel): @@ -334,8 +305,8 @@ class ElicitationSessionScope(BaseModel): @field_validator("tool_call_id", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ElicitationRequestScope(BaseModel): @@ -358,16 +329,6 @@ class ElicitationOtherPropertySchema(BaseModel): future ACP variants. """ - @field_validator("type", mode="before") - @classmethod - def _reject_known_type(cls, value: Any) -> Any: - # Restore the schema's `not` clause dropped for codegen: reject the known - # variants' discriminator values so a malformed known variant fails instead - # of silently parsing as this catch-all. - if value in ("string", "number", "integer", "boolean", "array"): - raise ValueError("type value is reserved by a known variant") - return value - class EnumOption(BaseModel): const: str @@ -393,8 +354,8 @@ class EnumOption(BaseModel): @field_validator("description", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class StringPropertySchema(BaseModel): @@ -418,7 +379,7 @@ class StringPropertySchema(BaseModel): """ Pattern the string must match. """ - format: Optional[str] = None + format: Optional[Literal["email", "uri", "date", "date-time"]] = None """ String format. """ @@ -445,8 +406,8 @@ class StringPropertySchema(BaseModel): @field_validator("default", "description", "title", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NumberPropertySchema(BaseModel): @@ -481,8 +442,8 @@ class NumberPropertySchema(BaseModel): @field_validator("default", "description", "title", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class IntegerPropertySchema(BaseModel): @@ -517,8 +478,8 @@ class IntegerPropertySchema(BaseModel): @field_validator("default", "description", "title", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class BooleanPropertySchema(BaseModel): @@ -545,8 +506,8 @@ class BooleanPropertySchema(BaseModel): @field_validator("default", "description", "title", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class OtherMultiSelectItems(BaseModel): @@ -562,18 +523,8 @@ class OtherMultiSelectItems(BaseModel): future ACP variants. """ - @field_validator("type", mode="before") - @classmethod - def _reject_known_type(cls, value: Any) -> Any: - # Restore the schema's `not` clause dropped for codegen: reject the known - # variants' discriminator values so a malformed known variant fails instead - # of silently parsing as this catch-all. - if value in ("string",): - raise ValueError("type value is reserved by a known variant") - return value - -class _StringMultiSelectItems(BaseModel): +class StringMultiSelectItemsBase(BaseModel): enum: List[str] """ Allowed enum values. @@ -626,8 +577,6 @@ class ElicitationUrlRequestMode(ElicitationRequestScope): class ElicitationUrlMode(RootModel[Union[ElicitationUrlSessionMode, ElicitationUrlRequestMode]]): - model_config = ConfigDict(use_attribute_docstrings=True) - root: Union[ElicitationUrlSessionMode, ElicitationUrlRequestMode] """ **UNSTABLE** @@ -680,8 +629,8 @@ class PromptCapabilities(BaseModel): @field_validator("audio", "embedded_context", "image", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: False) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class McpCapabilities(BaseModel): @@ -712,8 +661,8 @@ class McpCapabilities(BaseModel): @field_validator("acp", "http", "sse", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: False) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class SessionListCapabilities(BaseModel): @@ -864,8 +813,8 @@ class NesRecentFilesCapabilities(BaseModel): @field_validator("max_count", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NesRelatedSnippetsCapabilities(BaseModel): @@ -895,8 +844,8 @@ class NesEditHistoryCapabilities(BaseModel): @field_validator("max_count", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NesUserActionsCapabilities(BaseModel): @@ -915,8 +864,8 @@ class NesUserActionsCapabilities(BaseModel): @field_validator("max_count", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NesOpenFilesCapabilities(BaseModel): @@ -972,20 +921,10 @@ class AuthEnvVar(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("optional", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: False) - - @field_validator("label", mode="wrap") - @classmethod - def _salvage_on_error_1(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) - - @field_validator("secret", mode="wrap") + @field_validator("label", "optional", "secret", mode="wrap") @classmethod - def _salvage_on_error_2(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: True) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class AuthMethodEnvVar(BaseModel): @@ -1020,13 +959,13 @@ class AuthMethodEnvVar(BaseModel): @field_validator("description", "link", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("vars", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class AuthMethodTerminal(BaseModel): @@ -1061,13 +1000,13 @@ class AuthMethodTerminal(BaseModel): @field_validator("description", "env", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("args", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class AuthMethodAgent(BaseModel): @@ -1094,8 +1033,8 @@ class AuthMethodAgent(BaseModel): @field_validator("description", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class Implementation(BaseModel): @@ -1127,8 +1066,8 @@ class Implementation(BaseModel): @field_validator("title", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class AuthenticateResponse(BaseModel): @@ -1144,14 +1083,7 @@ class AuthenticateResponse(BaseModel): class ProviderCurrentConfig(BaseModel): api_type: Annotated[ - Union[ - Literal["anthropic"], - Literal["openai"], - Literal["azure"], - Literal["vertex"], - Literal["bedrock"], - Dict[str, Any], - ], + Union[Literal["anthropic"], Literal["openai"], Literal["azure"], Literal["vertex"], Literal["bedrock"], str], Field(alias="apiType"), ] """ @@ -1228,8 +1160,8 @@ class SessionMode(BaseModel): @field_validator("description", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class SessionConfigSelectOption(BaseModel): @@ -1256,8 +1188,8 @@ class SessionConfigSelectOption(BaseModel): @field_validator("description", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class SessionConfigBoolean(BaseModel): @@ -1303,13 +1235,13 @@ class SessionInfo(BaseModel): @field_validator("title", "updated_at", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("additional_directories", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class DeleteSessionResponse(BaseModel): @@ -1381,8 +1313,8 @@ class Usage(BaseModel): @field_validator("cached_read_tokens", "cached_write_tokens", "thought_tokens", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class StartNesResponse(BaseModel): @@ -1579,7 +1511,7 @@ class UnstructuredCommandInput(BaseModel): """ -class _CurrentModeUpdate(BaseModel): +class CurrentModeUpdateBase(BaseModel): current_mode_id: Annotated[str, Field(alias="currentModeId")] """ The ID of the current mode @@ -1594,7 +1526,7 @@ class _CurrentModeUpdate(BaseModel): """ -class _SessionInfoUpdate(BaseModel): +class SessionInfoUpdateBase(BaseModel): title: Optional[str] = None """ Human-readable title for the session. Set to null to clear. @@ -1612,11 +1544,6 @@ class _SessionInfoUpdate(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("title", "updated_at", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) - class Cost(BaseModel): amount: float @@ -1637,7 +1564,7 @@ class Cost(BaseModel): """ -class _UsageUpdate(BaseModel): +class UsageUpdateBase(BaseModel): used: Annotated[int, Field(ge=0)] """ Tokens currently in context. @@ -1659,11 +1586,6 @@ class _UsageUpdate(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("cost", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) - class CompleteElicitationNotification(BaseModel): elicitation_id: Annotated[str, Field(alias="elicitationId")] @@ -1706,8 +1628,8 @@ class MessageMcpNotification(BaseModel): @field_validator("params", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class FileSystemCapabilities(BaseModel): @@ -1730,8 +1652,8 @@ class FileSystemCapabilities(BaseModel): @field_validator("read_text_file", "write_text_file", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: False) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class BooleanConfigOptionCapabilities(BaseModel): @@ -1774,8 +1696,8 @@ class AuthCapabilities(BaseModel): @field_validator("terminal", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: False) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ElicitationFormCapabilities(BaseModel): @@ -1866,14 +1788,7 @@ class SetProviderRequest(BaseModel): Provider ID to configure. """ api_type: Annotated[ - Union[ - Literal["anthropic"], - Literal["openai"], - Literal["azure"], - Literal["vertex"], - Literal["bedrock"], - Dict[str, Any], - ], + Union[Literal["anthropic"], Literal["openai"], Literal["azure"], Literal["vertex"], Literal["bedrock"], str], Field(alias="apiType"), ] """ @@ -2127,7 +2042,7 @@ class SetSessionConfigOptionBooleanRequest(BaseModel): """ The boolean value. """ - type: Literal["boolean"] + type: Literal["boolean"] = "boolean" class SetSessionConfigOptionSelectRequest(BaseModel): @@ -2329,7 +2244,7 @@ class ReadTextFileResponse(BaseModel): class DeniedOutcome(BaseModel): - outcome: Literal["cancelled"] + outcome: Literal["cancelled"] = "cancelled" class SelectedPermissionOutcome(BaseModel): @@ -2382,8 +2297,8 @@ class TerminalExitStatus(BaseModel): @field_validator("exit_code", "signal", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ReleaseTerminalResponse(BaseModel): @@ -2417,8 +2332,8 @@ class WaitForTerminalExitResponse(BaseModel): @field_validator("exit_code", "signal", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class KillTerminalResponse(BaseModel): @@ -2441,7 +2356,7 @@ class DeclineElicitationResponse(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - action: Literal["decline"] + action: Literal["decline"] = "decline" class CancelElicitationResponse(BaseModel): @@ -2453,7 +2368,7 @@ class CancelElicitationResponse(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - action: Literal["cancel"] + action: Literal["cancel"] = "cancel" class OtherElicitationResponse(BaseModel): @@ -2477,20 +2392,8 @@ class OtherElicitationResponse(BaseModel): future ACP variants. """ - @field_validator("action", mode="before") - @classmethod - def _reject_known_action(cls, value: Any) -> Any: - # Restore the schema's `not` clause dropped for codegen: reject the known - # variants' discriminator values so a malformed known variant fails instead - # of silently parsing as this catch-all. - if value in ("accept", "decline", "cancel"): - raise ValueError("action value is reserved by a known variant") - return value - class ElicitationContentValue(RootModel[Union[str, int, float, bool, List[str]]]): - model_config = ConfigDict(use_attribute_docstrings=True) - root: Union[str, int, float, bool, List[str]] """ Allowed wire representations for [`ElicitationContentValue`]. @@ -2672,15 +2575,15 @@ class WriteTextFileRequest(BaseModel): class FileEditToolCallContent(Diff): - type: Literal["diff"] + type: Literal["diff"] = "diff" class TerminalToolCallContent(Terminal): - type: Literal["terminal"] + type: Literal["terminal"] = "terminal" class Annotations(BaseModel): - audience: Optional[List[str]] = None + audience: Optional[List[Literal["assistant", "user"]]] = None """ Intended recipients for this content, such as the user or assistant. """ @@ -2703,13 +2606,13 @@ class Annotations(BaseModel): @field_validator("last_modified", "priority", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("audience", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class TextContent(BaseModel): @@ -2732,8 +2635,8 @@ class TextContent(BaseModel): @field_validator("annotations", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ImageContent(BaseModel): @@ -2764,8 +2667,8 @@ class ImageContent(BaseModel): @field_validator("annotations", "uri", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class AudioContent(BaseModel): @@ -2792,8 +2695,8 @@ class AudioContent(BaseModel): @field_validator("annotations", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ResourceLink(BaseModel): @@ -2836,8 +2739,8 @@ class ResourceLink(BaseModel): @field_validator("annotations", "description", "mime_type", "size", "title", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class EmbeddedResource(BaseModel): @@ -2860,8 +2763,8 @@ class EmbeddedResource(BaseModel): @field_validator("annotations", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class PermissionOption(BaseModel): @@ -2873,7 +2776,7 @@ class PermissionOption(BaseModel): """ Human-readable label to display to the user. """ - kind: PermissionOptionKind + kind: Literal["allow_once", "allow_always", "reject_once", "reject_always"] """ Hint about the nature of this permission option. """ @@ -2930,21 +2833,38 @@ class CreateTerminalRequest(BaseModel): @field_validator("cwd", "output_byte_limit", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) - @field_validator("args", mode="wrap") + @field_validator("args", "env", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) - @field_validator("env", mode="wrap") - @classmethod - def _skip_invalid_items_1(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + +class CreateUrlSessionElicitationRequestBase(ElicitationSessionScope): + elicitation_id: Annotated[str, Field(alias="elicitationId")] + """ + The unique identifier for this elicitation. + """ + url: AnyUrl + """ + The URL to direct the user to. + """ + + +class CreateUrlRequestElicitationRequestBase(ElicitationRequestScope): + elicitation_id: Annotated[str, Field(alias="elicitationId")] + """ + The unique identifier for this elicitation. + """ + url: AnyUrl + """ + The URL to direct the user to. + """ -class CreateUrlSessionElicitationRequest(ElicitationSessionScope): +class CreateUrlSessionElicitationRequest(CreateUrlSessionElicitationRequestBase, CreateUrlElicitationRequestBase): message: str """ A human-readable message describing what input is needed. @@ -2957,18 +2877,29 @@ class CreateUrlSessionElicitationRequest(ElicitationSessionScope): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - mode: Literal["url"] - elicitation_id: Annotated[str, Field(alias="elicitationId")] + mode: Literal["url"] = "url" + + +class CreateUrlRequestElicitationRequest(CreateUrlRequestElicitationRequestBase, CreateUrlElicitationRequestBase): + message: str """ - The unique identifier for this elicitation. + A human-readable message describing what input is needed. """ - url: AnyUrl + field_meta: Annotated[Optional[Dict[str, Any]], Field(alias="_meta")] = None """ - The URL to direct the user to. + The _meta property is reserved by ACP to allow clients and agents to attach additional + metadata to their interactions. Implementations MUST NOT make assumptions about values at + these keys. + + See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ + mode: Literal["url"] = "url" -class CreateUrlRequestElicitationRequest(ElicitationRequestScope): +class CreateOtherSessionElicitationRequest(ElicitationSessionScope): + model_config = ConfigDict( + extra="allow", + ) message: str """ A human-readable message describing what input is needed. @@ -2981,35 +2912,60 @@ class CreateUrlRequestElicitationRequest(ElicitationRequestScope): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - mode: Literal["url"] - elicitation_id: Annotated[str, Field(alias="elicitationId")] + mode: str """ - The unique identifier for this elicitation. + Custom or future elicitation mode. + + Values beginning with `_` are reserved for implementation-specific + extensions. Unknown values that do not begin with `_` are reserved for + future ACP variants. """ - url: AnyUrl + + +class CreateOtherRequestElicitationRequest(ElicitationRequestScope): + model_config = ConfigDict( + extra="allow", + ) + message: str """ - The URL to direct the user to. + A human-readable message describing what input is needed. + """ + field_meta: Annotated[Optional[Dict[str, Any]], Field(alias="_meta")] = None + """ + The _meta property is reserved by ACP to allow clients and agents to attach additional + metadata to their interactions. Implementations MUST NOT make assumptions about values at + these keys. + + See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) + """ + mode: str + """ + Custom or future elicitation mode. + + Values beginning with `_` are reserved for implementation-specific + extensions. Unknown values that do not begin with `_` are reserved for + future ACP variants. """ class ElicitationStringPropertySchema(StringPropertySchema): - type: Literal["string"] + type: Literal["string"] = "string" class ElicitationNumberPropertySchema(NumberPropertySchema): - type: Literal["number"] + type: Literal["number"] = "number" class ElicitationIntegerPropertySchema(IntegerPropertySchema): - type: Literal["integer"] + type: Literal["integer"] = "integer" class ElicitationBooleanPropertySchema(BooleanPropertySchema): - type: Literal["boolean"] + type: Literal["boolean"] = "boolean" -class StringMultiSelectItems(_StringMultiSelectItems): - type: Literal["string"] +class StringMultiSelectItems(StringMultiSelectItemsBase): + type: Literal["string"] = "string" class MultiSelectPropertySchema(BaseModel): @@ -3048,13 +3004,13 @@ class MultiSelectPropertySchema(BaseModel): @field_validator("description", "title", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("default", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class ConnectMcpRequest(BaseModel): @@ -3113,8 +3069,7 @@ class SessionCapabilities(BaseModel): Supplying `{}` means the agent supports deleting sessions from `session/list`. """ additional_directories: Annotated[ - Optional[SessionAdditionalDirectoriesCapabilities], - Field(alias="additionalDirectories"), + Optional[SessionAdditionalDirectoriesCapabilities], Field(alias="additionalDirectories") ] = None """ Whether the agent supports `additionalDirectories` on supported session lifecycle requests. @@ -3163,8 +3118,8 @@ class SessionCapabilities(BaseModel): @field_validator("additional_directories", "close", "delete", "fork", "list", "resume", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class AgentAuthCapabilities(BaseModel): @@ -3186,12 +3141,12 @@ class AgentAuthCapabilities(BaseModel): @field_validator("logout", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NesDocumentDidChangeCapabilities(BaseModel): - sync_kind: Annotated[str, Field(alias="syncKind")] + sync_kind: Annotated[Literal["full", "incremental"], Field(alias="syncKind")] """ The sync kind the agent wants: `"full"` or `"incremental"`. """ @@ -3243,16 +3198,16 @@ class NesContextCapabilities(BaseModel): "diagnostics", "edit_history", "open_files", "recent_files", "related_snippets", "user_actions", mode="wrap" ) @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class EnvVarAuthMethod(AuthMethodEnvVar): - type: Literal["env_var"] + type: Literal["env_var"] = "env_var" class TerminalAuthMethod(AuthMethodTerminal): - type: Literal["terminal"] + type: Literal["terminal"] = "terminal" class ProviderInfo(BaseModel): @@ -3261,14 +3216,7 @@ class ProviderInfo(BaseModel): Provider identifier, for example "main" or "openai". """ supported: List[ - Union[ - Literal["anthropic"], - Literal["openai"], - Literal["azure"], - Literal["vertex"], - Literal["bedrock"], - Dict[str, Any], - ] + Union[Literal["anthropic"], Literal["openai"], Literal["azure"], Literal["vertex"], Literal["bedrock"], str] ] """ Supported protocol types for this provider. @@ -3294,8 +3242,8 @@ class ProviderInfo(BaseModel): @field_validator("supported", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class SessionModeState(BaseModel): @@ -3318,8 +3266,8 @@ class SessionModeState(BaseModel): @field_validator("available_modes", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class SessionConfigOptionBoolean(SessionConfigBoolean): @@ -3336,13 +3284,7 @@ class SessionConfigOptionBoolean(SessionConfigBoolean): Optional description for the Client to display to the user. """ category: Optional[ - Union[ - Literal["mode"], - Literal["model"], - Literal["model_config"], - Literal["thought_level"], - Dict[str, Any], - ] + Union[Literal["mode"], Literal["model"], Literal["model_config"], Literal["thought_level"], str] ] = None """ Optional semantic category for this option (UX only). @@ -3355,12 +3297,7 @@ class SessionConfigOptionBoolean(SessionConfigBoolean): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - type: Literal["boolean"] - - @field_validator("category", "description", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + type: Literal["boolean"] = "boolean" class SessionConfigSelectGroup(BaseModel): @@ -3387,8 +3324,8 @@ class SessionConfigSelectGroup(BaseModel): @field_validator("options", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class ListSessionsResponse(BaseModel): @@ -3412,17 +3349,19 @@ class ListSessionsResponse(BaseModel): @field_validator("next_cursor", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("sessions", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class PromptResponse(BaseModel): - stop_reason: Annotated[StopReason, Field(alias="stopReason")] + stop_reason: Annotated[ + Literal["end_turn", "max_tokens", "max_turn_requests", "refusal", "cancelled"], Field(alias="stopReason") + ] """ Indicates why the agent stopped processing the turn. """ @@ -3445,20 +3384,20 @@ class PromptResponse(BaseModel): @field_validator("usage", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NesJumpSuggestionVariant(NesJumpSuggestion): - kind: Literal["jump"] + kind: Literal["jump"] = "jump" class NesRenameSuggestionVariant(NesRenameSuggestion): - kind: Literal["rename"] + kind: Literal["rename"] = "rename" class NesSearchAndReplaceSuggestionVariant(NesSearchAndReplaceSuggestion): - kind: Literal["searchAndReplace"] + kind: Literal["searchAndReplace"] = "searchAndReplace" class Range(BaseModel): @@ -3509,24 +3448,34 @@ class Error(BaseModel): @field_validator("data", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class AgentPlanRemovedUpdate(PlanRemoved): - session_update: Annotated[Literal["plan_removed"], Field(alias="sessionUpdate")] + session_update: Annotated[Literal["plan_removed"], Field(alias="sessionUpdate")] = "plan_removed" -class CurrentModeUpdate(_CurrentModeUpdate): - session_update: Annotated[Literal["current_mode_update"], Field(alias="sessionUpdate")] +class CurrentModeUpdate(CurrentModeUpdateBase): + session_update: Annotated[Literal["current_mode_update"], Field(alias="sessionUpdate")] = "current_mode_update" -class SessionInfoUpdate(_SessionInfoUpdate): - session_update: Annotated[Literal["session_info_update"], Field(alias="sessionUpdate")] +class SessionInfoUpdate(SessionInfoUpdateBase): + session_update: Annotated[Literal["session_info_update"], Field(alias="sessionUpdate")] = "session_info_update" + @field_validator("title", "updated_at", mode="wrap") + @classmethod + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) -class UsageUpdate(_UsageUpdate): - session_update: Annotated[Literal["usage_update"], Field(alias="sessionUpdate")] + +class UsageUpdate(UsageUpdateBase): + session_update: Annotated[Literal["usage_update"], Field(alias="sessionUpdate")] = "usage_update" + + @field_validator("cost", mode="wrap") + @classmethod + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class PlanEntry(BaseModel): @@ -3534,12 +3483,12 @@ class PlanEntry(BaseModel): """ Human-readable description of what this task aims to accomplish. """ - priority: PlanEntryPriority + priority: Literal["high", "medium", "low"] """ The relative importance of this task. Used to indicate which tasks are most critical to the overall goal. """ - status: PlanEntryStatus + status: Literal["pending", "in_progress", "completed"] """ Current execution status of this task. """ @@ -3572,16 +3521,16 @@ class Plan(BaseModel): @field_validator("entries", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class PlanUpdateFile(PlanFile): - type: Literal["file"] + type: Literal["file"] = "file" class PlanUpdateMarkdown(PlanMarkdown): - type: Literal["markdown"] + type: Literal["markdown"] = "markdown" class PlanItems(BaseModel): @@ -3607,13 +3556,11 @@ class PlanItems(BaseModel): @field_validator("entries", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class AvailableCommandInput(RootModel[UnstructuredCommandInput]): - model_config = ConfigDict(use_attribute_docstrings=True) - root: UnstructuredCommandInput """ The input specification for a command. @@ -3641,8 +3588,8 @@ class SessionConfigOptionsCapabilities(BaseModel): @field_validator("boolean", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ElicitationCapabilities(BaseModel): @@ -3671,8 +3618,8 @@ class ElicitationCapabilities(BaseModel): @field_validator("form", "url", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ClientNesCapabilities(BaseModel): @@ -3699,26 +3646,25 @@ class ClientNesCapabilities(BaseModel): @field_validator("jump", "rename", "search_and_replace", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class HttpMcpServer(McpServerHttp): - type: Literal["http"] + type: Literal["http"] = "http" class SseMcpServer(McpServerSse): - type: Literal["sse"] + type: Literal["sse"] = "sse" class AcpMcpServer(McpServerAcp): - type: Literal["acp"] + type: Literal["acp"] = "acp" class LoadSessionRequest(BaseModel): mcp_servers: Annotated[ - List[Union[HttpMcpServer, SseMcpServer, AcpMcpServer, McpServerStdio]], - Field(alias="mcpServers"), + List[Union[HttpMcpServer, SseMcpServer, AcpMcpServer, McpServerStdio]], Field(alias="mcpServers") ] """ List of MCP servers to connect to for this session. @@ -3749,15 +3695,10 @@ class LoadSessionRequest(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("additional_directories", mode="wrap") - @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) - - @field_validator("mcp_servers", mode="wrap") + @field_validator("additional_directories", "mcp_servers", mode="wrap") @classmethod - def _skip_invalid_items_1(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class ForkSessionRequest(BaseModel): @@ -3778,8 +3719,7 @@ class ForkSessionRequest(BaseModel): session. """ mcp_servers: Annotated[ - Optional[List[Union[HttpMcpServer, SseMcpServer, AcpMcpServer, McpServerStdio]]], - Field(alias="mcpServers"), + Optional[List[Union[HttpMcpServer, SseMcpServer, AcpMcpServer, McpServerStdio]]], Field(alias="mcpServers") ] = None """ List of MCP servers to connect to for this session. @@ -3793,15 +3733,10 @@ class ForkSessionRequest(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("additional_directories", mode="wrap") - @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) - - @field_validator("mcp_servers", mode="wrap") + @field_validator("additional_directories", "mcp_servers", mode="wrap") @classmethod - def _skip_invalid_items_1(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class ResumeSessionRequest(BaseModel): @@ -3823,8 +3758,7 @@ class ResumeSessionRequest(BaseModel): the request `cwd` matches the session's `cwd`. """ mcp_servers: Annotated[ - Optional[List[Union[HttpMcpServer, SseMcpServer, AcpMcpServer, McpServerStdio]]], - Field(alias="mcpServers"), + Optional[List[Union[HttpMcpServer, SseMcpServer, AcpMcpServer, McpServerStdio]]], Field(alias="mcpServers") ] = None """ List of MCP servers to connect to for this session. @@ -3838,15 +3772,10 @@ class ResumeSessionRequest(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("additional_directories", mode="wrap") - @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) - - @field_validator("mcp_servers", mode="wrap") + @field_validator("additional_directories", "mcp_servers", mode="wrap") @classmethod - def _skip_invalid_items_1(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class StartNesRequest(BaseModel): @@ -3873,8 +3802,8 @@ class StartNesRequest(BaseModel): @field_validator("repository", "workspace_uri", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NesRelatedSnippet(BaseModel): @@ -3924,8 +3853,8 @@ class NesOpenFile(BaseModel): @field_validator("last_focused_ms", "visible_range", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NesDiagnostic(BaseModel): @@ -3937,7 +3866,7 @@ class NesDiagnostic(BaseModel): """ The range of the diagnostic. """ - severity: str + severity: Literal["error", "warning", "information", "hint"] """ The severity of the diagnostic. """ @@ -3967,7 +3896,7 @@ class ClientErrorMessage(BaseModel): class AllowedOutcome(SelectedPermissionOutcome): - outcome: Literal["selected"] + outcome: Literal["selected"] = "selected" class TerminalOutputResponse(BaseModel): @@ -3994,8 +3923,8 @@ class TerminalOutputResponse(BaseModel): @field_validator("exit_status", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class AcceptElicitationResponse(ElicitationAcceptAction): @@ -4007,7 +3936,7 @@ class AcceptElicitationResponse(ElicitationAcceptAction): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - action: Literal["accept"] + action: Literal["accept"] = "accept" class TextDocumentContentChangeEvent(BaseModel): @@ -4069,7 +3998,7 @@ class RejectNesNotification(BaseModel): """ The ID of the rejected suggestion. """ - reason: Optional[str] = None + reason: Optional[Literal["rejected", "ignored", "replaced", "cancelled"]] = None """ The reason for rejection. """ @@ -4084,28 +4013,28 @@ class RejectNesNotification(BaseModel): @field_validator("reason", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class TextContentBlock(TextContent): - type: Literal["text"] + type: Literal["text"] = "text" class ImageContentBlock(ImageContent): - type: Literal["image"] + type: Literal["image"] = "image" class AudioContentBlock(AudioContent): - type: Literal["audio"] + type: Literal["audio"] = "audio" class ResourceContentBlock(ResourceLink): - type: Literal["resource_link"] + type: Literal["resource_link"] = "resource_link" class EmbeddedResourceContentBlock(EmbeddedResource): - type: Literal["resource"] + type: Literal["resource"] = "resource" class Content(BaseModel): @@ -4129,7 +4058,7 @@ class Content(BaseModel): class ElicitationMultiSelectPropertySchema(MultiSelectPropertySchema): - type: Literal["array"] + type: Literal["array"] = "array" class AgentErrorMessage(BaseModel): @@ -4175,8 +4104,8 @@ class NesDocumentEventCapabilities(BaseModel): @field_validator("did_change", "did_close", "did_focus", "did_open", "did_save", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ListProvidersResponse(BaseModel): @@ -4252,12 +4181,12 @@ class NesEditSuggestion(BaseModel): @field_validator("cursor_position", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class AgentPlanUpdate(Plan): - session_update: Annotated[Literal["plan"], Field(alias="sessionUpdate")] + session_update: Annotated[Literal["plan"], Field(alias="sessionUpdate")] = "plan" class ContentChunk(BaseModel): @@ -4288,19 +4217,16 @@ class ContentChunk(BaseModel): @field_validator("message_id", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class PlanUpdateItems(PlanItems): - type: Literal["items"] + type: Literal["items"] = "items" class PlanUpdate(BaseModel): - plan: Annotated[ - Union[PlanUpdateItems, PlanUpdateFile, PlanUpdateMarkdown], - Field(discriminator="type"), - ] + plan: Annotated[Union[PlanUpdateItems, PlanUpdateFile, PlanUpdateMarkdown], Field(discriminator="type")] """ The updated plan content. """ @@ -4338,11 +4264,11 @@ class AvailableCommand(BaseModel): @field_validator("input", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) -class _AvailableCommandsUpdate(BaseModel): +class AvailableCommandsUpdateBase(BaseModel): available_commands: Annotated[List[AvailableCommand], Field(alias="availableCommands")] """ Commands the agent can execute @@ -4356,11 +4282,6 @@ class _AvailableCommandsUpdate(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("available_commands", mode="wrap") - @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) - class ClientSessionCapabilities(BaseModel): config_options: Annotated[Optional[SessionConfigOptionsCapabilities], Field(alias="configOptions")] = None @@ -4381,8 +4302,8 @@ class ClientSessionCapabilities(BaseModel): @field_validator("config_options", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NewSessionRequest(BaseModel): @@ -4399,8 +4320,7 @@ class NewSessionRequest(BaseModel): additional roots are activated for the new session. """ mcp_servers: Annotated[ - List[Union[HttpMcpServer, SseMcpServer, AcpMcpServer, McpServerStdio]], - Field(alias="mcpServers"), + List[Union[HttpMcpServer, SseMcpServer, AcpMcpServer, McpServerStdio]], Field(alias="mcpServers") ] """ List of MCP (Model Context Protocol) servers the agent should connect to. @@ -4414,15 +4334,10 @@ class NewSessionRequest(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("additional_directories", mode="wrap") - @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) - - @field_validator("mcp_servers", mode="wrap") + @field_validator("additional_directories", "mcp_servers", mode="wrap") @classmethod - def _skip_invalid_items_1(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class PromptRequest(BaseModel): @@ -4503,10 +4418,7 @@ class NesSuggestContext(BaseModel): class RequestPermissionResponse(BaseModel): - outcome: Annotated[ - Union[DeniedOutcome, AllowedOutcome], - Field(discriminator="outcome"), - ] + outcome: Annotated[Union[DeniedOutcome, AllowedOutcome], Field(discriminator="outcome")] """ The user's decision on the permission request. """ @@ -4548,16 +4460,16 @@ class DidChangeDocumentNotification(BaseModel): @field_validator("content_changes", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class ContentToolCallContent(Content): - type: Literal["content"] + type: Literal["content"] = "content" class ElicitationSchema(BaseModel): - type: Optional[str] = "object" + type: Optional[Literal["object"]] = "object" """ Type discriminator. Always `"object"`. """ @@ -4601,15 +4513,10 @@ class ElicitationSchema(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("type", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: "object") - - @field_validator("description", "title", mode="wrap") + @field_validator("description", "title", "type", mode="wrap") @classmethod - def _salvage_on_error_1(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class ElicitationFormSessionMode(ElicitationSessionScope): @@ -4627,8 +4534,6 @@ class ElicitationFormRequestMode(ElicitationRequestScope): class ElicitationFormMode(RootModel[Union[ElicitationFormSessionMode, ElicitationFormRequestMode]]): - model_config = ConfigDict(use_attribute_docstrings=True) - root: Union[ElicitationFormSessionMode, ElicitationFormRequestMode] """ **UNSTABLE** @@ -4655,8 +4560,8 @@ class NesEventCapabilities(BaseModel): @field_validator("document", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class SessionConfigOptionSelect(SessionConfigSelect): @@ -4673,13 +4578,7 @@ class SessionConfigOptionSelect(SessionConfigSelect): Optional description for the Client to display to the user. """ category: Optional[ - Union[ - Literal["mode"], - Literal["model"], - Literal["model_config"], - Literal["thought_level"], - Dict[str, Any], - ] + Union[Literal["mode"], Literal["model"], Literal["model_config"], Literal["thought_level"], str] ] = None """ Optional semantic category for this option (UX only). @@ -4692,12 +4591,7 @@ class SessionConfigOptionSelect(SessionConfigSelect): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - type: Literal["select"] - - @field_validator("category", "description", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + type: Literal["select"] = "select" class LoadSessionResponse(BaseModel): @@ -4709,12 +4603,7 @@ class LoadSessionResponse(BaseModel): """ config_options: Annotated[ Optional[ - List[ - Annotated[ - Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], - Field(discriminator="type"), - ] - ] + List[Annotated[Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], Field(discriminator="type")]] ], Field(alias="configOptions"), ] = None @@ -4732,13 +4621,13 @@ class LoadSessionResponse(BaseModel): @field_validator("modes", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("config_options", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class ForkSessionResponse(BaseModel): @@ -4754,12 +4643,7 @@ class ForkSessionResponse(BaseModel): """ config_options: Annotated[ Optional[ - List[ - Annotated[ - Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], - Field(discriminator="type"), - ] - ] + List[Annotated[Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], Field(discriminator="type")]] ], Field(alias="configOptions"), ] = None @@ -4777,13 +4661,13 @@ class ForkSessionResponse(BaseModel): @field_validator("modes", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("config_options", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class ResumeSessionResponse(BaseModel): @@ -4795,12 +4679,7 @@ class ResumeSessionResponse(BaseModel): """ config_options: Annotated[ Optional[ - List[ - Annotated[ - Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], - Field(discriminator="type"), - ] - ] + List[Annotated[Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], Field(discriminator="type")]] ], Field(alias="configOptions"), ] = None @@ -4818,23 +4697,18 @@ class ResumeSessionResponse(BaseModel): @field_validator("modes", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("config_options", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class SetSessionConfigOptionResponse(BaseModel): config_options: Annotated[ - List[ - Annotated[ - Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], - Field(discriminator="type"), - ] - ], + List[Annotated[Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], Field(discriminator="type")]], Field(alias="configOptions"), ] """ @@ -4851,32 +4725,39 @@ class SetSessionConfigOptionResponse(BaseModel): @field_validator("config_options", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class NesEditSuggestionVariant(NesEditSuggestion): - kind: Literal["edit"] + kind: Literal["edit"] = "edit" class UserMessageChunk(ContentChunk): - session_update: Annotated[Literal["user_message_chunk"], Field(alias="sessionUpdate")] + session_update: Annotated[Literal["user_message_chunk"], Field(alias="sessionUpdate")] = "user_message_chunk" class AgentMessageChunk(ContentChunk): - session_update: Annotated[Literal["agent_message_chunk"], Field(alias="sessionUpdate")] + session_update: Annotated[Literal["agent_message_chunk"], Field(alias="sessionUpdate")] = "agent_message_chunk" class AgentThoughtChunk(ContentChunk): - session_update: Annotated[Literal["agent_thought_chunk"], Field(alias="sessionUpdate")] + session_update: Annotated[Literal["agent_thought_chunk"], Field(alias="sessionUpdate")] = "agent_thought_chunk" class AgentPlanContentUpdate(PlanUpdate): - session_update: Annotated[Literal["plan_update"], Field(alias="sessionUpdate")] + session_update: Annotated[Literal["plan_update"], Field(alias="sessionUpdate")] = "plan_update" + +class AvailableCommandsUpdate(AvailableCommandsUpdateBase): + session_update: Annotated[Literal["available_commands_update"], Field(alias="sessionUpdate")] = ( + "available_commands_update" + ) -class AvailableCommandsUpdate(_AvailableCommandsUpdate): - session_update: Annotated[Literal["available_commands_update"], Field(alias="sessionUpdate")] + @field_validator("available_commands", mode="wrap") + @classmethod + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class ToolCall(BaseModel): @@ -4888,12 +4769,14 @@ class ToolCall(BaseModel): """ Human-readable title describing what the tool is doing. """ - kind: Optional[ToolKind] = None + kind: Optional[ + Literal["read", "edit", "delete", "move", "search", "execute", "think", "fetch", "switch_mode", "other"] + ] = None """ The category of tool being invoked. Helps clients choose appropriate icons and UI treatment. """ - status: Optional[ToolCallStatus] = None + status: Optional[Literal["pending", "in_progress", "completed", "failed"]] = None """ Current execution status of the tool call. """ @@ -4932,28 +4815,18 @@ class ToolCall(BaseModel): @field_validator("kind", "raw_input", "raw_output", "status", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) - - @field_validator("content", mode="wrap") - @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) - @field_validator("locations", mode="wrap") + @field_validator("content", "locations", mode="wrap") @classmethod - def _skip_invalid_items_1(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) -class _ConfigOptionUpdate(BaseModel): +class ConfigOptionUpdateBase(BaseModel): config_options: Annotated[ - List[ - Annotated[ - Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], - Field(discriminator="type"), - ] - ], + List[Annotated[Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], Field(discriminator="type")]], Field(alias="configOptions"), ] """ @@ -4968,14 +4841,12 @@ class _ConfigOptionUpdate(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("config_options", mode="wrap") - @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) - class ClientCapabilities(BaseModel): - fs: Annotated[Optional[FileSystemCapabilities], Field(validate_default=True)] = FileSystemCapabilities() + fs: Annotated[Optional[FileSystemCapabilities], Field(validate_default=True)] = { + "readTextFile": False, + "writeTextFile": False, + } """ File system capabilities supported by the client. Determines which file operations the agent can request. @@ -5035,7 +4906,9 @@ class ClientCapabilities(BaseModel): Optional. Omitted or `null` both mean the client does not advertise any NES suggestion-kind extensions. """ - position_encodings: Annotated[Optional[List[str]], Field(alias="positionEncodings")] = None + position_encodings: Annotated[ + Optional[List[Literal["utf-16", "utf-32", "utf-8"]]], Field(alias="positionEncodings") + ] = None """ **UNSTABLE** @@ -5052,30 +4925,15 @@ class ClientCapabilities(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("terminal", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: False) - - @field_validator("elicitation", "nes", "plan", "session", mode="wrap") - @classmethod - def _salvage_on_error_1(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) - - @field_validator("fs", mode="wrap") + @field_validator("auth", "elicitation", "fs", "nes", "plan", "session", "terminal", mode="wrap") @classmethod - def _salvage_on_error_2(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: {"readTextFile": False, "writeTextFile": False}) - - @field_validator("auth", mode="wrap") - @classmethod - def _salvage_on_error_3(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: {"terminal": False}) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("position_encodings", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class SuggestNesRequest(BaseModel): @@ -5099,7 +4957,7 @@ class SuggestNesRequest(BaseModel): """ The current text selection range, if any. """ - trigger_kind: Annotated[str, Field(alias="triggerKind")] + trigger_kind: Annotated[Literal["automatic", "diagnostic", "manual"], Field(alias="triggerKind")] """ What triggered this suggestion request. """ @@ -5134,10 +4992,7 @@ class ClientResponseMessage(BaseModel): ConnectMcpResponse, DisconnectMcpResponse, Union[ - AcceptElicitationResponse, - DeclineElicitationResponse, - CancelElicitationResponse, - OtherElicitationResponse, + AcceptElicitationResponse, DeclineElicitationResponse, CancelElicitationResponse, OtherElicitationResponse ], Any, ] @@ -5147,8 +5002,6 @@ class ClientResponseMessage(BaseModel): class ClientResponse(RootModel[Union[ClientResponseMessage, ClientErrorMessage]]): - model_config = ConfigDict(use_attribute_docstrings=True) - root: Union[ClientResponseMessage, ClientErrorMessage] """ A JSON-RPC response object. @@ -5184,11 +5037,13 @@ class ToolCallUpdate(BaseModel): """ The ID of the tool call being updated. """ - kind: Optional[ToolKind] = None + kind: Optional[ + Literal["read", "edit", "delete", "move", "search", "execute", "think", "fetch", "switch_mode", "other"] + ] = None """ Update the tool kind. """ - status: Optional[ToolCallStatus] = None + status: Optional[Literal["pending", "in_progress", "completed", "failed"]] = None """ Update the execution status. """ @@ -5230,21 +5085,30 @@ class ToolCallUpdate(BaseModel): @field_validator("kind", "raw_input", "raw_output", "status", "title", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) - @field_validator("content", mode="wrap") + @field_validator("content", "locations", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) - @field_validator("locations", mode="wrap") - @classmethod - def _skip_invalid_items_1(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + +class CreateFormSessionElicitationRequestBase(ElicitationSessionScope): + requested_schema: Annotated[ElicitationSchema, Field(alias="requestedSchema")] + """ + A JSON Schema describing the form fields to present to the user. + """ -class CreateFormSessionElicitationRequest(ElicitationSessionScope): +class CreateFormRequestElicitationRequestBase(ElicitationRequestScope): + requested_schema: Annotated[ElicitationSchema, Field(alias="requestedSchema")] + """ + A JSON Schema describing the form fields to present to the user. + """ + + +class CreateFormSessionElicitationRequest(CreateFormSessionElicitationRequestBase, CreateFormElicitationRequestBase): message: str """ A human-readable message describing what input is needed. @@ -5257,14 +5121,10 @@ class CreateFormSessionElicitationRequest(ElicitationSessionScope): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - mode: Literal["form"] - requested_schema: Annotated[ElicitationSchema, Field(alias="requestedSchema")] - """ - A JSON Schema describing the form fields to present to the user. - """ + mode: Literal["form"] = "form" -class CreateFormRequestElicitationRequest(ElicitationRequestScope): +class CreateFormRequestElicitationRequest(CreateFormRequestElicitationRequestBase, CreateFormElicitationRequestBase): message: str """ A human-readable message describing what input is needed. @@ -5277,38 +5137,7 @@ class CreateFormRequestElicitationRequest(ElicitationRequestScope): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - mode: Literal["form"] - requested_schema: Annotated[ElicitationSchema, Field(alias="requestedSchema")] - """ - A JSON Schema describing the form fields to present to the user. - """ - - -ElicitationMode = Union[ - ElicitationFormSessionMode, - ElicitationFormRequestMode, - ElicitationUrlSessionMode, - ElicitationUrlRequestMode, -] -CreateFormElicitationRequest = Union[ - CreateFormSessionElicitationRequest, - CreateFormRequestElicitationRequest, -] -CreateUrlElicitationRequest = Union[ - CreateUrlSessionElicitationRequest, - CreateUrlRequestElicitationRequest, -] -CreateElicitationRequest = Union[ - CreateFormElicitationRequest, - CreateUrlElicitationRequest, - CreateOtherElicitationRequest, -] -CreateElicitationResponse = Union[ - AcceptElicitationResponse, - DeclineElicitationResponse, - CancelElicitationResponse, - OtherElicitationResponse, -] + mode: Literal["form"] = "form" class NesCapabilities(BaseModel): @@ -5331,8 +5160,8 @@ class NesCapabilities(BaseModel): @field_validator("context", "events", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class NewSessionResponse(BaseModel): @@ -5350,12 +5179,7 @@ class NewSessionResponse(BaseModel): """ config_options: Annotated[ Optional[ - List[ - Annotated[ - Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], - Field(discriminator="type"), - ] - ] + List[Annotated[Union[SessionConfigOptionSelect, SessionConfigOptionBoolean], Field(discriminator="type")]] ], Field(alias="configOptions"), ] = None @@ -5373,13 +5197,13 @@ class NewSessionResponse(BaseModel): @field_validator("modes", mode="wrap") @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("config_options", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class SuggestNesResponse(BaseModel): @@ -5408,15 +5232,20 @@ class SuggestNesResponse(BaseModel): class ToolCallStart(ToolCall): - session_update: Annotated[Literal["tool_call"], Field(alias="sessionUpdate")] + session_update: Annotated[Literal["tool_call"], Field(alias="sessionUpdate")] = "tool_call" class ToolCallProgress(ToolCallUpdate): - session_update: Annotated[Literal["tool_call_update"], Field(alias="sessionUpdate")] + session_update: Annotated[Literal["tool_call_update"], Field(alias="sessionUpdate")] = "tool_call_update" -class ConfigOptionUpdate(_ConfigOptionUpdate): - session_update: Annotated[Literal["config_option_update"], Field(alias="sessionUpdate")] +class ConfigOptionUpdate(ConfigOptionUpdateBase): + session_update: Annotated[Literal["config_option_update"], Field(alias="sessionUpdate")] = "config_option_update" + + @field_validator("config_options", mode="wrap") + @classmethod + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class InitializeRequest(BaseModel): @@ -5425,9 +5254,12 @@ class InitializeRequest(BaseModel): The latest protocol version supported by the client. """ client_capabilities: Annotated[ - Optional[ClientCapabilities], - Field(alias="clientCapabilities", validate_default=True), - ] = ClientCapabilities() + Optional[ClientCapabilities], Field(alias="clientCapabilities", validate_default=True) + ] = { + "fs": {"readTextFile": False, "writeTextFile": False}, + "terminal": False, + "auth": {"terminal": False}, + } """ Capabilities supported by the client. """ @@ -5448,35 +5280,13 @@ class InitializeRequest(BaseModel): @field_validator("protocol_version", mode="before") @classmethod - def _coerce_protocol_version(cls, value: Any) -> int: - # Some clients (e.g. Zed) send a date string like "2024-11-05" instead - # of an integer. The Rust SDK treats legacy strings as version 0; this - # SDK maps unparsable values to 1 so the connection is not rejected. - # See: https://github.com/agentclientprotocol/rust-sdk/blob/main/crates/agent-client-protocol-schema/src/version.rs - if isinstance(value, int): - return value - try: - return int(value) - except (TypeError, ValueError): - return 1 - - @field_validator("client_info", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) + def coerce_protocol_version_validator(cls, v: Any, info: ValidationInfo) -> Any: + return coerce_protocol_version(v, info) - @field_validator("client_capabilities", mode="wrap") + @field_validator("client_capabilities", "client_info", mode="wrap") @classmethod - def _salvage_on_error_1(cls, value: Any, handler: Any) -> Any: - return salvage_on_error( - value, - handler, - lambda: { - "fs": {"readTextFile": False, "writeTextFile": False}, - "terminal": False, - "auth": {"terminal": False}, - }, - ) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class RequestPermissionRequest(BaseModel): @@ -5508,22 +5318,26 @@ class AgentCapabilities(BaseModel): Whether the agent supports `session/load`. """ prompt_capabilities: Annotated[ - Optional[PromptCapabilities], - Field(alias="promptCapabilities", validate_default=True), - ] = PromptCapabilities() + Optional[PromptCapabilities], Field(alias="promptCapabilities", validate_default=True) + ] = { + "image": False, + "audio": False, + "embeddedContext": False, + } """ Prompt capabilities supported by the agent. """ - mcp_capabilities: Annotated[Optional[McpCapabilities], Field(alias="mcpCapabilities", validate_default=True)] = ( - McpCapabilities() - ) + mcp_capabilities: Annotated[Optional[McpCapabilities], Field(alias="mcpCapabilities", validate_default=True)] = { + "http": False, + "sse": False, + "acp": False, + } """ MCP capabilities supported by the agent. """ session_capabilities: Annotated[ - Optional[SessionCapabilities], - Field(alias="sessionCapabilities", validate_default=True), - ] = SessionCapabilities() + Optional[SessionCapabilities], Field(alias="sessionCapabilities", validate_default=True) + ] = {} """ Session lifecycle and prompt capabilities advertised by the agent. """ @@ -5553,7 +5367,7 @@ class AgentCapabilities(BaseModel): Optional. Omitted or `null` both mean the agent does not advertise support for NES methods. """ - position_encoding: Annotated[Optional[str], Field(alias="positionEncoding")] = None + position_encoding: Annotated[Optional[Literal["utf-16", "utf-32", "utf-8"]], Field(alias="positionEncoding")] = None """ **UNSTABLE** @@ -5570,30 +5384,20 @@ class AgentCapabilities(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("load_session", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: False) - - @field_validator("nes", "position_encoding", "providers", mode="wrap") - @classmethod - def _salvage_on_error_1(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) - - @field_validator("mcp_capabilities", mode="wrap") - @classmethod - def _salvage_on_error_2(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: {"http": False, "sse": False, "acp": False}) - - @field_validator("prompt_capabilities", mode="wrap") - @classmethod - def _salvage_on_error_3(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: {"image": False, "audio": False, "embeddedContext": False}) - - @field_validator("auth", "session_capabilities", mode="wrap") + @field_validator( + "auth", + "load_session", + "mcp_capabilities", + "nes", + "position_encoding", + "prompt_capabilities", + "providers", + "session_capabilities", + mode="wrap", + ) @classmethod - def _salvage_on_error_4(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: {}) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) class SessionNotification(BaseModel): @@ -5694,11 +5498,9 @@ class AgentRequest(BaseModel): MessageMcpRequest, DisconnectMcpRequest, Union[ - CreateFormSessionElicitationRequest, - CreateFormRequestElicitationRequest, - CreateUrlSessionElicitationRequest, - CreateUrlRequestElicitationRequest, - CreateOtherElicitationRequest, + Union[CreateOtherSessionElicitationRequest, CreateOtherRequestElicitationRequest], + Union[CreateFormSessionElicitationRequest, CreateFormRequestElicitationRequest], + Union[CreateUrlSessionElicitationRequest, CreateUrlRequestElicitationRequest], ], Any, ] @@ -5717,9 +5519,14 @@ class InitializeResponse(BaseModel): The client should disconnect, if it doesn't support this version. """ agent_capabilities: Annotated[ - Optional[AgentCapabilities], - Field(alias="agentCapabilities", validate_default=True), - ] = AgentCapabilities() + Optional[AgentCapabilities], Field(alias="agentCapabilities", validate_default=True) + ] = { + "loadSession": False, + "promptCapabilities": {"image": False, "audio": False, "embeddedContext": False}, + "mcpCapabilities": {"http": False, "sse": False, "acp": False}, + "sessionCapabilities": {}, + "auth": {}, + } """ Capabilities supported by the agent. """ @@ -5745,30 +5552,15 @@ class InitializeResponse(BaseModel): See protocol docs: [Extensibility](https://agentclientprotocol.com/protocol/extensibility) """ - @field_validator("agent_info", mode="wrap") - @classmethod - def _salvage_on_error_0(cls, value: Any, handler: Any) -> Any: - return salvage_on_error(value, handler, lambda: None) - - @field_validator("agent_capabilities", mode="wrap") + @field_validator("agent_capabilities", "agent_info", mode="wrap") @classmethod - def _salvage_on_error_1(cls, value: Any, handler: Any) -> Any: - return salvage_on_error( - value, - handler, - lambda: { - "loadSession": False, - "promptCapabilities": {"image": False, "audio": False, "embeddedContext": False}, - "mcpCapabilities": {"http": False, "sse": False, "acp": False}, - "sessionCapabilities": {}, - "auth": {}, - }, - ) + def use_default_on_error_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return use_default_on_error(v, handler, info) @field_validator("auth_methods", mode="wrap") @classmethod - def _skip_invalid_items_0(cls, value: Any, handler: Any) -> Any: - return skip_invalid_items(value, handler) + def skip_invalid_items_validator(cls, v: Any, handler: ValidatorFunctionWrapHandler, info: ValidationInfo) -> Any: + return skip_invalid_items(v, handler, info) class AgentNotification(BaseModel): @@ -5776,14 +5568,7 @@ class AgentNotification(BaseModel): """ The notification method name. """ - params: Optional[ - Union[ - SessionNotification, - CompleteElicitationNotification, - MessageMcpNotification, - Any, - ] - ] = None + params: Optional[Union[SessionNotification, CompleteElicitationNotification, MessageMcpNotification, Any]] = None """ Method-specific notification parameters. """ @@ -5822,9 +5607,68 @@ class AgentResponseMessage(BaseModel): class AgentResponse(RootModel[Union[AgentResponseMessage, AgentErrorMessage]]): - model_config = ConfigDict(use_attribute_docstrings=True) - root: Union[AgentResponseMessage, AgentErrorMessage] """ A JSON-RPC response object. """ + + +PermissionOptionKind = Literal["allow_once", "allow_always", "reject_once", "reject_always"] +PlanEntryPriority = Literal["high", "medium", "low"] +PlanEntryStatus = Literal["pending", "in_progress", "completed"] +StopReason = Literal["end_turn", "max_tokens", "max_turn_requests", "refusal", "cancelled"] +ToolCallStatus = Literal["pending", "in_progress", "completed", "failed"] +ToolKind = Literal[ + "read", + "edit", + "delete", + "move", + "search", + "execute", + "think", + "fetch", + "switch_mode", + "other", +] + +CreateOtherElicitationRequest = Union[ + CreateOtherSessionElicitationRequest, + CreateOtherRequestElicitationRequest, +] +CreateFormElicitationRequest = Union[ + CreateFormSessionElicitationRequest, + CreateFormRequestElicitationRequest, +] +CreateUrlElicitationRequest = Union[ + CreateUrlSessionElicitationRequest, + CreateUrlRequestElicitationRequest, +] +CreateElicitationRequest = Union[ + CreateFormElicitationRequest, + CreateUrlElicitationRequest, + CreateOtherElicitationRequest, +] + +CreateElicitationResponse = Union[ + AcceptElicitationResponse, + DeclineElicitationResponse, + CancelElicitationResponse, + OtherElicitationResponse, +] +ElicitationMode = Union[ + ElicitationFormSessionMode, + ElicitationFormRequestMode, + ElicitationUrlSessionMode, + ElicitationUrlRequestMode, +] + +_AvailableCommandsUpdate = AvailableCommandsUpdateBase +_CurrentModeUpdate = CurrentModeUpdateBase +_ConfigOptionUpdate = ConfigOptionUpdateBase +_SessionInfoUpdate = SessionInfoUpdateBase +_UsageUpdate = UsageUpdateBase +_StringMultiSelectItems = StringMultiSelectItemsBase + + +class Jsonrpc(Enum): + field_2_0 = "2.0" diff --git a/tests/test_elicitation_catchall.py b/tests/test_elicitation_catchall.py index 331185f..aa94849 100644 --- a/tests/test_elicitation_catchall.py +++ b/tests/test_elicitation_catchall.py @@ -61,7 +61,11 @@ def test_unknown_elicitation_mode_dispatches_to_clean_request_error() -> None: # A custom mode parses (above); the client router must then reject it with a clean # RequestError (invalid params) rather than a bare TypeError that surfaces as an # opaque -32603 internal error. - request = CreateOtherElicitationRequest(message="hi", mode="x-voice") + request = _REQUEST.validate_python({ + "mode": "x-voice", + "message": "hi", + "sessionId": "sess-1", + }) with pytest.raises(RequestError) as exc_info: _mode_from_create_elicitation_request(request) assert isinstance(exc_info.value, RequestError) diff --git a/tests/test_gen_all.py b/tests/test_gen_all.py index 57a4be9..e1d586e 100644 --- a/tests/test_gen_all.py +++ b/tests/test_gen_all.py @@ -1,13 +1,8 @@ -from acp.schema import AvailableCommandInput, ReadTextFileRequest +from pathlib import Path + +from acp.schema import ReadTextFileRequest from scripts.gen_all import resolve_ref, schema_source_paths -from scripts.gen_schema import ( - _deserialize_field_specs, - _extensible_union_excluded_tags, - _fallback_expression, - _normalize_catchall_unions, - _preprocess_schema_for_codegen, - _restore_required_nullable_fields, -) +from scripts.gen_schema import generate_schema def test_generated_field_descriptions_are_introspectable() -> None: @@ -15,10 +10,6 @@ def test_generated_field_descriptions_are_introspectable() -> None: assert ReadTextFileRequest.model_fields["path"].description == path_description assert ReadTextFileRequest.model_json_schema()["properties"]["path"]["description"] == path_description - root_description = "The input specification for a command." - assert AvailableCommandInput.model_fields["root"].description == root_description - assert AvailableCommandInput.model_json_schema()["description"] == root_description - def test_resolve_ref_accepts_schema_release_tags() -> None: assert resolve_ref("schema-v1.16.0") == "refs/tags/schema-v1.16.0" @@ -57,159 +48,9 @@ def test_parse_args_can_skip_format(monkeypatch) -> None: assert gen_all.parse_args().format_output is False -def test_codegen_preprocess_distributes_common_object_properties() -> None: - schema = { - "$defs": { - "ScopeA": { - "type": "object", - "properties": {"scopeA": {"type": "string"}}, - "required": ["scopeA"], - }, - "ScopeB": { - "type": "object", - "properties": {"scopeB": {"type": "string"}}, - "required": ["scopeB"], - }, - "Mode": { - "type": "object", - "properties": {"payload": {"type": "string"}}, - "required": ["payload"], - "anyOf": [ - {"allOf": [{"$ref": "#/$defs/ScopeA"}]}, - {"allOf": [{"$ref": "#/$defs/ScopeB"}]}, - ], - }, - "Request": { - "type": "object", - "properties": {"message": {"type": "string"}}, - "required": ["message"], - "oneOf": [ - { - "type": "object", - "properties": {"kind": {"type": "string", "const": "mode"}}, - "required": ["kind"], - "allOf": [{"$ref": "#/$defs/Mode"}], - } - ], - }, - }, - "$ref": "#/$defs/Request", - } - - request = _preprocess_schema_for_codegen(schema)["$defs"]["Request"] - - assert len(request["oneOf"]) == 2 - assert request["oneOf"][0]["required"] == ["message", "kind", "payload"] - assert request["oneOf"][0]["properties"].keys() >= {"message", "kind", "payload"} - assert request["oneOf"][0]["allOf"] == [{"$ref": "#/$defs/ScopeA"}] - assert request["oneOf"][1]["allOf"] == [{"$ref": "#/$defs/ScopeB"}] - - -def test_codegen_preprocess_normalizes_catchall_unions() -> None: - schema = { - "anyOf": [ - { - "type": "object", - "properties": {"type": {"type": "string", "const": "known"}}, - "required": ["type"], - }, - { - "title": "other", - "description": "Custom or future.", - "type": "object", - "properties": {"type": {"type": "string"}}, - "required": ["type"], - "not": {"anyOf": [{"const": "known"}]}, - "unevaluatedProperties": True, - }, - ], - "discriminator": {"propertyName": "type"}, - } - - normalized = _normalize_catchall_unions(schema) - - assert "discriminator" not in normalized - known, other = normalized["anyOf"] - assert known["properties"]["type"]["const"] == "known" - assert other["additionalProperties"] is True - assert other["properties"] == {"type": {"type": "string"}} - assert other["required"] == ["type"] - assert "not" not in other - assert "unevaluatedProperties" not in other - - -def test_extensible_union_excluded_tags_reads_not_clause() -> None: - union_def = { - "discriminator": {"propertyName": "action"}, - "anyOf": [ - {"properties": {"action": {"const": "accept"}}, "required": ["action"]}, - { - "title": "other", - "properties": {"action": {"type": "string"}}, - "not": { - "anyOf": [ - {"properties": {"action": {"const": "accept"}}}, - {"properties": {"action": {"const": "decline"}}}, - ] - }, - }, - ], - } - - assert _extensible_union_excluded_tags(union_def, "action") == ("accept", "decline") - - -def test_deserialize_field_specs_groups_by_fallback_and_excludes_meta() -> None: - definition = { - "required": ["items"], - "properties": { - "_meta": {"x-deserialize-default-on-error": True}, - "note": {"type": "string", "x-deserialize-default-on-error": True}, - "flag": {"type": "boolean", "default": False, "x-deserialize-default-on-error": True}, - "items": {"type": "array", "x-deserialize-skip-invalid-items": True}, - }, - } - - salvage, skip = _deserialize_field_specs(definition) - - assert salvage == {"lambda: None": ["note"], "lambda: False": ["flag"]} - assert skip == ["items"] - - -def test_fallback_expression_matches_schema_default_rules() -> None: - assert _fallback_expression({"default": False}, is_required=False) == "lambda: False" - assert _fallback_expression({"type": "array"}, is_required=True) == "lambda: []" - assert _fallback_expression({"type": ["array", "null"]}, is_required=False) == "lambda: None" - assert _fallback_expression({"type": "string"}, is_required=False) == "lambda: None" - - -def test_codegen_postprocess_preserves_required_nullable_fields() -> None: - schema = { - "$defs": { - "Example": { - "type": "object", - "properties": { - "requiredId": {"anyOf": [{"type": "null"}, {"type": "string"}]}, - "optionalId": {"anyOf": [{"type": "null"}, {"type": "string"}]}, - }, - "required": ["requiredId"], - } - } - } - content = """\ -class Example(BaseModel): - required_id: Annotated[ - Optional[str], - Field(alias="requiredId"), - ] = None - optional_id: Annotated[ - Optional[str], - Field(alias="optionalId"), - ] = None -""" - - processed = _restore_required_nullable_fields(content, schema) - - assert 'Field(alias="requiredId"),\n ] = None' not in processed - assert 'Field(alias="requiredId"),\n ]' in processed - assert 'Field(alias="optionalId"),\n ] = None' in processed +def test_codegen_check_is_clean_and_read_only() -> None: + output = Path("src/acp/schema.py") + before = output.read_bytes() + + assert generate_schema(check=True) + assert output.read_bytes() == before diff --git a/tests/test_utils.py b/tests/test_utils.py index bf00257..436fd73 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -34,6 +34,15 @@ def test_serialize_params_omits_meta_when_absent() -> None: assert "_meta" not in payload["content"] +def test_serialize_params_includes_defaulted_discriminators() -> None: + chunk = AgentMessageChunk(content=TextContentBlock(text="demo")) + + assert serialize_params(chunk) == { + "sessionUpdate": "agent_message_chunk", + "content": {"type": "text", "text": "demo"}, + } + + def test_field_meta_can_be_set_by_name_on_models() -> None: chunk = AgentMessageChunk( session_update="agent_message_chunk", diff --git a/uv.lock b/uv.lock index c4d74e3..30e9799 100644 --- a/uv.lock +++ b/uv.lock @@ -12,6 +12,7 @@ version = "0.12.1" source = { editable = "." } dependencies = [ { name = "pydantic" }, + { name = "pydantic-core" }, ] [package.optional-dependencies] @@ -49,13 +50,14 @@ requires-dist = [ { name = "logfire", marker = "extra == 'logfire'", specifier = ">=0.14" }, { name = "opentelemetry-sdk", marker = "extra == 'logfire'", specifier = ">=1.28.0" }, { name = "pydantic", specifier = ">=2.7" }, + { name = "pydantic-core", specifier = ">=2.18.1" }, { name = "websockets", marker = "extra == 'http'", specifier = ">=12.0" }, ] provides-extras = ["logfire", "http"] [package.metadata.requires-dev] dev = [ - { name = "datamodel-code-generator", specifier = ">=0.71.0" }, + { name = "datamodel-code-generator", specifier = "==0.71.0" }, { name = "deptry", specifier = ">=0.23.0" }, { name = "httpx", extras = ["http2"], specifier = ">=0.27" }, { name = "mkdocs", specifier = ">=1.4.2" },