Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions backend/app/tools/tool_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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:
Expand Down Expand Up @@ -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 {}
)
Expand Down Expand Up @@ -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,
Expand Down
239 changes: 235 additions & 4 deletions backend/tests/test_tool_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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://",
Expand Down