mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 10:38:32 +08:00
feat(mcp): support MCP protocol 2025-06-18 for workflow-as-MCP server (version negotiation + structured output) (#37892)
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> Co-authored-by: yunlu.wen <yunlu.wen@dify.ai>
This commit is contained in:
parent
dbd3316615
commit
5a342f9258
@ -1,6 +1,6 @@
|
||||
from typing import Any, Union
|
||||
|
||||
from flask import Response
|
||||
from flask import Response, request
|
||||
from flask_restx import Resource
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
from sqlalchemy import select
|
||||
@ -9,7 +9,7 @@ from sqlalchemy.orm import Session, sessionmaker
|
||||
from controllers.common.schema import register_schema_model
|
||||
from controllers.mcp import mcp_ns
|
||||
from core.mcp import types as mcp_types
|
||||
from core.mcp.server.streamable_http import handle_mcp_request
|
||||
from core.mcp.server.streamable_http import handle_mcp_request, negotiate_protocol_version
|
||||
from extensions.ext_database import db
|
||||
from graphon.variables.input_entities import VariableEntity, VariableEntityType
|
||||
from libs import helper
|
||||
@ -68,6 +68,17 @@ class MCPAppApi(Resource):
|
||||
request_id: Union[int, str] | None = args.id
|
||||
mcp_request = self._parse_mcp_request(args.model_dump(exclude_none=True))
|
||||
|
||||
# Resolve the negotiated protocol version from the MCP-Protocol-Version header.
|
||||
is_initialize = isinstance(mcp_request.root, mcp_types.InitializeRequest)
|
||||
header_value = request.headers.get("MCP-Protocol-Version")
|
||||
protocol_version = negotiate_protocol_version(header_value, is_initialize)
|
||||
if protocol_version is None:
|
||||
# A notification never receives a response, even with an unsupported header.
|
||||
if isinstance(mcp_request, mcp_types.ClientNotification):
|
||||
protocol_version = mcp_types.DEFAULT_NEGOTIATED_VERSION
|
||||
else:
|
||||
return self._protocol_version_error_response(request_id, header_value)
|
||||
|
||||
with sessionmaker(db.engine, expire_on_commit=False).begin() as session:
|
||||
# Get MCP server and app
|
||||
mcp_server, app = self._get_mcp_server_and_app(server_code, session)
|
||||
@ -77,7 +88,28 @@ class MCPAppApi(Resource):
|
||||
user_input_form = self._get_user_input_form(app)
|
||||
|
||||
# Handle notification vs request differently
|
||||
return self._process_mcp_message(mcp_request, request_id, app, mcp_server, user_input_form, session)
|
||||
return self._process_mcp_message(
|
||||
mcp_request, request_id, app, mcp_server, user_input_form, session, protocol_version
|
||||
)
|
||||
|
||||
def _protocol_version_error_response(
|
||||
self, request_id: Union[int, str] | None, header_value: str | None
|
||||
) -> Response:
|
||||
"""Return a JSON-RPC error for an unsupported MCP-Protocol-Version header.
|
||||
|
||||
Per JSON-RPC 2.0, an error whose request id is unknown uses a null id, so we echo the
|
||||
offending request's id directly (None -> null) instead of fabricating a placeholder.
|
||||
"""
|
||||
error_data = mcp_types.ErrorData(
|
||||
code=mcp_types.INVALID_REQUEST,
|
||||
message=f"Unsupported MCP-Protocol-Version: {header_value}",
|
||||
)
|
||||
error_response = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"error": error_data.model_dump(by_alias=True, mode="json", exclude_none=True),
|
||||
}
|
||||
return helper.compact_generate_response(error_response)
|
||||
|
||||
def _get_mcp_server_and_app(self, server_code: str, session: Session) -> tuple[AppMCPServer, App]:
|
||||
"""Get and validate MCP server and app in one query session"""
|
||||
@ -104,12 +136,15 @@ class MCPAppApi(Resource):
|
||||
mcp_server: AppMCPServer,
|
||||
user_input_form: list[VariableEntity],
|
||||
session: Session,
|
||||
protocol_version: str,
|
||||
) -> Response:
|
||||
"""Process MCP message (notification or request)"""
|
||||
if isinstance(mcp_request, mcp_types.ClientNotification):
|
||||
return self._handle_notification(mcp_request)
|
||||
else:
|
||||
return self._handle_request(mcp_request, request_id, app, mcp_server, user_input_form, session)
|
||||
return self._handle_request(
|
||||
mcp_request, request_id, app, mcp_server, user_input_form, session, protocol_version
|
||||
)
|
||||
|
||||
def _handle_notification(self, mcp_request: mcp_types.ClientNotification) -> Response:
|
||||
"""Handle MCP notification"""
|
||||
@ -127,12 +162,15 @@ class MCPAppApi(Resource):
|
||||
mcp_server: AppMCPServer,
|
||||
user_input_form: list[VariableEntity],
|
||||
session: Session,
|
||||
protocol_version: str,
|
||||
) -> Response:
|
||||
"""Handle MCP request"""
|
||||
if request_id is None:
|
||||
raise MCPRequestError(mcp_types.INVALID_REQUEST, "Request ID is required")
|
||||
|
||||
result = self._handle_mcp_request(app, mcp_server, mcp_request, user_input_form, session, request_id)
|
||||
result = self._handle_mcp_request(
|
||||
app, mcp_server, mcp_request, user_input_form, session, request_id, protocol_version
|
||||
)
|
||||
if result is None:
|
||||
# This shouldn't happen for requests, but handle gracefully
|
||||
raise MCPRequestError(mcp_types.INTERNAL_ERROR, "No response generated for request")
|
||||
@ -229,6 +267,7 @@ class MCPAppApi(Resource):
|
||||
user_input_form: list[VariableEntity],
|
||||
session: Session,
|
||||
request_id: Union[int, str],
|
||||
protocol_version: str,
|
||||
) -> mcp_types.JSONRPCResponse | mcp_types.JSONRPCError | None:
|
||||
"""Handle MCP request and return response"""
|
||||
end_user = self._retrieve_end_user(mcp_server.tenant_id, mcp_server.id)
|
||||
@ -238,4 +277,6 @@ class MCPAppApi(Resource):
|
||||
client_name = f"{client_info.name}@{client_info.version}"
|
||||
end_user = self._create_end_user(client_name, app.tenant_id, app.id, mcp_server.id, session)
|
||||
|
||||
return handle_mcp_request(session, app, mcp_request, user_input_form, mcp_server, end_user, request_id)
|
||||
return handle_mcp_request(
|
||||
session, app, mcp_request, user_input_form, mcp_server, end_user, request_id, protocol_version
|
||||
)
|
||||
|
||||
@ -15,6 +15,36 @@ from services.app_generate_service import AppGenerateService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Structured tool output (outputSchema + structuredContent) was introduced in MCP 2025-06-18.
|
||||
STRUCTURED_OUTPUT_MIN_VERSION = "2025-06-18"
|
||||
|
||||
|
||||
def _supports_structured_output(protocol_version: str) -> bool:
|
||||
"""Return True when the negotiated protocol version supports structured tool output.
|
||||
|
||||
MCP protocol versions are YYYY-MM-DD strings, so lexical comparison equals chronological.
|
||||
"""
|
||||
return protocol_version >= STRUCTURED_OUTPUT_MIN_VERSION
|
||||
|
||||
|
||||
def negotiate_protocol_version(header_value: str | None, is_initialize: bool) -> str | None:
|
||||
"""Resolve the negotiated protocol version for an incoming MCP request.
|
||||
|
||||
The version is taken from the MCP-Protocol-Version header on post-initialize requests.
|
||||
Returns the version to use for behavior gating, or None when the client sent an explicit
|
||||
but unsupported header (the caller should reply with a JSON-RPC INVALID_REQUEST error).
|
||||
Initialize requests negotiate via the request body, so they always receive
|
||||
DEFAULT_NEGOTIATED_VERSION and their header is never validated or rejected.
|
||||
"""
|
||||
if is_initialize:
|
||||
return mcp_types.DEFAULT_NEGOTIATED_VERSION
|
||||
# Treat an absent or empty header as "not specified" -> default version.
|
||||
if not header_value:
|
||||
return mcp_types.DEFAULT_NEGOTIATED_VERSION
|
||||
if header_value not in mcp_types.SERVER_SUPPORTED_PROTOCOL_VERSIONS:
|
||||
return None
|
||||
return header_value
|
||||
|
||||
|
||||
class ToolParameterSchemaDict(TypedDict):
|
||||
type: str
|
||||
@ -35,6 +65,7 @@ def handle_mcp_request(
|
||||
mcp_server: AppMCPServer,
|
||||
end_user: EndUser | None = None,
|
||||
request_id: int | str = 1,
|
||||
protocol_version: str = mcp_types.DEFAULT_NEGOTIATED_VERSION,
|
||||
) -> mcp_types.JSONRPCResponse | mcp_types.JSONRPCError:
|
||||
"""
|
||||
Handle MCP request and return JSON-RPC response
|
||||
@ -77,15 +108,24 @@ def handle_mcp_request(
|
||||
# Dispatch request to appropriate handler based on instance type
|
||||
match request_root:
|
||||
case mcp_types.InitializeRequest():
|
||||
return create_success_response(handle_initialize(mcp_server.description))
|
||||
return create_success_response(
|
||||
handle_initialize(mcp_server.description, request_root.params.protocolVersion)
|
||||
)
|
||||
case mcp_types.ListToolsRequest():
|
||||
return create_success_response(
|
||||
handle_list_tools(
|
||||
app.name, app.mode, user_input_form, mcp_server.description, mcp_server.parameters_dict
|
||||
app.name,
|
||||
app.mode,
|
||||
user_input_form,
|
||||
mcp_server.description,
|
||||
mcp_server.parameters_dict,
|
||||
protocol_version,
|
||||
)
|
||||
)
|
||||
case mcp_types.CallToolRequest():
|
||||
return create_success_response(handle_call_tool(session, app, request, user_input_form, end_user))
|
||||
return create_success_response(
|
||||
handle_call_tool(session, app, request, user_input_form, end_user, protocol_version)
|
||||
)
|
||||
case mcp_types.PingRequest():
|
||||
return create_success_response(handle_ping())
|
||||
case _:
|
||||
@ -104,14 +144,22 @@ def handle_ping() -> mcp_types.EmptyResult:
|
||||
return mcp_types.EmptyResult()
|
||||
|
||||
|
||||
def handle_initialize(description: str) -> mcp_types.InitializeResult:
|
||||
"""Handle initialize request"""
|
||||
def handle_initialize(description: str, requested_version: str | int) -> mcp_types.InitializeResult:
|
||||
"""Handle initialize request, negotiating the protocol version with the client.
|
||||
|
||||
Echoes the client's requested version when the server supports it, otherwise returns the
|
||||
server's latest supported version (per the MCP lifecycle spec).
|
||||
"""
|
||||
negotiated_version: str = mcp_types.SERVER_LATEST_PROTOCOL_VERSION
|
||||
if isinstance(requested_version, str) and requested_version in mcp_types.SERVER_SUPPORTED_PROTOCOL_VERSIONS:
|
||||
negotiated_version = requested_version
|
||||
|
||||
capabilities = mcp_types.ServerCapabilities(
|
||||
tools=mcp_types.ToolsCapability(listChanged=False),
|
||||
)
|
||||
|
||||
return mcp_types.InitializeResult(
|
||||
protocolVersion=mcp_types.SERVER_LATEST_PROTOCOL_VERSION,
|
||||
protocolVersion=negotiated_version,
|
||||
capabilities=capabilities,
|
||||
serverInfo=mcp_types.Implementation(name="Dify", version=dify_config.project.version),
|
||||
instructions=description,
|
||||
@ -124,19 +172,23 @@ def handle_list_tools(
|
||||
user_input_form: list[VariableEntity],
|
||||
description: str,
|
||||
parameters_dict: dict[str, str],
|
||||
protocol_version: str = mcp_types.DEFAULT_NEGOTIATED_VERSION,
|
||||
) -> mcp_types.ListToolsResult:
|
||||
"""Handle list tools request"""
|
||||
parameter_schema = build_parameter_schema(app_mode, user_input_form, parameters_dict)
|
||||
supports_structured = _supports_structured_output(protocol_version)
|
||||
|
||||
return mcp_types.ListToolsResult(
|
||||
tools=[
|
||||
mcp_types.Tool(
|
||||
name=app_name,
|
||||
description=description,
|
||||
inputSchema=cast(dict[str, Any], parameter_schema),
|
||||
)
|
||||
],
|
||||
# For 2025-06-18+ clients, expose an explicit display title and a permissive output
|
||||
# schema. Both stay None (and are stripped by exclude_none serialization) for older
|
||||
# clients, so their tool definition is unchanged.
|
||||
tool = mcp_types.Tool(
|
||||
name=app_name,
|
||||
title=app_name if supports_structured else None,
|
||||
description=description,
|
||||
inputSchema=cast(dict[str, Any], parameter_schema),
|
||||
outputSchema={"type": "object"} if supports_structured else None,
|
||||
)
|
||||
return mcp_types.ListToolsResult(tools=[tool])
|
||||
|
||||
|
||||
def handle_call_tool(
|
||||
@ -145,6 +197,7 @@ def handle_call_tool(
|
||||
request: mcp_types.ClientRequest,
|
||||
user_input_form: list[VariableEntity],
|
||||
end_user: EndUser | None,
|
||||
protocol_version: str = mcp_types.DEFAULT_NEGOTIATED_VERSION,
|
||||
) -> mcp_types.CallToolResult:
|
||||
"""Handle call tool request"""
|
||||
request_obj = cast(mcp_types.CallToolRequest, request.root)
|
||||
@ -163,7 +216,13 @@ def handle_call_tool(
|
||||
)
|
||||
|
||||
answer = extract_answer_from_response(app, response)
|
||||
return mcp_types.CallToolResult(content=[mcp_types.TextContent(text=answer, type="text")])
|
||||
structured_content = None
|
||||
if _supports_structured_output(protocol_version):
|
||||
structured_content = extract_structured_output(app, response, answer)
|
||||
return mcp_types.CallToolResult(
|
||||
content=[mcp_types.TextContent(text=answer, type="text")],
|
||||
structuredContent=structured_content,
|
||||
)
|
||||
|
||||
|
||||
def build_parameter_schema(
|
||||
@ -204,6 +263,29 @@ def prepare_tool_arguments(app: App, arguments: dict[str, Any]) -> ToolArguments
|
||||
return {"query": query, "inputs": args_copy}
|
||||
|
||||
|
||||
def extract_structured_output(app: App, response: Any, answer: str) -> dict[str, Any] | None:
|
||||
"""Build MCP structured tool output (2025-06-18) from the app response.
|
||||
|
||||
WORKFLOW mode exposes the raw outputs mapping; chat/agent/completion modes expose the
|
||||
answer string under an "answer" key. Returns None when no structured output is available.
|
||||
"""
|
||||
match app.mode:
|
||||
case AppMode.WORKFLOW:
|
||||
if isinstance(response, Mapping):
|
||||
data = response.get("data")
|
||||
if isinstance(data, Mapping):
|
||||
outputs = data.get("outputs")
|
||||
# All three guards use Mapping for consistency; coerce to a concrete dict
|
||||
# because structuredContent must be a JSON object (dict[str, Any]).
|
||||
if isinstance(outputs, Mapping):
|
||||
return dict(outputs)
|
||||
return None
|
||||
case AppMode.ADVANCED_CHAT | AppMode.CHAT | AppMode.AGENT_CHAT | AppMode.COMPLETION:
|
||||
return {"answer": answer}
|
||||
case _:
|
||||
return None
|
||||
|
||||
|
||||
def extract_answer_from_response(app: App, response: Any) -> str:
|
||||
"""Extract answer from app generate response"""
|
||||
answer = ""
|
||||
|
||||
@ -22,10 +22,13 @@ for reference.
|
||||
* Define additional model classes instead of using dictionaries. Do this even if they're
|
||||
not separate types in the schema.
|
||||
"""
|
||||
# Client support both version, not support 2025-06-18 yet.
|
||||
# Latest protocol version the Dify MCP client negotiates with upstream MCP servers.
|
||||
LATEST_PROTOCOL_VERSION = "2025-06-18"
|
||||
# Server support 2024-11-05 to allow claude to use.
|
||||
SERVER_LATEST_PROTOCOL_VERSION = "2024-11-05"
|
||||
# Latest protocol version the Dify MCP server advertises to connecting clients.
|
||||
SERVER_LATEST_PROTOCOL_VERSION = "2025-06-18"
|
||||
# Protocol versions the Dify MCP server can negotiate down to (e.g. Claude on 2024-11-05).
|
||||
SERVER_SUPPORTED_PROTOCOL_VERSIONS: frozenset[str] = frozenset({"2024-11-05", "2025-03-26", "2025-06-18"})
|
||||
# Version assumed when a client omits the MCP-Protocol-Version header on post-initialize requests.
|
||||
DEFAULT_NEGOTIATED_VERSION = "2025-03-26"
|
||||
ProgressToken = str | int
|
||||
Cursor = str
|
||||
|
||||
@ -8,7 +8,7 @@ from unittest.mock import MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from flask import Response
|
||||
from flask import Flask, Response
|
||||
from pydantic import ValidationError
|
||||
|
||||
import controllers.mcp.mcp as module
|
||||
@ -37,12 +37,15 @@ class DummyServer:
|
||||
self.app_id = app_id
|
||||
self.tenant_id = tenant_id
|
||||
self.id = server_id
|
||||
self.description = "Test server"
|
||||
self.parameters_dict = {}
|
||||
|
||||
|
||||
class DummyApp:
|
||||
def __init__(self, mode, workflow=None, app_model_config=None):
|
||||
self.id = _APP_ID
|
||||
self.tenant_id = _TENANT_ID
|
||||
self.name = "test_app"
|
||||
self.mode = mode
|
||||
self.workflow = workflow
|
||||
self.app_model_config = app_model_config
|
||||
@ -494,3 +497,220 @@ class TestMCPAppApi:
|
||||
with pytest.raises(module.MCPRequestError) as exc_info:
|
||||
post_fn("server-1")
|
||||
assert "Invalid user_input_form" in str(exc_info.value)
|
||||
|
||||
|
||||
_UNSUPPORTED_VERSION = "1999-01-01"
|
||||
|
||||
|
||||
def _initialize_payload(protocol_version: str = "2024-11-05") -> dict[str, object]:
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"method": "initialize",
|
||||
"id": 1,
|
||||
"params": {
|
||||
"protocolVersion": protocol_version,
|
||||
"capabilities": {},
|
||||
"clientInfo": {"name": "test-client", "version": "1.0"},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _tools_list_payload(request_id: int | None = 1) -> dict[str, object]:
|
||||
payload: dict[str, object] = {"jsonrpc": "2.0", "method": "tools/list", "params": {}}
|
||||
if request_id is not None:
|
||||
payload["id"] = request_id
|
||||
return payload
|
||||
|
||||
|
||||
def _tools_call_payload() -> dict[str, object]:
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"method": "tools/call",
|
||||
"id": 1,
|
||||
"params": {"name": "test_app", "arguments": {"query": "test question"}},
|
||||
}
|
||||
|
||||
|
||||
class TestMCPProtocolVersionNegotiationApi:
|
||||
"""MCP protocol version negotiation exercised through the HTTP controller layer.
|
||||
|
||||
Covers the MCP-Protocol-Version header contract (resolution, rejection, threading)
|
||||
and the serialized JSON responses seen by modern (2025-06-18) vs legacy (2024-11-05)
|
||||
clients, including the back-compat guarantee that legacy responses carry none of the
|
||||
structured-output fields.
|
||||
"""
|
||||
|
||||
def _make_api(self) -> module.MCPAppApi:
|
||||
server = DummyServer(status=module.AppMCPServerStatus.ACTIVE)
|
||||
app = DummyApp(mode=module.AppMode.CHAT, app_model_config=DummyConfig())
|
||||
api = module.MCPAppApi()
|
||||
api._get_mcp_server_and_app = MagicMock(return_value=(server, app))
|
||||
api._retrieve_end_user = MagicMock(return_value=MagicMock())
|
||||
return api
|
||||
|
||||
def _post(
|
||||
self, flask_app: Flask, api: module.MCPAppApi, payload: dict[str, object], headers: dict[str, str] | None = None
|
||||
) -> Response:
|
||||
fake_payload(payload)
|
||||
post_fn = unwrap(api.post)
|
||||
with flask_app.test_request_context(headers=headers):
|
||||
return post_fn("server-1")
|
||||
|
||||
@pytest.mark.parametrize("version", sorted(module.mcp_types.SERVER_SUPPORTED_PROTOCOL_VERSIONS))
|
||||
def test_initialize_echoes_supported_body_version(self, flask_app_with_containers, version):
|
||||
"""Initialize echoes every supported client-requested version back unchanged."""
|
||||
api = self._make_api()
|
||||
|
||||
response = self._post(flask_app_with_containers, api, _initialize_payload(version))
|
||||
|
||||
body = response.get_json()
|
||||
assert body["result"]["protocolVersion"] == version
|
||||
|
||||
def test_initialize_falls_back_for_unsupported_body_version(self, flask_app_with_containers):
|
||||
"""An unsupported requested version falls back to the server latest."""
|
||||
api = self._make_api()
|
||||
|
||||
response = self._post(flask_app_with_containers, api, _initialize_payload(_UNSUPPORTED_VERSION))
|
||||
|
||||
body = response.get_json()
|
||||
assert body["result"]["protocolVersion"] == module.mcp_types.SERVER_LATEST_PROTOCOL_VERSION
|
||||
|
||||
def test_initialize_ignores_unsupported_header(self, flask_app_with_containers):
|
||||
"""Initialize negotiates via the request body, so its header is never rejected."""
|
||||
api = self._make_api()
|
||||
|
||||
response = self._post(
|
||||
flask_app_with_containers,
|
||||
api,
|
||||
_initialize_payload("2024-11-05"),
|
||||
headers={"MCP-Protocol-Version": _UNSUPPORTED_VERSION},
|
||||
)
|
||||
|
||||
body = response.get_json()
|
||||
assert "error" not in body
|
||||
assert body["result"]["protocolVersion"] == "2024-11-05"
|
||||
|
||||
@pytest.mark.parametrize("request_id", [5, None])
|
||||
def test_unsupported_header_returns_invalid_request_error(self, flask_app_with_containers, request_id):
|
||||
"""An unsupported header gets a JSON-RPC error echoing the request id (missing id -> null)."""
|
||||
api = self._make_api()
|
||||
|
||||
with patch.object(module, "handle_mcp_request", autospec=True) as mock_handle:
|
||||
response = self._post(
|
||||
flask_app_with_containers,
|
||||
api,
|
||||
_tools_list_payload(request_id=request_id),
|
||||
headers={"MCP-Protocol-Version": _UNSUPPORTED_VERSION},
|
||||
)
|
||||
|
||||
body = response.get_json()
|
||||
assert response.status_code == 200
|
||||
assert body["jsonrpc"] == "2.0"
|
||||
assert body["id"] == request_id
|
||||
assert body["error"]["code"] == module.mcp_types.INVALID_REQUEST
|
||||
assert _UNSUPPORTED_VERSION in body["error"]["message"]
|
||||
mock_handle.assert_not_called()
|
||||
|
||||
def test_notification_with_unsupported_header_is_accepted(self, flask_app_with_containers):
|
||||
"""A notification is accepted (202, no body) even with an unsupported header."""
|
||||
api = self._make_api()
|
||||
|
||||
response = self._post(
|
||||
flask_app_with_containers,
|
||||
api,
|
||||
{"jsonrpc": "2.0", "method": "notifications/initialized", "params": {}},
|
||||
headers={"MCP-Protocol-Version": _UNSUPPORTED_VERSION},
|
||||
)
|
||||
|
||||
assert response.status_code == 202
|
||||
|
||||
@pytest.mark.parametrize("version", sorted(module.mcp_types.SERVER_SUPPORTED_PROTOCOL_VERSIONS))
|
||||
def test_supported_header_is_threaded_to_handler(self, flask_app_with_containers, version):
|
||||
"""Every supported header value is passed through to handle_mcp_request."""
|
||||
api = self._make_api()
|
||||
|
||||
with patch.object(module, "handle_mcp_request", return_value=DummyResult(), autospec=True) as mock_handle:
|
||||
self._post(
|
||||
flask_app_with_containers,
|
||||
api,
|
||||
_tools_list_payload(),
|
||||
headers={"MCP-Protocol-Version": version},
|
||||
)
|
||||
|
||||
assert mock_handle.call_args.args[-1] == version
|
||||
|
||||
def test_absent_header_defaults_to_back_compat_version(self, flask_app_with_containers):
|
||||
"""An absent header resolves to the spec's default version (2025-03-26)."""
|
||||
api = self._make_api()
|
||||
|
||||
with patch.object(module, "handle_mcp_request", return_value=DummyResult(), autospec=True) as mock_handle:
|
||||
self._post(flask_app_with_containers, api, _tools_list_payload())
|
||||
|
||||
assert mock_handle.call_args.args[-1] == module.mcp_types.DEFAULT_NEGOTIATED_VERSION
|
||||
|
||||
def test_tools_list_json_advertises_structured_output_for_modern_client(self, flask_app_with_containers):
|
||||
"""A 2025-06-18 client sees outputSchema and title in the serialized tool JSON."""
|
||||
api = self._make_api()
|
||||
|
||||
response = self._post(
|
||||
flask_app_with_containers,
|
||||
api,
|
||||
_tools_list_payload(),
|
||||
headers={"MCP-Protocol-Version": "2025-06-18"},
|
||||
)
|
||||
|
||||
tool = response.get_json()["result"]["tools"][0]
|
||||
assert tool["outputSchema"] == {"type": "object"}
|
||||
assert tool["title"] == "test_app"
|
||||
|
||||
def test_tools_list_json_unchanged_for_legacy_client(self, flask_app_with_containers):
|
||||
"""A 2024-11-05 client sees exactly the pre-upgrade tool JSON keys."""
|
||||
api = self._make_api()
|
||||
|
||||
response = self._post(
|
||||
flask_app_with_containers,
|
||||
api,
|
||||
_tools_list_payload(),
|
||||
headers={"MCP-Protocol-Version": "2024-11-05"},
|
||||
)
|
||||
|
||||
tool = response.get_json()["result"]["tools"][0]
|
||||
assert set(tool) == {"name", "description", "inputSchema"}
|
||||
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_tools_call_json_includes_structured_content_for_modern_client(
|
||||
self, mock_app_generate, flask_app_with_containers
|
||||
):
|
||||
"""A 2025-06-18 client receives structuredContent alongside the text content."""
|
||||
api = self._make_api()
|
||||
mock_app_generate.generate.return_value = {"answer": "test answer"}
|
||||
|
||||
response = self._post(
|
||||
flask_app_with_containers,
|
||||
api,
|
||||
_tools_call_payload(),
|
||||
headers={"MCP-Protocol-Version": "2025-06-18"},
|
||||
)
|
||||
|
||||
result = response.get_json()["result"]
|
||||
assert result["structuredContent"] == {"answer": "test answer"}
|
||||
assert result["content"][0]["text"] == "test answer"
|
||||
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_tools_call_json_omits_structured_content_for_legacy_client(
|
||||
self, mock_app_generate, flask_app_with_containers
|
||||
):
|
||||
"""A 2024-11-05 client receives the pre-upgrade tools/call JSON without structuredContent."""
|
||||
api = self._make_api()
|
||||
mock_app_generate.generate.return_value = {"answer": "test answer"}
|
||||
|
||||
response = self._post(
|
||||
flask_app_with_containers,
|
||||
api,
|
||||
_tools_call_payload(),
|
||||
headers={"MCP-Protocol-Version": "2024-11-05"},
|
||||
)
|
||||
|
||||
result = response.get_json()["result"]
|
||||
assert "structuredContent" not in result
|
||||
assert result["content"][0]["text"] == "test answer"
|
||||
|
||||
@ -10,11 +10,13 @@ from core.mcp.server.streamable_http import (
|
||||
build_parameter_schema,
|
||||
convert_input_form_to_parameters,
|
||||
extract_answer_from_response,
|
||||
extract_structured_output,
|
||||
handle_call_tool,
|
||||
handle_initialize,
|
||||
handle_list_tools,
|
||||
handle_mcp_request,
|
||||
handle_ping,
|
||||
negotiate_protocol_version,
|
||||
prepare_tool_arguments,
|
||||
process_mapping_response,
|
||||
)
|
||||
@ -64,6 +66,8 @@ class TestHandleMCPRequest:
|
||||
# Setup initialize request
|
||||
self.mock_request.root = Mock(spec=types.InitializeRequest)
|
||||
self.mock_request.root.id = 123
|
||||
self.mock_request.root.params = Mock()
|
||||
self.mock_request.root.params.protocolVersion = "2025-06-18"
|
||||
request_type = Mock(return_value=types.InitializeRequest)
|
||||
|
||||
with patch("core.mcp.server.streamable_http.type", request_type):
|
||||
@ -91,6 +95,33 @@ class TestHandleMCPRequest:
|
||||
assert result.jsonrpc == "2.0"
|
||||
assert result.id == 123
|
||||
|
||||
def test_handle_list_tools_request_threads_protocol_version(self):
|
||||
"""The negotiated version reaches handle_list_tools through the dispatcher."""
|
||||
self.mock_request.root = Mock(spec=types.ListToolsRequest)
|
||||
self.mock_request.root.id = 123
|
||||
|
||||
result = handle_mcp_request(
|
||||
Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123, "2025-06-18"
|
||||
)
|
||||
|
||||
assert isinstance(result, types.JSONRPCResponse)
|
||||
tool = result.result["tools"][0]
|
||||
assert tool["outputSchema"] == {"type": "object"}
|
||||
assert tool["title"] == "test_app"
|
||||
|
||||
def test_handle_list_tools_request_legacy_serialization_unchanged(self):
|
||||
"""A 2024-11-05 tools/list response serializes without any 2025-06-18 fields."""
|
||||
self.mock_request.root = Mock(spec=types.ListToolsRequest)
|
||||
self.mock_request.root.id = 123
|
||||
|
||||
result = handle_mcp_request(
|
||||
Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123, "2024-11-05"
|
||||
)
|
||||
|
||||
assert isinstance(result, types.JSONRPCResponse)
|
||||
tool = result.result["tools"][0]
|
||||
assert set(tool) == {"name", "description", "inputSchema"}
|
||||
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool_request(self, mock_app_generate):
|
||||
"""Test handling call tool request"""
|
||||
@ -119,6 +150,43 @@ class TestHandleMCPRequest:
|
||||
# Verify AppGenerateService was called
|
||||
mock_app_generate.generate.assert_called_once()
|
||||
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool_request_threads_protocol_version(self, mock_app_generate):
|
||||
"""The negotiated version reaches handle_call_tool through the dispatcher."""
|
||||
mock_call_request = Mock(spec=types.CallToolRequest)
|
||||
mock_call_request.params = Mock()
|
||||
mock_call_request.params.arguments = {"query": "test question"}
|
||||
mock_call_request.id = 123
|
||||
self.mock_request.root = mock_call_request
|
||||
|
||||
mock_app_generate.generate.return_value = {"answer": "test answer"}
|
||||
|
||||
result = handle_mcp_request(
|
||||
Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123, "2025-06-18"
|
||||
)
|
||||
|
||||
assert isinstance(result, types.JSONRPCResponse)
|
||||
assert result.result["structuredContent"] == {"answer": "test answer"}
|
||||
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool_request_legacy_serialization_unchanged(self, mock_app_generate):
|
||||
"""A 2024-11-05 tools/call response serializes without structuredContent."""
|
||||
mock_call_request = Mock(spec=types.CallToolRequest)
|
||||
mock_call_request.params = Mock()
|
||||
mock_call_request.params.arguments = {"query": "test question"}
|
||||
mock_call_request.id = 123
|
||||
self.mock_request.root = mock_call_request
|
||||
|
||||
mock_app_generate.generate.return_value = {"answer": "test answer"}
|
||||
|
||||
result = handle_mcp_request(
|
||||
Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123, "2024-11-05"
|
||||
)
|
||||
|
||||
assert isinstance(result, types.JSONRPCResponse)
|
||||
assert "structuredContent" not in result.result
|
||||
assert result.result["content"][0]["text"] == "test answer"
|
||||
|
||||
def test_handle_unknown_request_type(self):
|
||||
"""Test handling unknown request type"""
|
||||
|
||||
@ -183,18 +251,49 @@ class TestIndividualHandlers:
|
||||
result = handle_ping()
|
||||
assert isinstance(result, types.EmptyResult)
|
||||
|
||||
def test_handle_initialize(self):
|
||||
"""Test initialize handler"""
|
||||
description = "Test server"
|
||||
|
||||
def test_handle_initialize_echoes_supported_version(self):
|
||||
"""A supported requested version is echoed back unchanged."""
|
||||
with patch("core.mcp.server.streamable_http.dify_config") as mock_config:
|
||||
mock_config.project.version = "1.0.0"
|
||||
result = handle_initialize(description)
|
||||
result = handle_initialize("Test server", "2024-11-05")
|
||||
|
||||
assert isinstance(result, types.InitializeResult)
|
||||
assert result.protocolVersion == types.SERVER_LATEST_PROTOCOL_VERSION
|
||||
assert result.protocolVersion == "2024-11-05"
|
||||
assert result.instructions == "Test server"
|
||||
|
||||
def test_handle_initialize_echoes_intermediate_version(self):
|
||||
"""The intermediate supported version (2025-03-26) is echoed back."""
|
||||
with patch("core.mcp.server.streamable_http.dify_config") as mock_config:
|
||||
mock_config.project.version = "1.0.0"
|
||||
result = handle_initialize("Test server", "2025-03-26")
|
||||
|
||||
assert result.protocolVersion == "2025-03-26"
|
||||
|
||||
def test_handle_initialize_negotiates_latest_for_modern_client(self):
|
||||
"""A 2025-06-18 client gets 2025-06-18 back."""
|
||||
with patch("core.mcp.server.streamable_http.dify_config") as mock_config:
|
||||
mock_config.project.version = "1.0.0"
|
||||
result = handle_initialize("Test server", "2025-06-18")
|
||||
|
||||
assert result.protocolVersion == "2025-06-18"
|
||||
|
||||
def test_handle_initialize_falls_back_for_unknown_version(self):
|
||||
"""An unsupported requested version falls back to the server latest."""
|
||||
with patch("core.mcp.server.streamable_http.dify_config") as mock_config:
|
||||
mock_config.project.version = "1.0.0"
|
||||
result = handle_initialize("Test server", "1999-01-01")
|
||||
|
||||
assert result.protocolVersion == types.SERVER_LATEST_PROTOCOL_VERSION
|
||||
assert result.protocolVersion == "2025-06-18"
|
||||
|
||||
def test_handle_initialize_non_string_version_falls_back(self):
|
||||
"""A malformed (non-string) requested version falls back to the server latest."""
|
||||
with patch("core.mcp.server.streamable_http.dify_config") as mock_config:
|
||||
mock_config.project.version = "1.0.0"
|
||||
result = handle_initialize("Test server", 20250618)
|
||||
|
||||
assert result.protocolVersion == types.SERVER_LATEST_PROTOCOL_VERSION
|
||||
|
||||
def test_handle_list_tools(self):
|
||||
"""Test list tools handler"""
|
||||
app_name = "test_app"
|
||||
@ -210,6 +309,30 @@ class TestIndividualHandlers:
|
||||
assert result.tools[0].name == "test_app"
|
||||
assert result.tools[0].description == "Test server"
|
||||
|
||||
def test_handle_list_tools_adds_structured_output_for_modern_client(self):
|
||||
"""Tool advertises outputSchema and title when negotiated >= 2025-06-18."""
|
||||
result = handle_list_tools("test_app", AppMode.CHAT, [], "Test server", {}, "2025-06-18")
|
||||
|
||||
tool = result.tools[0]
|
||||
assert tool.outputSchema == {"type": "object"}
|
||||
assert tool.title == "test_app"
|
||||
|
||||
def test_handle_list_tools_omits_structured_output_for_legacy_client(self):
|
||||
"""Tool stays unchanged (no outputSchema/title) for 2024-11-05 clients."""
|
||||
result = handle_list_tools("test_app", AppMode.CHAT, [], "Test server", {}, "2024-11-05")
|
||||
|
||||
tool = result.tools[0]
|
||||
assert tool.outputSchema is None
|
||||
assert tool.title is None
|
||||
|
||||
def test_handle_list_tools_omits_structured_output_for_intermediate_client(self):
|
||||
"""The 2025-03-26 negotiated version is below the structured-output threshold."""
|
||||
result = handle_list_tools("test_app", AppMode.CHAT, [], "Test server", {}, "2025-03-26")
|
||||
|
||||
tool = result.tools[0]
|
||||
assert tool.outputSchema is None
|
||||
assert tool.title is None
|
||||
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool(self, mock_app_generate):
|
||||
"""Test call tool handler"""
|
||||
@ -239,6 +362,44 @@ class TestIndividualHandlers:
|
||||
assert hasattr(text_content, "text")
|
||||
assert text_content.text == "test answer"
|
||||
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool_structured_output_modern_client(self, mock_app_generate):
|
||||
"""structuredContent is attached alongside TextContent for >= 2025-06-18."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
|
||||
mock_request = Mock()
|
||||
mock_call_request = Mock(spec=types.CallToolRequest)
|
||||
mock_call_request.params = Mock()
|
||||
mock_call_request.params.arguments = {"query": "test question"}
|
||||
mock_request.root = mock_call_request
|
||||
|
||||
mock_app_generate.generate.return_value = {"answer": "test answer"}
|
||||
|
||||
result = handle_call_tool(Mock(), app, mock_request, [], Mock(spec=EndUser), "2025-06-18")
|
||||
|
||||
assert result.structuredContent == {"answer": "test answer"}
|
||||
assert result.content[0].text == "test answer"
|
||||
|
||||
@patch("core.mcp.server.streamable_http.AppGenerateService")
|
||||
def test_handle_call_tool_no_structured_output_legacy_client(self, mock_app_generate):
|
||||
"""structuredContent is omitted for 2024-11-05 clients."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
|
||||
mock_request = Mock()
|
||||
mock_call_request = Mock(spec=types.CallToolRequest)
|
||||
mock_call_request.params = Mock()
|
||||
mock_call_request.params.arguments = {"query": "test question"}
|
||||
mock_request.root = mock_call_request
|
||||
|
||||
mock_app_generate.generate.return_value = {"answer": "test answer"}
|
||||
|
||||
result = handle_call_tool(Mock(), app, mock_request, [], Mock(spec=EndUser), "2024-11-05")
|
||||
|
||||
assert result.structuredContent is None
|
||||
assert result.content[0].text == "test answer"
|
||||
|
||||
def test_handle_call_tool_no_end_user(self):
|
||||
"""Test call tool handler without end user"""
|
||||
app = Mock(spec=App)
|
||||
@ -375,6 +536,65 @@ class TestUtilityFunctions:
|
||||
|
||||
assert result == "thinking...more thinking"
|
||||
|
||||
def test_extract_structured_output_workflow(self):
|
||||
"""Workflow mode exposes the raw outputs mapping as structured content."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
|
||||
response = {"data": {"outputs": {"result": "test result"}}}
|
||||
|
||||
assert extract_structured_output(app, response, "ignored") == {"result": "test result"}
|
||||
|
||||
def test_extract_structured_output_chat(self):
|
||||
"""Chat mode wraps the answer string under an 'answer' key."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHAT
|
||||
|
||||
assert extract_structured_output(app, {"answer": "hi"}, "hi") == {"answer": "hi"}
|
||||
|
||||
def test_extract_structured_output_workflow_missing_outputs(self):
|
||||
"""Missing or malformed outputs fall back to None."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
|
||||
assert extract_structured_output(app, {"data": {}}, "ignored") is None
|
||||
|
||||
def test_extract_structured_output_workflow_non_mapping_response(self):
|
||||
"""A non-mapping workflow response yields no structured output."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
|
||||
assert extract_structured_output(app, None, "ignored") is None
|
||||
|
||||
def test_extract_structured_output_workflow_non_mapping_data(self):
|
||||
"""A non-mapping 'data' entry yields no structured output."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
|
||||
assert extract_structured_output(app, {"data": "not a mapping"}, "ignored") is None
|
||||
|
||||
def test_extract_structured_output_workflow_non_mapping_outputs(self):
|
||||
"""A non-mapping 'outputs' entry yields no structured output."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
|
||||
assert extract_structured_output(app, {"data": {"outputs": ["not", "a", "mapping"]}}, "ignored") is None
|
||||
|
||||
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.AGENT_CHAT, AppMode.COMPLETION])
|
||||
def test_extract_structured_output_other_answer_modes(self, mode):
|
||||
"""Every chat-style mode wraps the answer string under an 'answer' key."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = mode
|
||||
|
||||
assert extract_structured_output(app, {"answer": "hi"}, "hi") == {"answer": "hi"}
|
||||
|
||||
def test_extract_structured_output_unknown_mode(self):
|
||||
"""Modes outside the MCP surface produce no structured output."""
|
||||
app = Mock(spec=App)
|
||||
app.mode = AppMode.CHANNEL
|
||||
|
||||
assert extract_structured_output(app, {"answer": "hi"}, "hi") is None
|
||||
|
||||
def test_process_mapping_response_invalid_mode(self):
|
||||
"""Test processing mapping response with invalid app mode"""
|
||||
app = Mock(spec=App)
|
||||
@ -578,3 +798,29 @@ class TestUtilityFunctions:
|
||||
# Or validation should also raise SchemaError
|
||||
with pytest.raises(jsonschema.exceptions.SchemaError):
|
||||
jsonschema.validate(instance={"count": 1.23}, schema=bad_schema)
|
||||
|
||||
|
||||
class TestNegotiateProtocolVersion:
|
||||
"""Test the MCP-Protocol-Version header resolver."""
|
||||
|
||||
def test_initialize_ignores_header(self):
|
||||
"""Initialize negotiates via the request body, so its header is ignored."""
|
||||
assert negotiate_protocol_version("anything", True) == types.DEFAULT_NEGOTIATED_VERSION
|
||||
|
||||
def test_absent_header_defaults(self):
|
||||
"""An absent header defaults to 2025-03-26 per the spec back-compat rule."""
|
||||
assert negotiate_protocol_version(None, False) == types.DEFAULT_NEGOTIATED_VERSION
|
||||
|
||||
def test_empty_header_treated_as_absent(self):
|
||||
"""An empty header value is treated as absent and defaults to 2025-03-26."""
|
||||
assert negotiate_protocol_version("", False) == types.DEFAULT_NEGOTIATED_VERSION
|
||||
|
||||
def test_supported_header_passes_through(self):
|
||||
"""All supported header values are used as the negotiated version."""
|
||||
assert negotiate_protocol_version("2025-06-18", False) == "2025-06-18"
|
||||
assert negotiate_protocol_version("2025-03-26", False) == "2025-03-26"
|
||||
assert negotiate_protocol_version("2024-11-05", False) == "2024-11-05"
|
||||
|
||||
def test_unsupported_header_returns_none(self):
|
||||
"""An explicit but unsupported header signals an error (None)."""
|
||||
assert negotiate_protocol_version("1999-01-01", False) is None
|
||||
|
||||
@ -4,6 +4,7 @@ import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from core.mcp.types import (
|
||||
DEFAULT_NEGOTIATED_VERSION,
|
||||
INTERNAL_ERROR,
|
||||
INVALID_PARAMS,
|
||||
INVALID_REQUEST,
|
||||
@ -11,6 +12,7 @@ from core.mcp.types import (
|
||||
METHOD_NOT_FOUND,
|
||||
PARSE_ERROR,
|
||||
SERVER_LATEST_PROTOCOL_VERSION,
|
||||
SERVER_SUPPORTED_PROTOCOL_VERSIONS,
|
||||
Annotations,
|
||||
CallToolRequest,
|
||||
CallToolRequestParams,
|
||||
@ -59,7 +61,9 @@ class TestConstants:
|
||||
def test_protocol_versions(self):
|
||||
"""Test protocol version constants."""
|
||||
assert LATEST_PROTOCOL_VERSION == "2025-06-18"
|
||||
assert SERVER_LATEST_PROTOCOL_VERSION == "2024-11-05"
|
||||
assert SERVER_LATEST_PROTOCOL_VERSION == "2025-06-18"
|
||||
assert DEFAULT_NEGOTIATED_VERSION == "2025-03-26"
|
||||
assert sorted(SERVER_SUPPORTED_PROTOCOL_VERSIONS) == ["2024-11-05", "2025-03-26", "2025-06-18"]
|
||||
|
||||
def test_error_codes(self):
|
||||
"""Test JSON-RPC error code constants."""
|
||||
|
||||
Loading…
Reference in New Issue
Block a user