From f84094a7fb3500a73071b576fea267cd8bc70409 Mon Sep 17 00:00:00 2001 From: CW5201 <19317334773@163.com> Date: Sat, 19 Sep 2026 17:05:33 +0800 Subject: [PATCH] fix(tools): validate HTTP tool arguments against input_schema before dispatch When a model-generated tool call omits a required field or passes a wrong-typed/nested argument, the executor previously forwarded the raw arguments straight to the downstream HTTP service, which failed with an opaque 4xx instead of a structured error the agent loop could use to self-correct. - ToolExecutor.execute() now validates arguments against the tool's input_schema before the HTTP dispatch (sync and detached-submit paths) - Returns a new SCHEMA_INVALID ToolError listing up to 5 validation issues (missing required fields, type mismatches, dotted JSON paths) without issuing the HTTP request - None / empty / non-dict / malformed schemas are logged and skipped, preserving existing behaviour for legacy tools - MCP and A2A paths are intentionally untouched: their schemas belong to the provider, which is contractually responsible for validation - No new dependencies (jsonschema is already a backend dep), no DB, API, or ToolError structure changes --- backend/app/tools/tool_executor.py | 75 +++++++++ backend/tests/test_tool_executor.py | 239 +++++++++++++++++++++++++++- 2 files changed, 310 insertions(+), 4 deletions(-) diff --git a/backend/app/tools/tool_executor.py b/backend/app/tools/tool_executor.py index 08f4a02ef..bbcde210e 100644 --- a/backend/app/tools/tool_executor.py +++ b/backend/app/tools/tool_executor.py @@ -2,14 +2,18 @@ import base64 import json +import logging import os import re +from collections.abc import Mapping from dataclasses import dataclass from datetime import timedelta from typing import Any from urllib.parse import urlsplit import httpx +from jsonschema import Draft202012Validator +from jsonschema.exceptions import SchemaError from sqlalchemy.exc import IntegrityError from sqlmodel import Session, select @@ -24,7 +28,10 @@ from app.tools.mcp_client import MCPClientError, execute_mcp_tool, execute_mcp_tool_result from app.tools.tool_schema import MCPAppDescriptor, ToolCall, ToolError, ToolResult +logger = logging.getLogger(__name__) + SECRET_PATTERN = re.compile(r"\$\{secret\.([A-Z0-9_]+)\}") +SCHEMA_INVALID_ERROR_LIMIT = 5 def _json_path(value: Any, path: str) -> Any: @@ -123,6 +130,10 @@ def execute( tool.name, "UNSUPPORTED_TOOL_TYPE", f"不支持的工具类型:{tool.tool_type}" ) + schema_error = self._validate_http_arguments(tool, tool_call.arguments) + if schema_error is not None: + return schema_error + execution = ( tool.config_json.get("execution", {}) if isinstance(tool.config_json, dict) else {} ) @@ -652,6 +663,70 @@ def repl(match: re.Match[str]) -> str: return SECRET_PATTERN.sub(repl, value) + def _validate_http_arguments(self, tool: Tool, arguments: dict[str, Any]) -> ToolResult | None: + """Validate HTTP tool arguments against the tool's input_schema. + + Returns an error ToolResult when arguments violate a well-formed + schema. Malformed or empty schemas are logged and skipped so a bad + schema never blocks a previously-working request. + """ + schema = tool.input_schema + if schema is None or schema == {}: + return None + if not isinstance(schema, Mapping): + logger.warning( + "Skipping argument validation for tool %s: input_schema is not a dict.", + tool.name, + ) + return None + try: + validator = Draft202012Validator(schema) + except (SchemaError, ValueError, TypeError, AttributeError) as exc: + logger.warning( + "Skipping argument validation for tool %s: malformed input_schema (%s).", + tool.name, + exc, + ) + return None + + try: + errors = sorted( + validator.iter_errors(arguments), + key=lambda e: (list(e.path), e.message), + ) + except (TypeError, ValueError, AttributeError) as exc: + # Malformed schemas (e.g. required set to a scalar) can only fail + # during evaluation, not at validator construction time. + logger.warning( + "Skipping argument validation for tool %s: input_schema is not evaluable (%s).", + tool.name, + exc, + ) + return None + if not errors: + return None + + detail = self._schema_error_message(errors) + return self._error(tool.name, "SCHEMA_INVALID", detail) + + @staticmethod + def _schema_error_message(errors: list) -> str: + parts = [] + for err in errors[:SCHEMA_INVALID_ERROR_LIMIT]: + path = "$." + ".".join(str(p) for p in err.path) if err.path else "$" + if err.validator == "required": + missing = str(err.message).split("is")[0] or "required field" + parts.append(f"{path} 缺少必填字段 {missing.strip()}") + else: + parts.append(f"{path}: {err.message}") + if len(errors) > SCHEMA_INVALID_ERROR_LIMIT: + parts.append(f"(另有 {len(errors) - SCHEMA_INVALID_ERROR_LIMIT} 项未显示)") + return ( + "工具参数不符合 input_schema:" + + ";".join(parts) + + "。请修正参数后重新调用。" + ) + def _error(self, tool_name: str, code: str, message: str) -> ToolResult: return ToolResult( tool_name=tool_name, diff --git a/backend/tests/test_tool_executor.py b/backend/tests/test_tool_executor.py index b4af3afac..9111cfba9 100644 --- a/backend/tests/test_tool_executor.py +++ b/backend/tests/test_tool_executor.py @@ -4,14 +4,14 @@ import httpx import pytest +from sqlalchemy.pool import StaticPool +from sqlmodel import Session, SQLModel, create_engine, select from app.agents.branching import ensure_private_resource_binding -from app.tools.tool_executor import ToolExecutor -from app.tools.tool_schema import ToolCall from app.db.models import A2ATaskEvent, A2ATaskRun, AgentProfile, MCPServer, Tenant, Tool from app.security.internal_service import INTERNAL_SERVICE_HEADER, internal_service_token -from sqlalchemy.pool import StaticPool -from sqlmodel import Session, SQLModel, create_engine, select +from app.tools.tool_executor import ToolExecutor +from app.tools.tool_schema import ToolCall def test_resolve_secret_header(monkeypatch): @@ -899,6 +899,237 @@ def _mock_mcp_server_path() -> Path: return Path(__file__).resolve().parents[1] / "mock_servers" / "mcp_stdio_server.py" +def _http_schema_session( + db, input_schema, *, extra_tool_kwargs=None +): + """Insert a minimal HTTP tool and return (db, captured, FakeClient-agnostic name).""" + db.add(Tenant(id="tenant_demo", name="Demo")) + kwargs = dict( + tenant_id="tenant_demo", + name="order.query", + tool_type="http", + method="POST", + url="https://example.test/order/query", + input_schema=input_schema, + enabled=True, + ) + kwargs.update(extra_tool_kwargs or {}) + db.add(Tool(**kwargs)) + db.commit() + return db + + +class _RecordingFakeClient: + """httpx.Client stand-in that records requests and returns 200 JSON.""" + + def __init__(self, *, timeout: float): + self.requests: list[tuple[str, str, dict | None]] = [] + + def __enter__(self): + return self + + def __exit__(self, *args): + return None + + def request(self, method, url, headers=None, json=None, params=None): + self.requests.append((method, url, json)) + return httpx.Response(200, json={"ok": True}, request=httpx.Request(method, url)) + + +def test_http_tool_missing_required_argument_returns_schema_invalid(monkeypatch) -> None: + recorded: list[_RecordingFakeClient] = [] + + class FakeClient(_RecordingFakeClient): + def __init__(self, **kwargs): + super().__init__(**kwargs) + recorded.append(self) + + monkeypatch.setattr(httpx, "Client", FakeClient) + with _test_session() as db: + _http_schema_session( + db, + {"type": "object", "properties": {"order_id": {"type": "string"}}, "required": ["order_id"]}, + ) + result = ToolExecutor(db).execute( + "tenant_demo", ToolCall(name="order.query", arguments={"wrong": 1}) + ) + + assert result.success is False + assert result.error.code == "SCHEMA_INVALID" + assert "order_id" in result.error.message + assert not any(len(client.requests) for client in recorded), "HTTP request must not fire" + + +def test_http_tool_wrong_type_argument_returns_schema_invalid(monkeypatch) -> None: + recorded: list[_RecordingFakeClient] = [] + + class FakeClient(_RecordingFakeClient): + def __init__(self, **kwargs): + super().__init__(**kwargs) + recorded.append(self) + + monkeypatch.setattr(httpx, "Client", FakeClient) + with _test_session() as db: + _http_schema_session( + db, + {"type": "object", "properties": {"amount": {"type": "integer"}}, "required": ["amount"]}, + ) + result = ToolExecutor(db).execute( + "tenant_demo", ToolCall(name="order.query", arguments={"amount": "not-a-number"}) + ) + + assert result.success is False + assert result.error.code == "SCHEMA_INVALID" + assert "amount" in result.error.message + assert not any(len(client.requests) for client in recorded) + + +def test_http_tool_nested_schema_error_includes_path(monkeypatch) -> None: + recorded: list[_RecordingFakeClient] = [] + + class FakeClient(_RecordingFakeClient): + def __init__(self, **kwargs): + super().__init__(**kwargs) + recorded.append(self) + + monkeypatch.setattr(httpx, "Client", FakeClient) + with _test_session() as db: + _http_schema_session( + db, + { + "type": "object", + "properties": {"spec": {"type": "object", "properties": {"size": {"type": "integer"}}}}, + "required": ["spec"], + }, + ) + result = ToolExecutor(db).execute( + "tenant_demo", ToolCall(name="order.query", arguments={"spec": {"size": "large"}}) + ) + + assert result.success is False + assert result.error.code == "SCHEMA_INVALID" + assert "spec.size" in result.error.message + assert not any(len(client.requests) for client in recorded) + + +def test_http_tool_valid_arguments_dispatch_normally(monkeypatch) -> None: + recorded: list[_RecordingFakeClient] = [] + + class FakeClient(_RecordingFakeClient): + def __init__(self, **kwargs): + super().__init__(**kwargs) + recorded.append(self) + + monkeypatch.setattr(httpx, "Client", FakeClient) + with _test_session() as db: + _http_schema_session( + db, + {"type": "object", "properties": {"order_id": {"type": "string"}}, "required": ["order_id"]}, + ) + result = ToolExecutor(db).execute( + "tenant_demo", ToolCall(name="order.query", arguments={"order_id": "O-1"}) + ) + + assert result.success is True + assert result.data == {"ok": True} + assert len(recorded) == 1 + assert recorded[0].requests[0][2] == {"order_id": "O-1"} + + +def test_http_tool_empty_schema_keeps_previous_behavior(monkeypatch) -> None: + recorded: list[_RecordingFakeClient] = [] + + class FakeClient(_RecordingFakeClient): + def __init__(self, **kwargs): + super().__init__(**kwargs) + recorded.append(self) + + monkeypatch.setattr(httpx, "Client", FakeClient) + with _test_session() as db: + _http_schema_session(db, {}) + result = ToolExecutor(db).execute( + "tenant_demo", ToolCall(name="order.query", arguments={"anything": 1}) + ) + + assert result.success is True + assert len(recorded) == 1 + + +def test_http_tool_none_schema_keeps_previous_behavior(monkeypatch) -> None: + recorded: list[_RecordingFakeClient] = [] + + class FakeClient(_RecordingFakeClient): + def __init__(self, **kwargs): + super().__init__(**kwargs) + recorded.append(self) + + monkeypatch.setattr(httpx, "Client", FakeClient) + with _test_session() as db: + _http_schema_session(db, None) + result = ToolExecutor(db).execute( + "tenant_demo", ToolCall(name="order.query", arguments={"anything": 1}) + ) + + assert result.success is True + assert len(recorded) == 1 + + +def test_http_tool_malformed_schema_is_skipped_not_blocking(monkeypatch) -> None: + recorded: list[_RecordingFakeClient] = [] + + class FakeClient(_RecordingFakeClient): + def __init__(self, **kwargs): + super().__init__(**kwargs) + recorded.append(self) + + monkeypatch.setattr(httpx, "Client", FakeClient) + with _test_session() as db: + # 'required' must be a list of unique items in a valid schema; a scalar is malformed. + _http_schema_session(db, {"type": "object", "required": 5}) + result = ToolExecutor(db).execute( + "tenant_demo", ToolCall(name="order.query", arguments={}) + ) + + assert result.success is True, "malformed schema must not block the request" + assert len(recorded) == 1 + + +def test_mcp_tool_with_input_schema_bypasses_http_validation(monkeypatch) -> None: + recorded: list[_RecordingFakeClient] = [] + + class FakeClient(_RecordingFakeClient): + def __init__(self, **kwargs): + super().__init__(**kwargs) + recorded.append(self) + + monkeypatch.setattr(httpx, "Client", FakeClient) + with _test_session() as db: + db.add(Tenant(id="tenant_demo", name="Demo")) + db.add(MCPServer(id="server_builtin", tenant_id="tenant_demo", name="builtin", transport="builtin")) + db.add( + Tool( + tenant_id="tenant_demo", + name="mcp.demo_echo", + tool_type="mcp", + method="POST", + url="mcp://builtin.demo/echo", + mcp_server_id="server_builtin", + config_json={"tool": "echo"}, + input_schema={"type": "object", "properties": {"text": {"type": "string"}}, "required": ["text"]}, + enabled=True, + ) + ) + db.commit() + result = ToolExecutor(db).execute( + "tenant_demo", + ToolCall(name="mcp.demo_echo", arguments={"text": "hello"}), + ) + + # MCP path is unchanged: succeeds via the builtin MCP server, no HTTP client is used. + assert result.success is True + assert not recorded, "MCP tool must not go through the HTTP client" + + def _test_session(): engine = create_engine( "sqlite://",