From eafe04eceae0435d1e7fbff4c30ec434fe5c27a0 Mon Sep 17 00:00:00 2001 From: harsh mahajan Date: Thu, 10 Sep 2026 17:52:05 +0530 Subject: [PATCH 1/4] fix: validate query string lists before SDK execution --- src/mcp_server_appwrite/server.py | 10 ++++ tests/unit/test_server.py | 96 +++++++++++++++++++++++++++++++ 2 files changed, 106 insertions(+) diff --git a/src/mcp_server_appwrite/server.py b/src/mcp_server_appwrite/server.py index c5ceb56..67d3bcc 100644 --- a/src/mcp_server_appwrite/server.py +++ b/src/mcp_server_appwrite/server.py @@ -686,6 +686,16 @@ def _coerce_argument(param_name: str, value: Any, param_type: Any) -> Any: origin = get_origin(param_type) args = get_args(param_type) + if param_name == "queries" and origin is list: + if not isinstance(value, list) or any( + not isinstance(query, str) for query in value + ): + raise ValueError( + "'queries' must be an array of JSON strings, not query objects. " + "JSON-encode each query object before adding it to the array." + ) + return value + if param_type is InputFile: return _coerce_input_file(value, param_name) diff --git a/tests/unit/test_server.py b/tests/unit/test_server.py index 1689bcc..03569be 100644 --- a/tests/unit/test_server.py +++ b/tests/unit/test_server.py @@ -16,6 +16,7 @@ from appwrite_console.exception import AppwriteException from appwrite_console.input_file import InputFile from appwrite_console.models.row_list import RowList +from requests import Response from mcp_server_appwrite import server as server_module from mcp_server_appwrite.catalog_policy import API_KEY_PROFILE, OAUTH_PROFILE @@ -444,6 +445,101 @@ def test_prepare_arguments_rejects_unsupported_copied_response_fields(self): }, ) + def test_call_tool_rejects_unencoded_queries_before_sdk_execution(self): + client = build_introspection_client() + manager = register_services(client, profile=API_KEY_PROFILE) + server = build_mcp_server(build_operator(manager, client), transport="stdio") + entry = server.get_request_handler("tools/call") + self.assertIsNotNone(entry) + + async def run_check(): + ctx = Mock() + ctx.protocol_version = "2026-07-28" + ctx.meta = None + ctx.session.client_params = None + for queries in ( + [{"method": "limit", "values": [1]}], + [ + '{"method":"limit","values":[1]}', + {"method": "offset", "values": [1]}, + ], + {"method": "limit", "values": [1]}, + '{"method":"limit","values":[1]}', + [1], + ): + with self.subTest(queries=queries): + arguments = { + "database_id": "database-id", + "table_id": "table-id", + "queries": queries, + } + with patch("appwrite_console.client.requests.request") as request: + with self.assertRaises(ValueError) as error: + execute_registered_tool( + manager, + "tables_db_list_rows", + arguments, + client=client, + ) + request.assert_not_called() + self.assertIn("queries", str(error.exception)) + params = types.CallToolRequestParams( + name="appwrite_call_tool", + arguments={ + "tool_name": "tables_db_list_rows", + "arguments": arguments, + }, + ) + with patch("appwrite_console.client.requests.request") as request: + result = await entry.handler(ctx, params) + + self.assertTrue(result.is_error) + self.assertIn("queries", result.content[0].text) + self.assertIn("JSON", result.content[0].text) + request.assert_not_called() + + asyncio.run(run_check()) + + def test_execute_sdk_tool_preserves_encoded_queries(self): + client = build_introspection_client() + manager = register_services(client, profile=API_KEY_PROFILE) + response = Response() + response.status_code = 200 + response.headers["Content-Type"] = "application/json" + response._content = b'{"total":0,"rows":[]}' + query = '{"method":"equal","attribute":"name","values":["Zoë"]}' + + for queries, expected in ( + ([query], {"queries[0]": query}), + ([], {}), + (None, {}), + ): + with ( + self.subTest(queries=queries), + patch( + "appwrite_console.client.requests.request", return_value=response + ) as request, + ): + result = execute_registered_tool( + manager, + "tables_db_list_rows", + { + "database_id": "database-id", + "table_id": "table-id", + "queries": queries, + }, + client=client, + ) + + request.assert_called_once() + self.assertTrue( + request.call_args.kwargs["url"].endswith( + "/tablesdb/database-id/tables/table-id/rows" + ) + ) + self.assertEqual(request.call_args.kwargs["params"], expected) + self.assertEqual(json.loads(result[0].text), {"total": 0, "rows": []}) + def test_format_tool_result_serializes_json(self): result = _format_tool_result( "tables_db_list_rows", {"total": 1, "rows": []}, {} From 15cd7b887d2dd21b1a75ed02762c8b1424e2c048 Mon Sep 17 00:00:00 2001 From: harsh mahajan Date: Thu, 10 Sep 2026 18:04:16 +0530 Subject: [PATCH 2/4] test: cover query validation through public MCP calls --- tests/unit/test_server.py | 111 ++++++++++++++++++++------------------ 1 file changed, 58 insertions(+), 53 deletions(-) diff --git a/tests/unit/test_server.py b/tests/unit/test_server.py index 03569be..3dfd9e7 100644 --- a/tests/unit/test_server.py +++ b/tests/unit/test_server.py @@ -16,7 +16,6 @@ from appwrite_console.exception import AppwriteException from appwrite_console.input_file import InputFile from appwrite_console.models.row_list import RowList -from requests import Response from mcp_server_appwrite import server as server_module from mcp_server_appwrite.catalog_policy import API_KEY_PROFILE, OAUTH_PROFILE @@ -452,6 +451,13 @@ def test_call_tool_rejects_unencoded_queries_before_sdk_execution(self): entry = server.get_request_handler("tools/call") self.assertIsNotNone(entry) + class TablesDbService: + def __init__(self, client): + pass + + def list_rows(self, **arguments): + return {"total": 0, "rows": []} + async def run_check(): ctx = Mock() ctx.protocol_version = "2026-07-28" @@ -468,77 +474,76 @@ async def run_check(): [1], ): with self.subTest(queries=queries): - arguments = { - "database_id": "database-id", - "table_id": "table-id", - "queries": queries, - } - with patch("appwrite_console.client.requests.request") as request: - with self.assertRaises(ValueError) as error: - execute_registered_tool( - manager, - "tables_db_list_rows", - arguments, - client=client, - ) - request.assert_not_called() - self.assertIn("queries", str(error.exception)) params = types.CallToolRequestParams( name="appwrite_call_tool", arguments={ "tool_name": "tables_db_list_rows", - "arguments": arguments, + "arguments": { + "database_id": "database-id", + "table_id": "table-id", + "queries": queries, + }, }, ) - with patch("appwrite_console.client.requests.request") as request: - result = await entry.handler(ctx, params) + result = await entry.handler(ctx, params) self.assertTrue(result.is_error) self.assertIn("queries", result.content[0].text) self.assertIn("JSON", result.content[0].text) - request.assert_not_called() - asyncio.run(run_check()) + with patch.dict(server_module.SERVICE_CLASSES, {"tables_db": TablesDbService}): + asyncio.run(run_check()) - def test_execute_sdk_tool_preserves_encoded_queries(self): + def test_call_tool_preserves_encoded_queries(self): client = build_introspection_client() manager = register_services(client, profile=API_KEY_PROFILE) - response = Response() - response.status_code = 200 - response.headers["Content-Type"] = "application/json" - response._content = b'{"total":0,"rows":[]}' + server = build_mcp_server(build_operator(manager, client), transport="stdio") + entry = server.get_request_handler("tools/call") + self.assertIsNotNone(entry) query = '{"method":"equal","attribute":"name","values":["Zoë"]}' + received_queries = [] - for queries, expected in ( - ([query], {"queries[0]": query}), - ([], {}), - (None, {}), - ): - with ( - self.subTest(queries=queries), - patch( - "appwrite_console.client.requests.request", return_value=response - ) as request, + class TablesDbService: + def __init__(self, client): + pass + + def list_rows(self, database_id, table_id, queries=None): + received_queries.append(queries) + return {"total": 0, "rows": []} + + async def run_check(): + ctx = Mock() + ctx.protocol_version = "2026-07-28" + ctx.meta = None + ctx.session.client_params = None + for arguments, expected in ( + ({"queries": [query]}, [query]), + ({"queries": []}, []), + ({"queries": None}, None), + ({}, None), ): - result = execute_registered_tool( - manager, - "tables_db_list_rows", - { - "database_id": "database-id", - "table_id": "table-id", - "queries": queries, - }, - client=client, - ) + with self.subTest(arguments=arguments): + params = types.CallToolRequestParams( + name="appwrite_call_tool", + arguments={ + "tool_name": "tables_db_list_rows", + "arguments": { + "database_id": "database-id", + "table_id": "table-id", + **arguments, + }, + }, + ) + result = await entry.handler(ctx, params) - request.assert_called_once() - self.assertTrue( - request.call_args.kwargs["url"].endswith( - "/tablesdb/database-id/tables/table-id/rows" + self.assertFalse(result.is_error) + self.assertEqual(received_queries[-1], expected) + self.assertEqual( + json.loads(result.content[0].text), {"total": 0, "rows": []} ) - ) - self.assertEqual(request.call_args.kwargs["params"], expected) - self.assertEqual(json.loads(result[0].text), {"total": 0, "rows": []}) + + with patch.dict(server_module.SERVICE_CLASSES, {"tables_db": TablesDbService}): + asyncio.run(run_check()) def test_format_tool_result_serializes_json(self): result = _format_tool_result( From 29eab193d46c09770e2ae384e720a454299e9d3a Mon Sep 17 00:00:00 2001 From: harsh mahajan Date: Fri, 11 Sep 2026 15:01:22 +0530 Subject: [PATCH 3/4] test: cover upstream SDK query error handling --- src/mcp_server_appwrite/server.py | 10 ------ tests/unit/test_server.py | 54 +++++++++++++++++++++++++------ 2 files changed, 44 insertions(+), 20 deletions(-) diff --git a/src/mcp_server_appwrite/server.py b/src/mcp_server_appwrite/server.py index 4f711b8..05e0150 100644 --- a/src/mcp_server_appwrite/server.py +++ b/src/mcp_server_appwrite/server.py @@ -686,16 +686,6 @@ def _coerce_argument(param_name: str, value: Any, param_type: Any) -> Any: origin = get_origin(param_type) args = get_args(param_type) - if param_name == "queries" and origin is list: - if not isinstance(value, list) or any( - not isinstance(query, str) for query in value - ): - raise ValueError( - "'queries' must be an array of JSON strings, not query objects. " - "JSON-encode each query object before adding it to the array." - ) - return value - if param_type is InputFile: return _coerce_input_file(value, param_name) diff --git a/tests/unit/test_server.py b/tests/unit/test_server.py index 954e167..5f18fd2 100644 --- a/tests/unit/test_server.py +++ b/tests/unit/test_server.py @@ -5,8 +5,10 @@ import os import sys import tempfile +import threading import time import unittest +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from unittest.mock import Mock, patch @@ -100,6 +102,34 @@ def stream(self, method, url, **kwargs): return _FakeStream(self._response) +class _AppwriteServer(ThreadingHTTPServer): + def __init__(self, status, payload): + body = json.dumps(payload).encode() + + class Handler(BaseHTTPRequestHandler): + def do_GET(self): + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, *args): + pass + + super().__init__(("127.0.0.1", 0), Handler) + self.thread = threading.Thread(target=self.serve_forever, daemon=True) + + def __enter__(self): + self.thread.start() + return self + + def __exit__(self, *args): + self.shutdown() + self.thread.join() + super().__exit__(*args) + + class ServerHelperTests(unittest.TestCase): def test_parse_args_defaults_to_stdio(self): with patch.dict(os.environ, {}, clear=True): @@ -537,20 +567,13 @@ async def run_check(): with patch.dict(server_module.SERVICE_CLASSES, {"functions": FunctionsService}): asyncio.run(run_check()) - def test_call_tool_rejects_unencoded_queries_before_sdk_execution(self): + def test_call_tool_returns_api_error_for_unencoded_queries(self): client = build_introspection_client() manager = register_services(client, profile=API_KEY_PROFILE) server = build_mcp_server(build_operator(manager, client), transport="stdio") entry = server.get_request_handler("tools/call") self.assertIsNotNone(entry) - class TablesDbService: - def __init__(self, client): - pass - - def list_rows(self, **arguments): - return {"total": 0, "rows": []} - async def run_check(): ctx = Mock() ctx.protocol_version = "2026-07-28" @@ -581,10 +604,21 @@ async def run_check(): result = await entry.handler(ctx, params) self.assertTrue(result.is_error) + self.assertIn("code=400", result.content[0].text) + self.assertIn( + "type=general_argument_invalid", result.content[0].text + ) self.assertIn("queries", result.content[0].text) - self.assertIn("JSON", result.content[0].text) - with patch.dict(server_module.SERVICE_CLASSES, {"tables_db": TablesDbService}): + with _AppwriteServer( + 400, + { + "message": "Invalid queries: query values must be strings", + "code": 400, + "type": "general_argument_invalid", + }, + ) as api: + client.set_endpoint(f"http://127.0.0.1:{api.server_port}/v1") asyncio.run(run_check()) def test_call_tool_preserves_encoded_queries(self): From eabc50d7f679461bcee9dc0de7f74cbfa619f8c1 Mon Sep 17 00:00:00 2001 From: harsh mahajan Date: Fri, 11 Sep 2026 15:16:38 +0530 Subject: [PATCH 4/4] fix: classify generated SDK input validation errors --- .../error_classification.py | 8 +++ src/mcp_server_appwrite/error_monitoring.py | 2 +- tests/unit/test_error_classification.py | 27 ++++++++- tests/unit/test_error_monitoring.py | 43 ++++++++++++- tests/unit/test_server.py | 60 +++++-------------- 5 files changed, 92 insertions(+), 48 deletions(-) diff --git a/src/mcp_server_appwrite/error_classification.py b/src/mcp_server_appwrite/error_classification.py index 66d1141..5289fbb 100644 --- a/src/mcp_server_appwrite/error_classification.py +++ b/src/mcp_server_appwrite/error_classification.py @@ -17,6 +17,7 @@ "write_confirmation", "appwrite_4xx", "appwrite_5xx", + "sdk_input_validation", "sdk_validation", "response_too_large", "internal", @@ -27,6 +28,7 @@ "write_confirmation", "appwrite_4xx", "appwrite_5xx", + "sdk_input_validation", "sdk_validation", "response_too_large", "internal", @@ -83,6 +85,12 @@ def classify_tool_error(exc: BaseException) -> ErrorCategory: ) if appwrite_error is not None: code = _appwrite_status_code(appwrite_error) + if ( + code == 0 + and appwrite_error.type == "sdk_input_validation" + and appwrite_error.response is None + ): + return "sdk_input_validation" if code is not None and 400 <= code < 500: return "appwrite_4xx" if code is not None and 500 <= code < 600: diff --git a/src/mcp_server_appwrite/error_monitoring.py b/src/mcp_server_appwrite/error_monitoring.py index d04c066..cd72237 100644 --- a/src/mcp_server_appwrite/error_monitoring.py +++ b/src/mcp_server_appwrite/error_monitoring.py @@ -171,7 +171,7 @@ def _should_capture(exc: BaseException) -> bool: return False category = classify_tool_error(exc) - if category in {"write_confirmation", "appwrite_4xx"}: + if category in {"write_confirmation", "appwrite_4xx", "sdk_input_validation"}: return False # Pydantic validation errors are ValueError subclasses, but SDK response # validation is actionable model drift and must remain visible. diff --git a/tests/unit/test_error_classification.py b/tests/unit/test_error_classification.py index 6d657d4..fc6be25 100644 --- a/tests/unit/test_error_classification.py +++ b/tests/unit/test_error_classification.py @@ -41,6 +41,31 @@ def test_appwrite_5xx(self): "appwrite_5xx", ) + def test_sdk_input_validation(self): + error = AppwriteException( + 'Invalid parameter: "filters" must be an array of strings', + type="sdk_input_validation", + ) + wrapped = RuntimeError("wrapped") + wrapped.__cause__ = error + + for failure in (error, wrapped): + with self.subTest(failure=type(failure).__name__): + self.assertEqual(classify_tool_error(failure), "sdk_input_validation") + + def test_sdk_input_validation_type_does_not_hide_responses(self): + for code, response, expected in ( + (503, None, "appwrite_5xx"), + (0, {}, "internal"), + (0, "", "internal"), + ): + with self.subTest(code=code, response=response): + error = AppwriteException( + "upstream failed", code, "sdk_input_validation", response + ) + + self.assertEqual(classify_tool_error(error), expected) + def test_sdk_validation_takes_precedence_over_code(self): class Provider(BaseModel): options: dict @@ -49,7 +74,7 @@ class Provider(BaseModel): Provider.model_validate({"options": []}) except ValidationError as validation_error: appwrite_error = AppwriteException( - "Unable to parse response into Provider", 0, None + "Unable to parse response into Provider", 0, "sdk_input_validation" ) appwrite_error.__cause__ = validation_error else: # pragma: no cover - defensive diff --git a/tests/unit/test_error_monitoring.py b/tests/unit/test_error_monitoring.py index 5f67368..587028a 100644 --- a/tests/unit/test_error_monitoring.py +++ b/tests/unit/test_error_monitoring.py @@ -70,16 +70,55 @@ def test_wrapped_value_errors_are_not_captured(self): capture.assert_not_called() def test_wrapped_sdk_validation_errors_are_captured(self): + error_monitoring._enabled = True + class Payload(BaseModel): required: str try: Payload.model_validate({}) except ValidationError as exc: + error = AppwriteException("invalid response", type="sdk_input_validation") + error.__cause__ = exc wrapped = RuntimeError("wrapped") - wrapped.__cause__ = exc + wrapped.__cause__ = error + + with patch("sentry_sdk.capture_exception") as capture: + captured = error_monitoring.capture_exception(wrapped) + + self.assertTrue(captured) + capture.assert_called_once_with(wrapped) + + def test_sdk_input_validation_is_not_captured(self): + error_monitoring._enabled = True + + for wrapped in (False, True): + with self.subTest(wrapped=wrapped): + error = AppwriteException( + "invalid filters", type="sdk_input_validation" + ) + failure = RuntimeError("wrapped") if wrapped else error + if wrapped: + failure.__cause__ = error + with patch("sentry_sdk.capture_exception") as capture: + captured = error_monitoring.capture_exception(failure) + + self.assertFalse(captured) + capture.assert_not_called() + + def test_sdk_input_validation_does_not_hide_unexpected_failures(self): + error_monitoring._enabled = True + for error in ( + AppwriteException("network down"), + AppwriteException("upstream failed", 503, "sdk_input_validation"), + AppwriteException("invalid response", 0, "sdk_input_validation", {}), + ): + with self.subTest(error=error): + with patch("sentry_sdk.capture_exception") as capture: + captured = error_monitoring.capture_exception(error) - self.assertTrue(error_monitoring._should_capture(wrapped)) + self.assertTrue(captured) + capture.assert_called_once_with(error) def test_client_disconnects_are_not_captured(self): error_monitoring._enabled = True diff --git a/tests/unit/test_server.py b/tests/unit/test_server.py index 5f18fd2..95db322 100644 --- a/tests/unit/test_server.py +++ b/tests/unit/test_server.py @@ -5,10 +5,8 @@ import os import sys import tempfile -import threading import time import unittest -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from unittest.mock import Mock, patch @@ -102,34 +100,6 @@ def stream(self, method, url, **kwargs): return _FakeStream(self._response) -class _AppwriteServer(ThreadingHTTPServer): - def __init__(self, status, payload): - body = json.dumps(payload).encode() - - class Handler(BaseHTTPRequestHandler): - def do_GET(self): - self.send_response(status) - self.send_header("Content-Type", "application/json") - self.send_header("Content-Length", str(len(body))) - self.end_headers() - self.wfile.write(body) - - def log_message(self, *args): - pass - - super().__init__(("127.0.0.1", 0), Handler) - self.thread = threading.Thread(target=self.serve_forever, daemon=True) - - def __enter__(self): - self.thread.start() - return self - - def __exit__(self, *args): - self.shutdown() - self.thread.join() - super().__exit__(*args) - - class ServerHelperTests(unittest.TestCase): def test_parse_args_defaults_to_stdio(self): with patch.dict(os.environ, {}, clear=True): @@ -567,7 +537,7 @@ async def run_check(): with patch.dict(server_module.SERVICE_CLASSES, {"functions": FunctionsService}): asyncio.run(run_check()) - def test_call_tool_returns_api_error_for_unencoded_queries(self): + def test_call_tool_returns_sdk_input_validation_for_unencoded_queries(self): client = build_introspection_client() manager = register_services(client, profile=API_KEY_PROFILE) server = build_mcp_server(build_operator(manager, client), transport="stdio") @@ -604,23 +574,25 @@ async def run_check(): result = await entry.handler(ctx, params) self.assertTrue(result.is_error) - self.assertIn("code=400", result.content[0].text) - self.assertIn( - "type=general_argument_invalid", result.content[0].text - ) + self.assertIn("type=sdk_input_validation", result.content[0].text) self.assertIn("queries", result.content[0].text) + self.assertIn("string", result.content[0].text) - with _AppwriteServer( - 400, - { - "message": "Invalid queries: query values must be strings", - "code": 400, - "type": "general_argument_invalid", - }, - ) as api: - client.set_endpoint(f"http://127.0.0.1:{api.server_port}/v1") + with ( + patch( + "requests.sessions.Session.request", + side_effect=AssertionError( + "Invalid inputs must not send HTTP requests" + ), + ) as request, + patch.object(server_module.error_monitoring, "_enabled", True), + patch("sentry_sdk.capture_exception") as capture, + ): asyncio.run(run_check()) + request.assert_not_called() + capture.assert_not_called() + def test_call_tool_preserves_encoded_queries(self): client = build_introspection_client() manager = register_services(client, profile=API_KEY_PROFILE)