1use std::sync::Arc;
29
30use anyhow::anyhow;
31use axum::Extension;
32use axum::Json;
33use axum::response::IntoResponse;
34use http::{HeaderMap, HeaderValue, StatusCode};
35use mz_adapter_types::dyncfgs::{
36 ENABLE_MCP_AGENT, ENABLE_MCP_AGENT_QUERY_TOOL, ENABLE_MCP_AGENT_READ_DATA_PRODUCT_TOOL,
37 ENABLE_MCP_DEVELOPER, ENABLE_MCP_DEVELOPER_QUERY_TOOL, MCP_MAX_RESPONSE_SIZE,
38 MCP_REQUEST_TIMEOUT,
39};
40use mz_ore::cast::CastLossy;
41use mz_repr::namespaces::{self, SYSTEM_SCHEMAS};
42use mz_sql::parse::{parse_item_name_with_limit, parse_with_limit};
43use mz_sql::session::metadata::SessionMetadata;
44use mz_sql::session::vars::{APPLICATION_NAME, Var, VarInput};
45use mz_sql_parser::ast::display::{AstDisplay, escaped_string_literal};
46use mz_sql_parser::ast::visit::{self, Visit};
47use mz_sql_parser::ast::{Raw, RawItemName};
48use serde::{Deserialize, Serialize};
49use serde_json::json;
50use thiserror::Error;
51use tracing::{debug, warn};
52
53use crate::http::AuthedClient;
54use crate::http::mcp_metrics::{McpCallStatus, McpMetrics, ToolCallGuard};
55use crate::http::sql::{SqlRequest, SqlResponse, SqlResult, execute_request};
56
57const JSONRPC_VERSION: &str = "2.0";
61
62const MCP_PROTOCOL_VERSION: &str = "2025-11-25";
65
66const DISCOVERY_QUERY: &str = "SELECT * FROM mz_internal.mz_mcp_data_products";
68const DETAILS_QUERY_PREFIX: &str =
69 "SELECT * FROM mz_internal.mz_mcp_data_product_details WHERE object_name = ";
70
71#[derive(Debug, Error)]
73enum McpRequestError {
74 #[error("Invalid JSON-RPC version: expected 2.0")]
75 InvalidJsonRpcVersion,
76 #[error("Method not found: {0}")]
77 MethodNotFound(String),
78 #[error("Tool not found: {0}")]
79 ToolNotFound(String),
80 #[error("Data product not found: {0}")]
81 DataProductNotFound(String),
82 #[error("Query validation failed: {0}")]
83 QueryValidationFailed(String),
84 #[error("Query execution failed: {0}")]
85 QueryExecutionFailed(String),
86 #[error("Internal error: {0}")]
87 Internal(#[from] anyhow::Error),
88}
89
90impl McpRequestError {
91 fn error_code(&self) -> i32 {
92 match self {
93 Self::InvalidJsonRpcVersion => error_codes::INVALID_REQUEST,
94 Self::MethodNotFound(_) => error_codes::METHOD_NOT_FOUND,
95 Self::ToolNotFound(_) => error_codes::INVALID_PARAMS,
96 Self::DataProductNotFound(_) => error_codes::INVALID_PARAMS,
97 Self::QueryValidationFailed(_) => error_codes::INVALID_PARAMS,
98 Self::QueryExecutionFailed(_) | Self::Internal(_) => error_codes::INTERNAL_ERROR,
99 }
100 }
101
102 fn error_type(&self) -> &'static str {
103 match self {
104 Self::InvalidJsonRpcVersion => "InvalidRequest",
105 Self::MethodNotFound(_) => "MethodNotFound",
106 Self::ToolNotFound(_) => "ToolNotFound",
107 Self::DataProductNotFound(_) => "DataProductNotFound",
108 Self::QueryValidationFailed(_) => "ValidationError",
109 Self::QueryExecutionFailed(_) => "ExecutionError",
110 Self::Internal(_) => "InternalError",
111 }
112 }
113}
114
115#[derive(Debug, Deserialize)]
117pub(crate) struct McpRequest {
118 jsonrpc: String,
119 id: Option<serde_json::Value>,
120 #[serde(flatten)]
121 method: McpMethod,
122}
123
124#[derive(Debug, Deserialize)]
126#[serde(tag = "method", content = "params")]
127enum McpMethod {
128 #[serde(rename = "initialize")]
130 Initialize(#[allow(dead_code)] InitializeParams),
131 #[serde(rename = "tools/list")]
134 ToolsList(#[allow(dead_code)] Option<serde_json::Value>),
135 #[serde(rename = "tools/call")]
136 ToolsCall(#[serde(deserialize_with = "deserialize_tools_call")] ToolsCallParams),
137 #[serde(rename = "ping")]
141 Ping(#[allow(dead_code)] Option<serde_json::Value>),
142 #[serde(rename = "notifications/initialized")]
143 NotificationsInitialized(#[allow(dead_code)] Option<serde_json::Value>),
144 #[serde(other)]
146 Unknown,
147}
148
149impl std::fmt::Display for McpMethod {
150 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
151 match self {
152 McpMethod::Initialize(_) => write!(f, "initialize"),
153 McpMethod::ToolsList(_) => write!(f, "tools/list"),
154 McpMethod::ToolsCall(_) => write!(f, "tools/call"),
155 McpMethod::Ping(_) => write!(f, "ping"),
156 McpMethod::NotificationsInitialized(_) => write!(f, "notifications/initialized"),
157 McpMethod::Unknown => write!(f, "unknown"),
158 }
159 }
160}
161
162#[derive(Debug, Deserialize)]
163struct InitializeParams {
164 #[serde(rename = "protocolVersion")]
166 #[allow(dead_code)]
167 protocol_version: String,
168 #[serde(default)]
170 #[allow(dead_code)]
171 capabilities: serde_json::Value,
172 #[serde(rename = "clientInfo")]
174 #[allow(dead_code)]
175 client_info: Option<ClientInfo>,
176}
177
178#[derive(Debug, Deserialize)]
179struct ClientInfo {
180 #[allow(dead_code)]
181 name: String,
182 #[allow(dead_code)]
183 version: String,
184}
185
186#[derive(Debug, Deserialize)]
189#[serde(tag = "name", content = "arguments")]
190#[serde(rename_all = "snake_case")]
191enum ToolsCallParams {
192 GetDataProducts(NoArguments),
194 GetDataProductDetails(GetDataProductDetailsParams),
195 ReadDataProduct(ReadDataProductParams),
196 Query(QueryParams),
197 QuerySystemCatalog(QuerySystemCatalogParams),
199}
200
201#[derive(Debug, Deserialize)]
206#[serde(deny_unknown_fields)]
207struct NoArguments {}
208
209fn deserialize_tools_call<'de, D>(deserializer: D) -> Result<ToolsCallParams, D::Error>
212where
213 D: serde::Deserializer<'de>,
214{
215 let mut params = serde_json::Value::deserialize(deserializer)?;
216 if let serde_json::Value::Object(fields) = &mut params {
217 let arguments = fields.entry("arguments").or_insert(serde_json::Value::Null);
218 if arguments.is_null() {
219 *arguments = json!({});
220 }
221 }
222 serde_json::from_value(params).map_err(serde::de::Error::custom)
223}
224
225impl std::fmt::Display for ToolsCallParams {
226 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
227 match self {
228 ToolsCallParams::GetDataProducts(_) => write!(f, "get_data_products"),
229 ToolsCallParams::GetDataProductDetails(_) => write!(f, "get_data_product_details"),
230 ToolsCallParams::ReadDataProduct(_) => write!(f, "read_data_product"),
231 ToolsCallParams::Query(_) => write!(f, "query"),
232 ToolsCallParams::QuerySystemCatalog(_) => write!(f, "query_system_catalog"),
233 }
234 }
235}
236
237#[derive(Debug, Deserialize)]
238struct GetDataProductDetailsParams {
239 name: String,
240}
241
242#[derive(Debug, Deserialize)]
243struct ReadDataProductParams {
244 name: String,
245 #[serde(default = "default_read_limit")]
246 limit: u32,
247 cluster: Option<String>,
248}
249
250const DEFAULT_READ_LIMIT: u32 = 500;
252
253fn default_read_limit() -> u32 {
254 DEFAULT_READ_LIMIT
255}
256
257#[derive(Debug, Deserialize)]
258struct QueryParams {
259 cluster: String,
260 cluster_replica: Option<String>,
263 sql_query: String,
264}
265
266#[derive(Debug, Deserialize)]
267struct QuerySystemCatalogParams {
268 sql_query: String,
269}
270
271#[derive(Debug, Serialize)]
272struct McpResponse {
273 jsonrpc: String,
274 id: serde_json::Value,
275 #[serde(skip_serializing_if = "Option::is_none")]
276 result: Option<McpResult>,
277 #[serde(skip_serializing_if = "Option::is_none")]
278 error: Option<McpError>,
279}
280
281impl McpResponse {
282 fn success(id: serde_json::Value, result: McpResult) -> Self {
284 Self {
285 jsonrpc: JSONRPC_VERSION.to_string(),
286 id,
287 result: Some(result),
288 error: None,
289 }
290 }
291
292 fn error(id: serde_json::Value, error: McpError) -> Self {
294 Self {
295 jsonrpc: JSONRPC_VERSION.to_string(),
296 id,
297 result: None,
298 error: Some(error),
299 }
300 }
301}
302
303#[derive(Debug, Serialize)]
305#[serde(untagged)]
306enum McpResult {
307 Initialize(InitializeResult),
308 ToolsList(ToolsListResult),
309 ToolContent(ToolContentResult),
310}
311
312#[derive(Debug, Serialize)]
313struct InitializeResult {
314 #[serde(rename = "protocolVersion")]
315 protocol_version: String,
316 capabilities: Capabilities,
317 #[serde(rename = "serverInfo")]
318 server_info: ServerInfo,
319 #[serde(skip_serializing_if = "Option::is_none")]
320 instructions: Option<String>,
321}
322
323#[derive(Debug, Serialize)]
324struct Capabilities {
325 tools: serde_json::Value,
326}
327
328#[derive(Debug, Serialize)]
329struct ServerInfo {
330 name: String,
331 version: String,
332}
333
334#[derive(Debug, Serialize)]
335struct ToolsListResult {
336 tools: Vec<ToolDefinition>,
337}
338
339#[derive(Debug, Serialize)]
340struct ToolDefinition {
341 name: String,
342 #[serde(skip_serializing_if = "Option::is_none")]
343 title: Option<String>,
344 description: String,
345 #[serde(rename = "inputSchema")]
346 input_schema: serde_json::Value,
347 #[serde(skip_serializing_if = "Option::is_none")]
348 annotations: Option<ToolAnnotations>,
349}
350
351#[derive(Debug, Serialize)]
354struct ToolAnnotations {
355 #[serde(rename = "readOnlyHint", skip_serializing_if = "Option::is_none")]
356 read_only_hint: Option<bool>,
357 #[serde(rename = "destructiveHint", skip_serializing_if = "Option::is_none")]
358 destructive_hint: Option<bool>,
359 #[serde(rename = "idempotentHint", skip_serializing_if = "Option::is_none")]
360 idempotent_hint: Option<bool>,
361 #[serde(rename = "openWorldHint", skip_serializing_if = "Option::is_none")]
362 open_world_hint: Option<bool>,
363}
364
365const READ_ONLY_ANNOTATIONS: ToolAnnotations = ToolAnnotations {
367 read_only_hint: Some(true),
368 destructive_hint: Some(false),
369 idempotent_hint: Some(true),
370 open_world_hint: Some(false),
371};
372
373#[derive(Debug, Serialize)]
374struct ToolContentResult {
375 content: Vec<ContentBlock>,
376 #[serde(rename = "isError")]
378 is_error: bool,
379}
380
381#[derive(Debug, Serialize)]
382struct ContentBlock {
383 #[serde(rename = "type")]
384 content_type: String,
385 text: String,
386}
387
388mod error_codes {
390 pub const INVALID_REQUEST: i32 = -32600;
391 pub const METHOD_NOT_FOUND: i32 = -32601;
392 pub const INVALID_PARAMS: i32 = -32602;
393 pub const INTERNAL_ERROR: i32 = -32603;
394}
395
396#[derive(Debug, Serialize)]
397struct McpError {
398 code: i32,
399 message: String,
400 #[serde(skip_serializing_if = "Option::is_none")]
401 data: Option<serde_json::Value>,
402}
403
404impl From<McpRequestError> for McpError {
405 fn from(err: McpRequestError) -> Self {
406 McpError {
407 code: err.error_code(),
408 message: err.to_string(),
409 data: Some(json!({
410 "error_type": err.error_type(),
411 })),
412 }
413 }
414}
415
416#[derive(Debug, Clone, Copy)]
417enum McpEndpointType {
418 Agent,
419 Developer,
420}
421
422impl McpEndpointType {
423 fn as_label(self) -> &'static str {
426 match self {
427 McpEndpointType::Agent => "agent",
428 McpEndpointType::Developer => "developer",
429 }
430 }
431}
432
433impl std::fmt::Display for McpEndpointType {
434 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
435 f.write_str(self.as_label())
436 }
437}
438
439pub async fn handle_mcp_method_not_allowed() -> impl IntoResponse {
442 StatusCode::METHOD_NOT_ALLOWED
443}
444
445pub async fn handle_mcp_agent(
447 headers: HeaderMap,
448 Extension(allowed_origins): Extension<Arc<Vec<HeaderValue>>>,
449 Extension(metrics): Extension<McpMetrics>,
450 client: AuthedClient,
451 Json(body): Json<McpRequest>,
452) -> axum::response::Response {
453 if let Some(resp) = validate_origin(&headers, &allowed_origins) {
454 return resp;
455 }
456 handle_mcp_request(client, body, McpEndpointType::Agent, metrics)
457 .await
458 .into_response()
459}
460
461pub async fn handle_mcp_developer(
463 headers: HeaderMap,
464 Extension(allowed_origins): Extension<Arc<Vec<HeaderValue>>>,
465 Extension(metrics): Extension<McpMetrics>,
466 client: AuthedClient,
467 Json(body): Json<McpRequest>,
468) -> axum::response::Response {
469 if let Some(resp) = validate_origin(&headers, &allowed_origins) {
470 return resp;
471 }
472 handle_mcp_request(client, body, McpEndpointType::Developer, metrics)
473 .await
474 .into_response()
475}
476
477fn validate_origin(
486 headers: &HeaderMap,
487 allowed: &[HeaderValue],
488) -> Option<axum::response::Response> {
489 let origin = headers.get(http::header::ORIGIN)?;
490 if mz_http_util::origin_is_allowed(origin, allowed) {
491 return None;
492 }
493 warn!(
494 origin = ?origin,
495 "MCP request rejected: origin not in allowlist",
496 );
497 Some(StatusCode::FORBIDDEN.into_response())
498}
499
500async fn handle_mcp_request(
501 mut client: AuthedClient,
502 request: McpRequest,
503 endpoint_type: McpEndpointType,
504 metrics: McpMetrics,
505) -> impl IntoResponse {
506 let endpoint_label = endpoint_type.as_label();
507 let method_label = request.method.to_string();
508 let record_request =
509 |status: McpCallStatus| metrics.record_request(endpoint_label, &method_label, status);
510
511 let catalog = match tokio::time::timeout(
516 *MCP_REQUEST_TIMEOUT.default(),
517 client.client.catalog_snapshot("mcp"),
518 )
519 .await
520 {
521 Ok(catalog) => catalog,
522 Err(_elapsed) => {
523 warn!(endpoint = %endpoint_type, "MCP catalog snapshot timed out");
524 record_request(McpCallStatus::Timeout);
525 return StatusCode::SERVICE_UNAVAILABLE.into_response();
526 }
527 };
528 let dyncfgs = catalog.system_config().dyncfgs();
529 let enabled = match endpoint_type {
530 McpEndpointType::Agent => ENABLE_MCP_AGENT.get(dyncfgs),
531 McpEndpointType::Developer => ENABLE_MCP_DEVELOPER.get(dyncfgs),
532 };
533 if !enabled {
534 debug!(endpoint = %endpoint_type, "MCP endpoint disabled by feature flag");
535 record_request(McpCallStatus::EndpointDisabled);
536 return StatusCode::SERVICE_UNAVAILABLE.into_response();
537 }
538
539 let query_tool_enabled = match endpoint_type {
544 McpEndpointType::Agent => ENABLE_MCP_AGENT_QUERY_TOOL.get(dyncfgs),
545 McpEndpointType::Developer => ENABLE_MCP_DEVELOPER_QUERY_TOOL.get(dyncfgs),
546 };
547 let read_data_product_tool_enabled = ENABLE_MCP_AGENT_READ_DATA_PRODUCT_TOOL.get(dyncfgs);
551 let max_response_size = MCP_MAX_RESPONSE_SIZE.get(dyncfgs);
552 let request_timeout = MCP_REQUEST_TIMEOUT.get(dyncfgs);
553
554 let app_name = match endpoint_type {
558 McpEndpointType::Agent => "mz_mcp_agents",
559 McpEndpointType::Developer => "mz_mcp_developer",
560 };
561 client
562 .client
563 .session()
564 .vars_mut()
565 .set_default(APPLICATION_NAME.name(), VarInput::Flat(app_name))
566 .expect("application_name is a known session var");
567
568 let user = client.client.session().user().name.clone();
569 let is_notification = request.id.is_none();
570
571 debug!(
572 method = %request.method,
573 endpoint = %endpoint_type,
574 user = %user,
575 is_notification = is_notification,
576 "MCP request received"
577 );
578
579 if is_notification {
582 debug!(method = %request.method, "Received notification (no response will be sent)");
583 record_request(McpCallStatus::Ok);
584 return StatusCode::ACCEPTED.into_response();
585 }
586
587 let request_id = request.id.clone().unwrap_or(serde_json::Value::Null);
588
589 let metrics_inner = metrics.clone();
594 let result = tokio::time::timeout(
595 request_timeout,
596 mz_ore::task::spawn(|| "mcp_request", async move {
597 handle_mcp_request_inner(
598 &mut client,
599 request,
600 endpoint_type,
601 query_tool_enabled,
602 read_data_product_tool_enabled,
603 max_response_size,
604 metrics_inner,
605 )
606 .await
607 })
608 .abort_on_drop(),
609 )
610 .await;
611
612 let (response, status_label): (McpResponse, McpCallStatus) = match result {
613 Ok(inner) => inner,
614 Err(_elapsed) => {
615 warn!(
616 endpoint = %endpoint_type,
617 timeout = ?request_timeout,
618 "MCP request timed out",
619 );
620 let response = McpResponse::error(
621 request_id,
622 McpRequestError::QueryExecutionFailed(format!(
623 "Request timed out after {} seconds.",
624 request_timeout.as_secs(),
625 ))
626 .into(),
627 );
628 (response, McpCallStatus::Timeout)
629 }
630 };
631
632 record_request(status_label);
633 (StatusCode::OK, Json(response)).into_response()
634}
635
636async fn handle_mcp_request_inner(
637 client: &mut AuthedClient,
638 request: McpRequest,
639 endpoint_type: McpEndpointType,
640 query_tool_enabled: bool,
641 read_data_product_tool_enabled: bool,
642 max_response_size: usize,
643 metrics: McpMetrics,
644) -> (McpResponse, McpCallStatus) {
645 let request_id = request.id.clone().unwrap_or(serde_json::Value::Null);
647
648 let result = handle_mcp_method(
649 client,
650 &request,
651 endpoint_type,
652 query_tool_enabled,
653 read_data_product_tool_enabled,
654 max_response_size,
655 &metrics,
656 )
657 .await;
658
659 let status_label = call_status(&result);
660
661 let response = match result {
662 Ok(result_value) => McpResponse::success(request_id, result_value),
663 Err(e) => {
664 if !matches!(
666 e,
667 McpRequestError::MethodNotFound(_) | McpRequestError::InvalidJsonRpcVersion
668 ) {
669 warn!(error = %e, method = %request.method, "MCP method execution failed");
670 }
671 McpResponse::error(request_id, e.into())
672 }
673 };
674
675 (response, status_label)
676}
677
678async fn handle_mcp_method(
679 client: &mut AuthedClient,
680 request: &McpRequest,
681 endpoint_type: McpEndpointType,
682 query_tool_enabled: bool,
683 read_data_product_tool_enabled: bool,
684 max_response_size: usize,
685 metrics: &McpMetrics,
686) -> Result<McpResult, McpRequestError> {
687 if request.jsonrpc != JSONRPC_VERSION {
689 return Err(McpRequestError::InvalidJsonRpcVersion);
690 }
691
692 match &request.method {
694 McpMethod::Initialize(_) => {
695 debug!(endpoint = %endpoint_type, "Processing initialize");
696 handle_initialize(
697 endpoint_type,
698 query_tool_enabled,
699 read_data_product_tool_enabled,
700 )
701 }
702 McpMethod::ToolsList(_) => {
703 debug!(endpoint = %endpoint_type, "Processing tools/list");
704 handle_tools_list(
705 endpoint_type,
706 query_tool_enabled,
707 read_data_product_tool_enabled,
708 max_response_size,
709 )
710 }
711 McpMethod::ToolsCall(params) => {
712 debug!(tool = %params, endpoint = %endpoint_type, "Processing tools/call");
713 handle_tools_call(
714 client,
715 params,
716 endpoint_type,
717 query_tool_enabled,
718 read_data_product_tool_enabled,
719 max_response_size,
720 metrics,
721 )
722 .await
723 }
724 McpMethod::Ping(_) | McpMethod::NotificationsInitialized(_) | McpMethod::Unknown => Err(
725 McpRequestError::MethodNotFound("unknown method".to_string()),
726 ),
727 }
728}
729
730fn endpoint_instructions(
733 endpoint_type: McpEndpointType,
734 query_tool_enabled: bool,
735 read_data_product_tool_enabled: bool,
736) -> Option<String> {
737 match endpoint_type {
738 McpEndpointType::Agent => {
739 let read_paragraph = match (read_data_product_tool_enabled, query_tool_enabled) {
744 (true, _) => {
745 "`read_data_product` automatically routes the \
746 read to the cluster recorded in the data product catalog so indexes are used; \
747 you only need to set the `cluster` parameter if you intentionally want the \
748 read to run on a different cluster (e.g. one with larger or more replicas). \
749 A null `cluster` in discovery means your role lacks USAGE on the object's \
750 index/compute cluster; `read_data_product` without an override then reads on \
751 your session's default cluster (safe: only materialized views appear this way, \
752 and they serve from persist). "
753 }
754 (false, true) => {
755 "Use the `query` tool to read data products, passing the cluster from \
756 `get_data_product_details` so indexed reads hit the arrangement. If \
757 `get_data_product_details` returns a null `cluster`, your role lacks USAGE on \
758 the object's index/compute cluster; run the `query` against any cluster your \
759 role can use (the read still works, just without the index arrangement). "
760 }
761 (false, false) => {
762 "This server is configured for discovery only: no read tool is exposed. \
763 Use `get_data_products` and `get_data_product_details` to inspect what \
764 is available. "
765 }
766 };
767 Some(format!(
768 "You have access to Materialize data products via MCP. \
769 Prefer indexed objects (served from memory) over unindexed materialized views \
770 (read from persistent storage). {read_paragraph}\
771 `get_data_product_details` returns a `hydration` object with `hydrated`, \
772 `replica_count`, and `hydrated_replica_count` fields. Reads never return \
773 partial data: a read against a not-yet-hydrated product blocks until the \
774 dataflow catches up, and may hit the request timeout. Check `hydrated` \
775 before reading: if it is false and `replica_count` is greater than 0, the \
776 dataflow is still warming up, so wait and retry; if `replica_count` is 0 the \
777 cluster has no replicas and the read cannot make progress until one is added.",
778 ))
779 }
780 McpEndpointType::Developer => {
781 let query_tool_line = if query_tool_enabled {
785 "- query: read-only SELECT/SHOW/EXPLAIN that can also reach user objects on a named cluster. Use this for EXPLAIN ANALYZE and for inspecting user objects directly.\n"
786 } else {
787 ""
788 };
789 let introspection_rule = if query_tool_enabled {
793 "- mz_introspection relations (for example mz_dataflow_arrangement_sizes) are cluster-scoped and query_system_catalog answers about the session's default cluster: read them only through the query tool with its cluster argument (and cluster_replica on a multi-replica cluster), never through query_system_catalog"
794 } else {
795 "- mz_introspection relations (for example mz_dataflow_arrangement_sizes) are cluster-scoped and query_system_catalog answers about the session's default cluster only; cluster-targeted introspection is unavailable on this server (the query tool is disabled), so say so rather than reporting another cluster's numbers"
796 };
797 Some(format!(
798 "You are connected to the Materialize developer MCP server for troubleshooting and observability.\n\n\
799 Tools:\n\
800 - query_system_catalog: read-only SELECT/SHOW/EXPLAIN restricted to system catalog tables (mz_*, pg_catalog, information_schema). No cluster argument; prefer this for most catalog lookups.\n\
801 {query_tool_line}\n\
802 IMPORTANT: Before writing queries, discover table schemas using the mz_ontology tables:\n\
803 - mz_internal.mz_ontology_entity_types: what catalog entities exist and which tables they map to\n\
804 - mz_internal.mz_ontology_link_types: relationships between entities (foreign keys, metrics, etc.)\n\
805 - mz_internal.mz_ontology_properties: column names, types, and descriptions for each entity\n\
806 - mz_internal.mz_ontology_semantic_types: typed ID domains (CatalogItemId, ReplicaId, etc.)\n\n\
807 Use these to find the correct tables, join paths, and column names instead of guessing.\n\n\
808 Key rules:\n\
809 - mz_source_statuses and mz_sink_statuses use `last_status_change_at` (NOT `updated_at`)\n\
810 - mz_cluster_replica_utilization only has `replica_id` — JOIN with mz_cluster_replicas and mz_clusters to get names\n\
811 {introspection_rule}\n\
812 - mz_dataflow_arrangement_sizes.id and mz_dataflows.id are dataflow ids (uint8), not catalog ids: do not JOIN them to mz_catalog.mz_objects.id (text); join through mz_introspection.mz_compute_exports (dataflow_id to export_id) instead, or match the name column, which reads Dataflow: <database>.<schema>.<object>\n\
813 - Use SHOW COLUMNS FROM <table> to verify column names if unsure",
814 ))
815 }
816 }
817}
818
819fn handle_initialize(
820 endpoint_type: McpEndpointType,
821 query_tool_enabled: bool,
822 read_data_product_tool_enabled: bool,
823) -> Result<McpResult, McpRequestError> {
824 Ok(McpResult::Initialize(InitializeResult {
825 protocol_version: MCP_PROTOCOL_VERSION.to_string(),
826 capabilities: Capabilities { tools: json!({}) },
827 server_info: ServerInfo {
828 name: format!("materialize-mcp-{}", endpoint_type),
829 version: env!("CARGO_PKG_VERSION").to_string(),
830 },
831 instructions: endpoint_instructions(
832 endpoint_type,
833 query_tool_enabled,
834 read_data_product_tool_enabled,
835 ),
836 }))
837}
838
839fn handle_tools_list(
840 endpoint_type: McpEndpointType,
841 query_tool_enabled: bool,
842 read_data_product_tool_enabled: bool,
843 max_response_size: usize,
844) -> Result<McpResult, McpRequestError> {
845 let size_hint = format!(
846 "Response limit: {:.1} MB.",
847 f64::cast_lossy(max_response_size) / 1_000_000.0
848 );
849
850 let tools = match endpoint_type {
851 McpEndpointType::Agent => {
852 let mut tools = vec![
853 ToolDefinition {
854 name: "get_data_products".to_string(),
855 title: Some("List Data Products".to_string()),
856 description: "Discover all available real-time data views (data products) that represent business entities like customers, orders, products, etc. Each data product provides fresh, queryable data with defined schemas. Use this first to see what data is available before querying specific information.".to_string(),
857 input_schema: json!({
858 "type": "object",
859 "properties": {},
860 "required": []
861 }),
862 annotations: Some(READ_ONLY_ANNOTATIONS),
863 },
864 ToolDefinition {
865 name: "get_data_product_details".to_string(),
866 title: Some("Get Data Product Details".to_string()),
867 description: "Get the complete schema and structure of a specific data product, plus a `hydration` object reporting whether the dataflow is ready across the cluster's replicas (`{hydrated, replica_count, hydrated_replica_count}`). This shows you exactly what fields are available, their types, and what data you can query. Reads never return partial data, so check `hydration` before reading: if `hydrated` is false and `replica_count` is greater than 0 the dataflow is still warming up (a read would block until it catches up, possibly hitting the request timeout), so wait and retry; if `replica_count` is 0 the cluster has no replicas and the read cannot make progress until one is added.".to_string(),
868 input_schema: json!({
869 "type": "object",
870 "properties": {
871 "name": {
872 "type": "string",
873 "description": "Exact name of the data product from get_data_products() list"
874 }
875 },
876 "required": ["name"]
877 }),
878 annotations: Some(READ_ONLY_ANNOTATIONS),
879 },
880 ];
881 if read_data_product_tool_enabled {
882 tools.push(ToolDefinition {
883 name: "read_data_product".to_string(),
884 title: Some("Read Data Product".to_string()),
885 description: format!("Read rows from a specific data product. Returns up to `limit` rows (default {DEFAULT_READ_LIMIT}). The data product must exist in the catalog (use get_data_products() to discover available products). Use this to retrieve actual data from a known data product. {size_hint}"),
886 input_schema: json!({
887 "type": "object",
888 "properties": {
889 "name": {
890 "type": "string",
891 "description": "Exact fully-qualified name of the data product (e.g. '\"materialize\".\"schema\".\"view_name\"')"
892 },
893 "limit": {
894 "type": "integer",
895 "description": format!("Maximum number of rows to return (default {DEFAULT_READ_LIMIT})"),
896 "default": DEFAULT_READ_LIMIT
897 },
898 "cluster": {
899 "type": "string",
900 "description": "Optional override. By default, the read runs on the cluster recorded in the data product catalog (where the index or materialized view dataflow lives), so indexed reads actually hit their arrangement. A null `cluster` in discovery means your role lacks USAGE on the object's index/compute cluster; without an override, the read then runs on your session's default cluster (safe: only materialized views appear this way, and they serve from persist). Set this only to intentionally run the same read on a different cluster — e.g. one with more or larger replicas, or to compare cost/latency."
901 }
902 },
903 "required": ["name"]
904 }),
905 annotations: Some(READ_ONLY_ANNOTATIONS),
906 });
907 }
908 if query_tool_enabled {
909 tools.push(ToolDefinition {
910 name: "query".to_string(),
911 title: Some("Query Data Products".to_string()),
912 description: format!("Execute SQL queries against real-time data products to retrieve current business information. Use standard PostgreSQL syntax. You can JOIN multiple data products together, but ONLY if they are all hosted on the same cluster. Always specify the cluster parameter from the data product details. This provides fresh, up-to-date results from materialized views. {size_hint}"),
913 input_schema: json!({
914 "type": "object",
915 "properties": {
916 "cluster": {
917 "type": "string",
918 "description": "Exact cluster name from the data product details - required for query execution"
919 },
920 "sql_query": {
921 "type": "string",
922 "description": "PostgreSQL-compatible SELECT statement to retrieve data. Use the fully qualified data product name exactly as provided (with double quotes). You can JOIN multiple data products, but only those on the same cluster."
923 }
924 },
925 "required": ["cluster", "sql_query"]
926 }),
927 annotations: Some(READ_ONLY_ANNOTATIONS),
928 });
929 }
930 tools
931 }
932 McpEndpointType::Developer => {
933 let mut tools = vec![ToolDefinition {
934 name: "query_system_catalog".to_string(),
935 title: Some("Query System Catalog".to_string()),
936 description: concat!(
937 "Query Materialize system catalog tables for troubleshooting and observability. ",
938 "Only mz_*, pg_catalog, and information_schema tables are accessible. ",
939 "Use the mz_internal.mz_ontology_* tables to discover tables, columns, and join paths before writing queries.",
940 ).to_owned() + &format!(" {size_hint}"),
941 input_schema: json!({
942 "type": "object",
943 "properties": {
944 "sql_query": {
945 "type": "string",
946 "description": "PostgreSQL-compatible SELECT, SHOW, or EXPLAIN query referencing mz_*, pg_catalog, or information_schema tables"
947 }
948 },
949 "required": ["sql_query"]
950 }),
951 annotations: Some(READ_ONLY_ANNOTATIONS),
952 }];
953 if query_tool_enabled {
954 tools.push(ToolDefinition {
955 name: "query".to_string(),
956 title: Some("Query".to_string()),
957 description: format!(
958 "Execute a read-only SQL query (SELECT, SHOW, or EXPLAIN) against any object the role can access, including system catalog and user objects. Requires a cluster, which is what enables EXPLAIN ANALYZE and queries against indexed user objects. For pure system catalog lookups that do not need a cluster, prefer `query_system_catalog`. {size_hint}",
959 ),
960 input_schema: json!({
961 "type": "object",
962 "properties": {
963 "cluster": {
964 "type": "string",
965 "description": "Exact cluster name the query should run on. Required: EXPLAIN ANALYZE and queries against indexed user objects need a specific cluster to execute on."
966 },
967 "cluster_replica": {
968 "type": "string",
969 "description": "Optional replica name (e.g. 'r1') to target one replica of the cluster. Required for EXPLAIN ANALYZE on clusters with more than one replica. Find replica names in mz_catalog.mz_cluster_replicas."
970 },
971 "sql_query": {
972 "type": "string",
973 "description": "PostgreSQL-compatible SELECT, SHOW, or EXPLAIN statement. Multi-statement queries are rejected."
974 }
975 },
976 "required": ["cluster", "sql_query"]
977 }),
978 annotations: Some(READ_ONLY_ANNOTATIONS),
979 });
980 }
981 tools
982 }
983 };
984
985 Ok(McpResult::ToolsList(ToolsListResult { tools }))
986}
987
988async fn handle_tools_call(
989 client: &mut AuthedClient,
990 params: &ToolsCallParams,
991 endpoint_type: McpEndpointType,
992 query_tool_enabled: bool,
993 read_data_product_tool_enabled: bool,
994 max_response_size: usize,
995 metrics: &McpMetrics,
996) -> Result<McpResult, McpRequestError> {
997 let mut guard = ToolCallGuard::new(metrics, endpoint_type.as_label(), params.to_string());
999
1000 let result = match (endpoint_type, params) {
1001 (McpEndpointType::Agent, ToolsCallParams::GetDataProducts(_)) => {
1002 get_data_products(client, max_response_size).await
1003 }
1004 (McpEndpointType::Agent, ToolsCallParams::GetDataProductDetails(p)) => {
1005 get_data_product_details(client, &p.name, max_response_size).await
1006 }
1007 (McpEndpointType::Agent, ToolsCallParams::ReadDataProduct(_))
1008 if !read_data_product_tool_enabled =>
1009 {
1010 Err(McpRequestError::ToolNotFound(
1011 "read_data_product tool is not available. Use the query tool to read data products."
1012 .to_string(),
1013 ))
1014 }
1015 (McpEndpointType::Agent, ToolsCallParams::ReadDataProduct(p)) => {
1016 read_data_product(
1017 client,
1018 &p.name,
1019 p.limit,
1020 p.cluster.as_deref(),
1021 max_response_size,
1022 )
1023 .await
1024 }
1025 (McpEndpointType::Agent, ToolsCallParams::Query(_)) if !query_tool_enabled => {
1026 Err(McpRequestError::ToolNotFound(
1027 "query tool is not available. Use get_data_products, get_data_product_details, and read_data_product instead.".to_string(),
1028 ))
1029 }
1030 (McpEndpointType::Agent, ToolsCallParams::Query(p)) => {
1031 execute_query(client, &p.cluster, None, &p.sql_query, max_response_size).await
1034 }
1035 (McpEndpointType::Developer, ToolsCallParams::QuerySystemCatalog(p)) => {
1036 query_system_catalog(client, &p.sql_query, max_response_size).await
1037 }
1038 (McpEndpointType::Developer, ToolsCallParams::Query(_)) if !query_tool_enabled => {
1039 Err(McpRequestError::ToolNotFound(
1040 "query tool is not available. Use query_system_catalog instead.".to_string(),
1041 ))
1042 }
1043 (McpEndpointType::Developer, ToolsCallParams::Query(p)) => {
1044 execute_query(
1045 client,
1046 &p.cluster,
1047 p.cluster_replica.as_deref(),
1048 &p.sql_query,
1049 max_response_size,
1050 )
1051 .await
1052 }
1053 (endpoint, tool) => Err(McpRequestError::ToolNotFound(format!(
1055 "{} is not available on {} endpoint",
1056 tool, endpoint
1057 ))),
1058 };
1059
1060 guard.set_status(call_status(&result));
1061
1062 result
1063}
1064
1065fn call_status<T>(result: &Result<T, McpRequestError>) -> McpCallStatus {
1068 match result {
1069 Ok(_) => McpCallStatus::Ok,
1070 Err(e) => McpCallStatus::Error(e.error_type()),
1071 }
1072}
1073
1074async fn execute_sql(
1076 client: &mut AuthedClient,
1077 query: &str,
1078) -> Result<Vec<Box<serde_json::value::RawValue>>, McpRequestError> {
1079 let mut response = SqlResponse::new();
1080
1081 execute_request(
1082 client,
1083 SqlRequest::Simple {
1084 query: mz_ore::sql::Sql::trusted_external_request(query.to_string()),
1085 },
1086 &mut response,
1087 )
1088 .await
1089 .map_err(|e| McpRequestError::QueryExecutionFailed(e.to_string()))?;
1090
1091 select_single_rows(response.results)
1092}
1093
1094fn select_single_rows(
1100 results: Vec<SqlResult>,
1101) -> Result<Vec<Box<serde_json::value::RawValue>>, McpRequestError> {
1102 let mut rows = None;
1103 for result in results {
1104 match result {
1105 SqlResult::Rows { rows: r, .. } => {
1106 if rows.is_some() {
1107 return Err(McpRequestError::Internal(anyhow!(
1108 "MCP query returned multiple row-producing statements"
1109 )));
1110 }
1111 rows = Some(r);
1112 }
1113 SqlResult::Err { error, .. } => {
1114 return Err(McpRequestError::QueryExecutionFailed(error.message));
1115 }
1116 SqlResult::Ok { .. } => continue,
1117 }
1118 }
1119
1120 rows.ok_or_else(|| {
1121 McpRequestError::QueryExecutionFailed("Query did not return any results".to_string())
1122 })
1123}
1124
1125fn format_rows_response(
1132 rows: Vec<Box<serde_json::value::RawValue>>,
1133 max_size: usize,
1134) -> Result<McpResult, McpRequestError> {
1135 let text =
1138 serde_json::to_string_pretty(&rows).map_err(|e| McpRequestError::Internal(anyhow!(e)))?;
1139
1140 if text.len() > max_size {
1141 return Err(McpRequestError::QueryExecutionFailed(format!(
1142 "Response size ({} bytes) exceeds the {} byte limit. \
1143 Use LIMIT or WHERE to narrow your query.",
1144 text.len(),
1145 max_size,
1146 )));
1147 }
1148
1149 Ok(McpResult::ToolContent(ToolContentResult {
1150 content: vec![ContentBlock {
1151 content_type: "text".to_string(),
1152 text,
1153 }],
1154 is_error: false,
1155 }))
1156}
1157
1158async fn get_data_products(
1159 client: &mut AuthedClient,
1160 max_response_size: usize,
1161) -> Result<McpResult, McpRequestError> {
1162 debug!("Executing get_data_products");
1163 let rows = execute_sql(client, DISCOVERY_QUERY).await?;
1164 debug!("get_data_products returned {} rows", rows.len());
1165
1166 format_rows_response(rows, max_response_size)
1167}
1168
1169async fn get_data_product_details(
1170 client: &mut AuthedClient,
1171 name: &str,
1172 max_response_size: usize,
1173) -> Result<McpResult, McpRequestError> {
1174 debug!(name = %name, "Executing get_data_product_details");
1175
1176 let query = format!("{}{}", DETAILS_QUERY_PREFIX, escaped_string_literal(name));
1177
1178 let rows = execute_sql(client, &query).await?;
1179
1180 if rows.is_empty() {
1181 return Err(McpRequestError::DataProductNotFound(name.to_string()));
1182 }
1183
1184 format_rows_response(rows, max_response_size)
1185}
1186
1187fn safe_data_product_name(name: &str) -> Result<String, McpRequestError> {
1194 let name = name.trim();
1195 if name.is_empty() {
1196 return Err(McpRequestError::QueryValidationFailed(
1197 "Data product name cannot be empty".to_string(),
1198 ));
1199 }
1200
1201 let parsed = parse_item_name_with_limit(name)
1204 .map_err(McpRequestError::QueryValidationFailed)?
1205 .map_err(|_| {
1206 McpRequestError::QueryValidationFailed(format!(
1207 "Invalid data product name: {}. Expected a valid object name, \
1208 e.g. '\"database\".\"schema\".\"name\"' or 'my_view'",
1209 name
1210 ))
1211 })?;
1212
1213 Ok(parsed.to_ast_string_stable())
1216}
1217
1218async fn read_data_product(
1237 client: &mut AuthedClient,
1238 name: &str,
1239 limit: u32,
1240 cluster_override: Option<&str>,
1241 max_response_size: usize,
1242) -> Result<McpResult, McpRequestError> {
1243 debug!(name = %name, limit = limit, cluster_override = ?cluster_override, "Executing read_data_product");
1244
1245 let safe_name = safe_data_product_name(name)?;
1247
1248 let lookup_query = format!(
1254 "SELECT dp.cluster FROM mz_internal.mz_mcp_data_products dp \
1255 WHERE dp.object_name = {} \
1256 ORDER BY dp.cluster NULLS LAST \
1257 LIMIT 1",
1258 escaped_string_literal(name)
1259 );
1260 let lookup_rows = execute_sql(client, &lookup_query).await?;
1261 if lookup_rows.is_empty() {
1262 return Err(McpRequestError::DataProductNotFound(name.to_string()));
1263 }
1264 let catalog_cluster: Option<String> = lookup_rows
1267 .first()
1268 .and_then(|row| serde_json::from_str::<Vec<serde_json::Value>>(row.get()).ok())
1269 .and_then(|row| row.into_iter().next())
1270 .and_then(|v| v.as_str().map(|s| s.to_string()));
1271
1272 let target_cluster: Option<&str> = cluster_override.or(catalog_cluster.as_deref());
1278
1279 let read_query = build_read_query(&safe_name, limit, target_cluster);
1284
1285 let rows = execute_sql(client, &read_query).await?;
1286
1287 format_rows_response(rows, max_response_size)
1288}
1289
1290fn build_read_query(safe_name: &str, limit: u32, target_cluster: Option<&str>) -> String {
1299 let body = format!("SELECT * FROM {safe_name} LIMIT {limit}");
1300 match target_cluster {
1301 Some(cluster) => read_only_txn(
1302 &format!("SET CLUSTER = {}", escaped_string_literal(cluster)),
1303 &body,
1304 ),
1305 None => format!("BEGIN READ ONLY; {body}\n; COMMIT;"),
1306 }
1307}
1308
1309fn read_only_txn(set_clause: &str, body: &str) -> String {
1315 format!("BEGIN READ ONLY; {set_clause}; {body}\n; COMMIT;")
1316}
1317
1318fn validate_readonly_query(sql: &str) -> Result<(), McpRequestError> {
1320 let sql = sql.trim();
1321 if sql.is_empty() {
1322 return Err(McpRequestError::QueryValidationFailed(
1323 "Empty query".to_string(),
1324 ));
1325 }
1326
1327 let stmts = parse_with_limit(sql)
1331 .map_err(McpRequestError::QueryValidationFailed)?
1332 .map_err(|e| {
1333 McpRequestError::QueryValidationFailed(format!("Failed to parse SQL: {}", e))
1334 })?;
1335
1336 if stmts.len() != 1 {
1338 return Err(McpRequestError::QueryValidationFailed(format!(
1339 "Only one query allowed at a time. Found {} statements.",
1340 stmts.len()
1341 )));
1342 }
1343
1344 let stmt = &stmts[0];
1350 use mz_sql_parser::ast::Statement;
1351
1352 match &stmt.ast {
1353 Statement::Select(_)
1354 | Statement::Show(_)
1355 | Statement::ExplainPlan(_)
1356 | Statement::ExplainPushdown(_)
1357 | Statement::ExplainTimestamp(_)
1358 | Statement::ExplainSinkSchema(_)
1359 | Statement::ExplainAnalyzeObject(_)
1360 | Statement::ExplainAnalyzeCluster(_) => Ok(()),
1361 _ => Err(McpRequestError::QueryValidationFailed(
1362 "Only SELECT, SHOW, and EXPLAIN statements are allowed".to_string(),
1363 )),
1364 }
1365}
1366
1367async fn execute_query(
1368 client: &mut AuthedClient,
1369 cluster: &str,
1370 cluster_replica: Option<&str>,
1371 sql_query: &str,
1372 max_response_size: usize,
1373) -> Result<McpResult, McpRequestError> {
1374 debug!(cluster = %cluster, cluster_replica = ?cluster_replica, "Executing user query");
1375
1376 validate_readonly_query(sql_query)?;
1377 validate_cluster_replica(cluster_replica)?;
1378
1379 let combined_query = read_only_txn(&query_set_clause(cluster, cluster_replica), sql_query);
1382
1383 let rows = execute_sql(client, &combined_query).await?;
1384
1385 format_rows_response(rows, max_response_size)
1386}
1387
1388fn query_set_clause(cluster: &str, cluster_replica: Option<&str>) -> String {
1394 let mut set_clause = format!("SET CLUSTER = {}", escaped_string_literal(cluster));
1395 if let Some(replica) = cluster_replica {
1396 set_clause.push_str(&format!(
1397 "; SET CLUSTER_REPLICA = {}",
1398 escaped_string_literal(replica)
1399 ));
1400 }
1401 set_clause
1402}
1403
1404fn validate_cluster_replica(cluster_replica: Option<&str>) -> Result<(), McpRequestError> {
1409 if let Some(replica) = cluster_replica {
1410 if replica.trim().is_empty() {
1411 return Err(McpRequestError::QueryValidationFailed(
1412 "cluster_replica must not be empty or whitespace-only".to_string(),
1413 ));
1414 }
1415 }
1416 Ok(())
1417}
1418
1419async fn query_system_catalog(
1420 client: &mut AuthedClient,
1421 sql_query: &str,
1422 max_response_size: usize,
1423) -> Result<McpResult, McpRequestError> {
1424 debug!("Executing query_system_catalog");
1425
1426 validate_readonly_query(sql_query)?;
1428
1429 validate_system_catalog_query(sql_query)?;
1431
1432 let combined_query = read_only_txn(
1438 "SET search_path = mz_catalog, mz_internal, pg_catalog, information_schema",
1439 sql_query,
1440 );
1441
1442 let rows = execute_sql(client, &combined_query).await?;
1443
1444 format_rows_response(rows, max_response_size)
1445}
1446
1447struct TableReferenceCollector {
1449 tables: Vec<(Option<String>, String)>,
1451 cte_names: std::collections::BTreeSet<String>,
1453}
1454
1455impl TableReferenceCollector {
1456 fn new() -> Self {
1457 Self {
1458 tables: Vec::new(),
1459 cte_names: std::collections::BTreeSet::new(),
1460 }
1461 }
1462}
1463
1464impl<'ast> Visit<'ast, Raw> for TableReferenceCollector {
1465 fn visit_cte(&mut self, cte: &'ast mz_sql_parser::ast::Cte<Raw>) {
1466 self.cte_names
1468 .insert(cte.alias.name.as_str().to_lowercase());
1469 visit::visit_cte(self, cte);
1470 }
1471
1472 fn visit_table_factor(&mut self, table_factor: &'ast mz_sql_parser::ast::TableFactor<Raw>) {
1473 if let mz_sql_parser::ast::TableFactor::Table { name, .. } = table_factor {
1475 match name {
1476 RawItemName::Name(n) | RawItemName::Id(_, n, _) => {
1477 let parts = &n.0;
1478 if !parts.is_empty() {
1479 let table_name = parts.last().unwrap().as_str().to_lowercase();
1480
1481 if self.cte_names.contains(&table_name) {
1483 visit::visit_table_factor(self, table_factor);
1484 return;
1485 }
1486
1487 let schema = if parts.len() >= 2 {
1489 Some(parts[parts.len() - 2].as_str().to_lowercase())
1490 } else {
1491 None
1492 };
1493 self.tables.push((schema, table_name));
1494 }
1495 }
1496 }
1497 }
1498 visit::visit_table_factor(self, table_factor);
1499 }
1500}
1501
1502fn validate_system_catalog_query(sql: &str) -> Result<(), McpRequestError> {
1510 let stmts = parse_with_limit(sql)
1513 .map_err(McpRequestError::QueryValidationFailed)?
1514 .map_err(|e| {
1515 McpRequestError::QueryValidationFailed(format!("Failed to parse SQL: {}", e))
1516 })?;
1517
1518 if stmts.is_empty() {
1519 return Err(McpRequestError::QueryValidationFailed(
1520 "Empty query".to_string(),
1521 ));
1522 }
1523
1524 let mut collector = TableReferenceCollector::new();
1526 for stmt in &stmts {
1527 collector.visit_statement(&stmt.ast);
1528 }
1529
1530 let is_allowed_schema =
1533 |s: &str| SYSTEM_SCHEMAS.contains(&s) && s != namespaces::MZ_UNSAFE_SCHEMA;
1534
1535 let is_system_table = |(schema, table_name): &(Option<String>, String)| match schema {
1540 Some(s) => is_allowed_schema(s.as_str()),
1541 None => table_name.starts_with("mz_") || table_name.starts_with("pg_"),
1542 };
1543
1544 let non_system_tables: Vec<String> = collector
1546 .tables
1547 .iter()
1548 .filter(|t| !is_system_table(t))
1549 .map(|(schema, table)| match schema {
1550 Some(s) => format!("{}.{}", s, table),
1551 None => table.clone(),
1552 })
1553 .collect();
1554
1555 if !non_system_tables.is_empty() {
1556 return Err(McpRequestError::QueryValidationFailed(format!(
1557 "Query references non-system tables: {}. Only system catalog tables (mz_*, pg_catalog, information_schema) are allowed.",
1558 non_system_tables.join(", ")
1559 )));
1560 }
1561
1562 use mz_sql_parser::ast::Statement;
1565 let is_select = stmts.iter().any(|s| matches!(&s.ast, Statement::Select(_)));
1566
1567 if is_select && (collector.tables.is_empty() || !collector.tables.iter().any(is_system_table)) {
1568 return Err(McpRequestError::QueryValidationFailed(
1569 "Query must reference at least one system catalog table".to_string(),
1570 ));
1571 }
1572
1573 Ok(())
1574}
1575
1576#[cfg(test)]
1577mod tests {
1578 use super::*;
1579 use crate::http::sql::{Description, SqlError};
1580
1581 fn raw_rows(rows: Vec<Vec<serde_json::Value>>) -> Vec<Box<serde_json::value::RawValue>> {
1584 rows.iter()
1585 .map(|r| serde_json::value::to_raw_value(r).unwrap())
1586 .collect()
1587 }
1588
1589 fn rows_result(rows: Vec<Vec<serde_json::Value>>) -> SqlResult {
1590 SqlResult::Rows {
1591 tag: String::new(),
1592 rows: raw_rows(rows),
1593 desc: Description { columns: vec![] },
1594 notices: vec![],
1595 }
1596 }
1597
1598 fn ok_result() -> SqlResult {
1599 SqlResult::Ok {
1600 ok: String::new(),
1601 notices: vec![],
1602 parameters: vec![],
1603 }
1604 }
1605
1606 fn err_result(message: &str) -> SqlResult {
1607 SqlResult::Err {
1608 error: SqlError {
1609 message: message.to_string(),
1610 code: String::new(),
1611 detail: None,
1612 hint: None,
1613 position: None,
1614 },
1615 notices: vec![],
1616 }
1617 }
1618
1619 #[mz_ore::test]
1622 fn test_select_single_rows_extracts_rows() {
1623 let rows = vec![vec![serde_json::json!(1)]];
1624 let results = vec![
1625 ok_result(),
1626 ok_result(),
1627 rows_result(rows.clone()),
1628 ok_result(),
1629 ];
1630 let got = select_single_rows(results).unwrap();
1631 let got: Vec<&str> = got.iter().map(|r| r.get()).collect();
1632 assert_eq!(got, vec!["[1]"]);
1633 }
1634
1635 #[mz_ore::test]
1637 fn test_select_single_rows_requires_rows() {
1638 let err = select_single_rows(vec![ok_result(), ok_result()]).unwrap_err();
1639 assert!(
1640 matches!(err, McpRequestError::QueryExecutionFailed(_)),
1641 "{err:?}"
1642 );
1643 }
1644
1645 #[mz_ore::test]
1648 fn test_select_single_rows_rejects_multiple() {
1649 let results = vec![rows_result(vec![]), rows_result(vec![])];
1650 let err = select_single_rows(results).unwrap_err();
1651 assert!(matches!(err, McpRequestError::Internal(_)), "{err:?}");
1652 }
1653
1654 #[mz_ore::test]
1656 fn test_select_single_rows_surfaces_error() {
1657 let err = select_single_rows(vec![ok_result(), err_result("boom")]).unwrap_err();
1658 match err {
1659 McpRequestError::QueryExecutionFailed(msg) => assert_eq!(msg, "boom"),
1660 other => panic!("unexpected error: {other:?}"),
1661 }
1662 }
1663
1664 #[mz_ore::test]
1667 fn test_validate_origin() {
1668 let allowed = [HeaderValue::from_static("https://good.example")];
1669
1670 assert!(validate_origin(&HeaderMap::new(), &allowed).is_none());
1671
1672 let mut ok = HeaderMap::new();
1673 ok.insert(http::header::ORIGIN, allowed[0].clone());
1674 assert!(validate_origin(&ok, &allowed).is_none());
1675
1676 let mut bad = HeaderMap::new();
1677 bad.insert(
1678 http::header::ORIGIN,
1679 HeaderValue::from_static("https://evil.example"),
1680 );
1681 let rejected = validate_origin(&bad, &allowed);
1682 assert_eq!(
1683 rejected
1684 .expect("disallowed origin must be rejected")
1685 .status(),
1686 StatusCode::FORBIDDEN,
1687 );
1688 }
1689
1690 #[mz_ore::test]
1693 fn test_mcp_response_constructors() {
1694 let id = serde_json::json!(1);
1695
1696 let ok = McpResponse::success(
1697 id.clone(),
1698 McpResult::ToolContent(ToolContentResult {
1699 content: vec![],
1700 is_error: false,
1701 }),
1702 );
1703 assert_eq!(ok.jsonrpc, JSONRPC_VERSION);
1704 assert!(ok.result.is_some());
1705 assert!(ok.error.is_none());
1706
1707 let err = McpResponse::error(id, McpRequestError::ToolNotFound("t".to_string()).into());
1708 assert_eq!(err.jsonrpc, JSONRPC_VERSION);
1709 assert!(err.result.is_none());
1710 assert!(err.error.is_some());
1711 }
1712
1713 #[mz_ore::test]
1714 fn test_validate_readonly_query_select() {
1715 assert!(validate_readonly_query("SELECT * FROM mz_tables").is_ok());
1716 assert!(validate_readonly_query("SELECT 1 + 2").is_ok());
1717 assert!(validate_readonly_query(" SELECT 1 ").is_ok());
1718 }
1719
1720 #[mz_ore::test]
1721 fn test_validate_readonly_query_subqueries() {
1722 assert!(
1724 validate_readonly_query(
1725 "SELECT * FROM mz_tables WHERE id IN (SELECT id FROM mz_columns)"
1726 )
1727 .is_ok()
1728 );
1729
1730 assert!(
1732 validate_readonly_query(
1733 "SELECT * FROM (SELECT * FROM mz_tables WHERE name LIKE 'test%') AS t"
1734 )
1735 .is_ok()
1736 );
1737
1738 assert!(validate_readonly_query(
1740 "SELECT * FROM mz_tables t WHERE EXISTS (SELECT 1 FROM mz_columns c WHERE c.id = t.id)"
1741 )
1742 .is_ok());
1743
1744 assert!(validate_readonly_query(
1746 "SELECT * FROM mz_tables WHERE id IN (SELECT id FROM mz_columns WHERE type IN (SELECT name FROM mz_types))"
1747 )
1748 .is_ok());
1749
1750 assert!(
1752 validate_readonly_query(
1753 "SELECT * FROM mz_tables WHERE id = (SELECT MAX(id) FROM mz_columns)"
1754 )
1755 .is_ok()
1756 );
1757 }
1758
1759 #[mz_ore::test]
1760 fn test_validate_readonly_query_show() {
1761 assert!(validate_readonly_query("SHOW CLUSTERS").is_ok());
1762 assert!(validate_readonly_query("SHOW TABLES").is_ok());
1763 }
1764
1765 #[mz_ore::test]
1766 fn test_validate_readonly_query_explain() {
1767 assert!(validate_readonly_query("EXPLAIN SELECT 1").is_ok());
1772 assert!(
1773 validate_readonly_query("EXPLAIN FILTER PUSHDOWN FOR SELECT * FROM mz_tables").is_ok()
1774 );
1775 assert!(validate_readonly_query("EXPLAIN TIMESTAMP FOR SELECT 1").is_ok());
1776 assert!(validate_readonly_query("EXPLAIN ANALYZE MEMORY FOR INDEX foo").is_ok());
1777 assert!(validate_readonly_query("EXPLAIN ANALYZE MEMORY FOR MATERIALIZED VIEW mv").is_ok());
1778 assert!(validate_readonly_query("EXPLAIN ANALYZE CLUSTER MEMORY").is_ok());
1779 }
1780
1781 #[mz_ore::test]
1782 fn test_validate_readonly_query_rejects_writes() {
1783 assert!(validate_readonly_query("INSERT INTO t VALUES (1)").is_err());
1784 assert!(validate_readonly_query("UPDATE t SET a = 1").is_err());
1785 assert!(validate_readonly_query("DELETE FROM t").is_err());
1786 assert!(validate_readonly_query("CREATE TABLE t (a INT)").is_err());
1787 assert!(validate_readonly_query("DROP TABLE t").is_err());
1788 }
1789
1790 #[mz_ore::test]
1791 fn test_validate_readonly_query_rejects_multiple() {
1792 assert!(validate_readonly_query("SELECT 1; SELECT 2").is_err());
1793 }
1794
1795 #[mz_ore::test]
1796 fn test_validate_readonly_query_rejects_empty() {
1797 assert!(validate_readonly_query("").is_err());
1798 assert!(validate_readonly_query(" ").is_err());
1799 }
1800
1801 #[mz_ore::test]
1808 fn test_validate_readonly_query_enforces_size_limit() {
1809 use mz_sql_parser::parser::MAX_STATEMENT_BATCH_SIZE;
1810 let oversized: String =
1812 "SELECT 1;".repeat((MAX_STATEMENT_BATCH_SIZE / "SELECT 1;".len()) + 1);
1813 let err = validate_readonly_query(&oversized).expect_err("should be rejected");
1814 let msg = err.to_string();
1815 assert!(
1816 msg.contains("statement batch size cannot exceed"),
1817 "expected size-guard error, got: {msg}"
1818 );
1819 }
1820
1821 #[mz_ore::test]
1822 fn test_validate_system_catalog_query_enforces_size_limit() {
1823 use mz_sql_parser::parser::MAX_STATEMENT_BATCH_SIZE;
1824 let stmt = "SELECT * FROM mz_tables;";
1825 let oversized: String = stmt.repeat((MAX_STATEMENT_BATCH_SIZE / stmt.len()) + 1);
1826 let err = validate_system_catalog_query(&oversized).expect_err("should be rejected");
1827 let msg = err.to_string();
1828 assert!(
1829 msg.contains("statement batch size cannot exceed"),
1830 "expected size-guard error, got: {msg}"
1831 );
1832 }
1833
1834 #[mz_ore::test]
1835 fn test_validate_system_catalog_query_accepts_mz_tables() {
1836 assert!(validate_system_catalog_query("SELECT * FROM mz_tables").is_ok());
1837 assert!(validate_system_catalog_query("SELECT * FROM mz_internal.mz_comments").is_ok());
1838 assert!(
1839 validate_system_catalog_query(
1840 "SELECT * FROM mz_tables t JOIN mz_columns c ON t.id = c.id"
1841 )
1842 .is_ok()
1843 );
1844 }
1845
1846 #[mz_ore::test]
1847 fn test_validate_system_catalog_query_subqueries() {
1848 assert!(
1850 validate_system_catalog_query(
1851 "SELECT * FROM mz_tables WHERE id IN (SELECT id FROM mz_columns)"
1852 )
1853 .is_ok()
1854 );
1855
1856 assert!(validate_system_catalog_query(
1858 "SELECT * FROM mz_tables WHERE id IN (SELECT table_id FROM mz_columns WHERE type IN (SELECT id FROM mz_types))"
1859 )
1860 .is_ok());
1861
1862 assert!(
1864 validate_system_catalog_query(
1865 "SELECT * FROM (SELECT * FROM mz_tables WHERE name LIKE 'test%') AS t"
1866 )
1867 .is_ok()
1868 );
1869
1870 assert!(
1872 validate_system_catalog_query(
1873 "SELECT * FROM mz_tables WHERE id IN (SELECT table_id FROM user_data)"
1874 )
1875 .is_err()
1876 );
1877
1878 assert!(validate_system_catalog_query(
1880 "SELECT * FROM mz_tables WHERE id IN (SELECT id FROM (SELECT id FROM user_table) AS t)"
1881 )
1882 .is_err());
1883 }
1884
1885 #[mz_ore::test]
1886 fn test_validate_system_catalog_query_rejects_user_tables() {
1887 assert!(validate_system_catalog_query("SELECT * FROM user_data").is_err());
1888 assert!(validate_system_catalog_query("SELECT * FROM my_table").is_err());
1889 assert!(
1891 validate_system_catalog_query("SELECT * FROM user_data WHERE 'mz_' IS NOT NULL")
1892 .is_err()
1893 );
1894 }
1895
1896 #[mz_ore::test]
1897 fn test_validate_system_catalog_query_allows_functions() {
1898 assert!(
1900 validate_system_catalog_query(
1901 "SELECT date_part('year', now())::int4 AS y FROM mz_tables LIMIT 1"
1902 )
1903 .is_ok()
1904 );
1905 assert!(validate_system_catalog_query("SELECT length(name) FROM mz_tables").is_ok());
1906 assert!(
1907 validate_system_catalog_query(
1908 "SELECT count(*) FROM mz_sources WHERE now() > created_at"
1909 )
1910 .is_ok()
1911 );
1912 }
1913
1914 #[mz_ore::test]
1915 fn test_validate_system_catalog_query_unqualified_pg_catalog() {
1916 assert!(validate_system_catalog_query("SELECT * FROM pg_class").is_ok());
1919 assert!(validate_system_catalog_query("SELECT nspname FROM pg_namespace").is_ok());
1920 }
1921
1922 #[mz_ore::test]
1923 fn test_validate_system_catalog_query_schema_qualified() {
1924 assert!(validate_system_catalog_query("SELECT * FROM mz_catalog.mz_tables").is_ok());
1926 assert!(validate_system_catalog_query("SELECT * FROM mz_internal.mz_sessions").is_ok());
1927 assert!(validate_system_catalog_query("SELECT * FROM pg_catalog.pg_type").is_ok());
1928 assert!(validate_system_catalog_query("SELECT * FROM information_schema.tables").is_ok());
1929
1930 assert!(validate_system_catalog_query("SELECT * FROM public.user_table").is_err());
1932 assert!(validate_system_catalog_query("SELECT * FROM myschema.mytable").is_err());
1933
1934 assert!(
1936 validate_system_catalog_query("SELECT * FROM mz_unsafe.mz_some_table").is_err(),
1937 "mz_unsafe schema should be blocked even though it is a system schema"
1938 );
1939
1940 assert!(
1942 validate_system_catalog_query(
1943 "SELECT * FROM mz_catalog.mz_tables JOIN public.user_data ON true"
1944 )
1945 .is_err()
1946 );
1947 }
1948
1949 #[mz_ore::test]
1950 fn test_validate_system_catalog_query_adversarial_cases() {
1951 assert!(
1953 validate_system_catalog_query(
1954 "WITH user_cte AS (SELECT * FROM user_data) \
1955 SELECT * FROM mz_tables, user_cte"
1956 )
1957 .is_err(),
1958 "Should reject CTE referencing user table"
1959 );
1960
1961 assert!(
1963 validate_system_catalog_query(
1964 "WITH \
1965 cte1 AS (SELECT * FROM mz_tables), \
1966 cte2 AS (SELECT * FROM cte1), \
1967 cte3 AS (SELECT * FROM user_data) \
1968 SELECT * FROM cte2"
1969 )
1970 .is_err(),
1971 "Should reject CTE chain with user table"
1972 );
1973
1974 assert!(
1976 validate_system_catalog_query(
1977 "SELECT * FROM mz_tables t1 \
1978 JOIN user_data u ON t1.id = u.id \
1979 JOIN mz_sources s ON t1.id = s.id"
1980 )
1981 .is_err(),
1982 "Should reject multi-join with user table"
1983 );
1984
1985 assert!(
1987 validate_system_catalog_query(
1988 "SELECT * FROM mz_tables t \
1989 LEFT JOIN user_data u ON t.id = u.table_id \
1990 WHERE u.id IS NULL"
1991 )
1992 .is_err(),
1993 "Should reject LEFT JOIN with user table"
1994 );
1995
1996 assert!(
1998 validate_system_catalog_query(
1999 "SELECT * FROM mz_tables WHERE id IN \
2000 (SELECT table_id FROM (SELECT * FROM user_data) AS u)"
2001 )
2002 .is_err(),
2003 "Should reject nested subquery with user table"
2004 );
2005
2006 assert!(
2008 validate_system_catalog_query(
2009 "SELECT name FROM mz_tables \
2010 UNION \
2011 SELECT name FROM user_data"
2012 )
2013 .is_err(),
2014 "Should reject UNION with user table"
2015 );
2016
2017 assert!(
2019 validate_system_catalog_query(
2020 "SELECT id FROM mz_sources \
2021 UNION ALL \
2022 SELECT id FROM products"
2023 )
2024 .is_err(),
2025 "Should reject UNION ALL with user table"
2026 );
2027
2028 assert!(
2030 validate_system_catalog_query("SELECT * FROM mz_tables CROSS JOIN user_data").is_err(),
2031 "Should reject CROSS JOIN with user table"
2032 );
2033
2034 assert!(
2036 validate_system_catalog_query(
2037 "SELECT t.*, (SELECT COUNT(*) FROM user_data) AS cnt FROM mz_tables t"
2038 )
2039 .is_err(),
2040 "Should reject subquery in SELECT with user table"
2041 );
2042
2043 assert!(
2045 validate_system_catalog_query("SELECT * FROM mz_catalogg.fake_table").is_err(),
2046 "Should reject typo-squatting schema name"
2047 );
2048 assert!(
2049 validate_system_catalog_query("SELECT * FROM mz_catalog_hack.fake_table").is_err(),
2050 "Should reject fake schema with mz_catalog prefix"
2051 );
2052
2053 assert!(
2055 validate_system_catalog_query(
2056 "SELECT * FROM mz_tables t, LATERAL (SELECT * FROM user_data WHERE id = t.id) u"
2057 )
2058 .is_err(),
2059 "Should reject LATERAL join with user table"
2060 );
2061
2062 assert!(
2064 validate_system_catalog_query(
2065 "WITH \
2066 tables AS (SELECT * FROM mz_tables), \
2067 sources AS (SELECT * FROM mz_sources) \
2068 SELECT t.name, s.name \
2069 FROM tables t \
2070 JOIN sources s ON t.id = s.id \
2071 WHERE t.id IN (SELECT id FROM mz_columns)"
2072 )
2073 .is_ok(),
2074 "Should allow complex query with only system tables"
2075 );
2076
2077 assert!(
2079 validate_system_catalog_query(
2080 "SELECT name FROM mz_tables \
2081 UNION \
2082 SELECT name FROM mz_sources"
2083 )
2084 .is_ok(),
2085 "Should allow UNION of system tables"
2086 );
2087 }
2088
2089 #[mz_ore::test]
2090 fn test_validate_system_catalog_query_rejects_constant_queries() {
2091 assert!(
2094 validate_system_catalog_query("SELECT 1").is_err(),
2095 "Should reject constant SELECT with no table references"
2096 );
2097 assert!(
2098 validate_system_catalog_query("SELECT 1 + 2, 'hello'").is_err(),
2099 "Should reject constant expression SELECT"
2100 );
2101 assert!(
2102 validate_system_catalog_query("SELECT now()").is_err(),
2103 "Should reject function-only SELECT with no table references"
2104 );
2105 }
2106
2107 #[mz_ore::test]
2108 fn test_validate_system_catalog_query_rejects_mixed_tables() {
2109 assert!(
2110 validate_system_catalog_query(
2111 "SELECT * FROM mz_tables t JOIN user_data u ON t.id = u.table_id"
2112 )
2113 .is_err()
2114 );
2115 }
2116
2117 #[mz_ore::test]
2118 fn test_validate_system_catalog_query_allows_show() {
2119 assert!(
2121 validate_system_catalog_query("SHOW TABLES FROM mz_internal").is_ok(),
2122 "SHOW TABLES FROM mz_internal should be allowed"
2123 );
2124 assert!(
2125 validate_system_catalog_query("SHOW TABLES FROM mz_catalog").is_ok(),
2126 "SHOW TABLES FROM mz_catalog should be allowed"
2127 );
2128 assert!(
2129 validate_system_catalog_query("SHOW CLUSTERS").is_ok(),
2130 "SHOW CLUSTERS should be allowed"
2131 );
2132 assert!(
2133 validate_system_catalog_query("SHOW SOURCES").is_ok(),
2134 "SHOW SOURCES should be allowed"
2135 );
2136 assert!(
2137 validate_system_catalog_query("SHOW TABLES").is_ok(),
2138 "SHOW TABLES should be allowed"
2139 );
2140 }
2141
2142 #[mz_ore::test]
2143 fn test_validate_system_catalog_query_allows_explain() {
2144 assert!(
2145 validate_system_catalog_query("EXPLAIN SELECT * FROM mz_tables").is_ok(),
2146 "EXPLAIN of system table query should be allowed"
2147 );
2148 assert!(
2149 validate_system_catalog_query("EXPLAIN SELECT 1").is_ok(),
2150 "EXPLAIN SELECT 1 should be allowed"
2151 );
2152 }
2153
2154 #[mz_ore::test(tokio::test)]
2157 async fn test_tools_list_agent_query_tool_disabled() {
2158 let result = handle_tools_list(McpEndpointType::Agent, false, true, 1_000_000).unwrap();
2159 let McpResult::ToolsList(list) = result else {
2160 panic!("Expected ToolsList result");
2161 };
2162 let tool_names: Vec<&str> = list.tools.iter().map(|t| t.name.as_str()).collect();
2163 assert!(
2164 tool_names.contains(&"get_data_products"),
2165 "get_data_products should always be present"
2166 );
2167 assert!(
2168 tool_names.contains(&"get_data_product_details"),
2169 "get_data_product_details should always be present"
2170 );
2171 assert!(
2172 tool_names.contains(&"read_data_product"),
2173 "read_data_product should be present when its flag is on"
2174 );
2175 assert!(
2176 !tool_names.contains(&"query"),
2177 "query tool should be hidden when disabled"
2178 );
2179 }
2180
2181 #[mz_ore::test(tokio::test)]
2182 async fn test_tools_list_agent_query_tool_enabled() {
2183 let result = handle_tools_list(McpEndpointType::Agent, true, true, 1_000_000).unwrap();
2184 let McpResult::ToolsList(list) = result else {
2185 panic!("Expected ToolsList result");
2186 };
2187 let tool_names: Vec<&str> = list.tools.iter().map(|t| t.name.as_str()).collect();
2188 assert!(
2189 tool_names.contains(&"get_data_products"),
2190 "get_data_products should always be present"
2191 );
2192 assert!(
2193 tool_names.contains(&"get_data_product_details"),
2194 "get_data_product_details should always be present"
2195 );
2196 assert!(
2197 tool_names.contains(&"read_data_product"),
2198 "read_data_product should be present when its flag is on"
2199 );
2200 assert!(
2201 tool_names.contains(&"query"),
2202 "query tool should be present when enabled"
2203 );
2204 }
2205
2206 #[mz_ore::test(tokio::test)]
2207 async fn test_tools_list_agent_read_data_product_tool_disabled() {
2208 let result = handle_tools_list(McpEndpointType::Agent, true, false, 1_000_000).unwrap();
2209 let McpResult::ToolsList(list) = result else {
2210 panic!("Expected ToolsList result");
2211 };
2212 let tool_names: Vec<&str> = list.tools.iter().map(|t| t.name.as_str()).collect();
2213 assert!(
2214 tool_names.contains(&"get_data_products"),
2215 "get_data_products should always be present"
2216 );
2217 assert!(
2218 tool_names.contains(&"get_data_product_details"),
2219 "get_data_product_details should always be present"
2220 );
2221 assert!(
2222 !tool_names.contains(&"read_data_product"),
2223 "read_data_product should be hidden when disabled"
2224 );
2225 assert!(
2226 tool_names.contains(&"query"),
2227 "query tool should remain present when enabled"
2228 );
2229 }
2230
2231 #[mz_ore::test(tokio::test)]
2236 async fn test_tools_list_agent_both_read_tools_disabled() {
2237 let result = handle_tools_list(McpEndpointType::Agent, false, false, 1_000_000).unwrap();
2238 let McpResult::ToolsList(list) = result else {
2239 panic!("Expected ToolsList result");
2240 };
2241 let tool_names: Vec<&str> = list.tools.iter().map(|t| t.name.as_str()).collect();
2242 assert_eq!(
2243 tool_names
2244 .iter()
2245 .copied()
2246 .collect::<std::collections::BTreeSet<_>>(),
2247 ["get_data_product_details", "get_data_products"]
2248 .into_iter()
2249 .collect(),
2250 "only discovery tools should be advertised when both read flags are off",
2251 );
2252
2253 let instructions = endpoint_instructions(McpEndpointType::Agent, false, false)
2254 .expect("agent instructions must be present");
2255 assert!(
2256 !instructions.contains("Use the `query` tool"),
2257 "instructions must not point at query when it is hidden: {instructions}",
2258 );
2259 assert!(
2260 !instructions.contains("`read_data_product` automatically"),
2261 "instructions must not point at read_data_product when it is hidden: {instructions}",
2262 );
2263 assert!(
2264 instructions.contains("discovery only"),
2265 "instructions must tell the agent it is discovery-only: {instructions}",
2266 );
2267 }
2268
2269 #[mz_ore::test(tokio::test)]
2270 async fn test_tools_list_developer_query_tool_disabled() {
2271 let result = handle_tools_list(McpEndpointType::Developer, false, true, 1_000_000).unwrap();
2274 let McpResult::ToolsList(list) = result else {
2275 panic!("Expected ToolsList result");
2276 };
2277 let tool_names: Vec<&str> = list.tools.iter().map(|t| t.name.as_str()).collect();
2278 assert!(
2279 tool_names.contains(&"query_system_catalog"),
2280 "query_system_catalog should always be present on developer"
2281 );
2282 assert!(
2283 !tool_names.contains(&"query"),
2284 "query tool should be hidden when disabled"
2285 );
2286 }
2287
2288 #[mz_ore::test(tokio::test)]
2289 async fn test_tools_list_developer_query_tool_enabled() {
2290 let result = handle_tools_list(McpEndpointType::Developer, true, true, 1_000_000).unwrap();
2291 let McpResult::ToolsList(list) = result else {
2292 panic!("Expected ToolsList result");
2293 };
2294 let tool_names: Vec<&str> = list.tools.iter().map(|t| t.name.as_str()).collect();
2295 assert!(
2296 tool_names.contains(&"query_system_catalog"),
2297 "query_system_catalog should always be present on developer"
2298 );
2299 assert!(
2300 tool_names.contains(&"query"),
2301 "query tool should be present on developer when enabled"
2302 );
2303
2304 let instructions = endpoint_instructions(McpEndpointType::Developer, true, true)
2305 .expect("developer instructions must be present");
2306 assert!(
2307 instructions.contains("through the query tool with its cluster argument"),
2308 "instructions must route mz_introspection through query: {instructions}",
2309 );
2310 }
2311
2312 #[mz_ore::test]
2313 fn test_developer_instructions_query_tool_disabled() {
2314 let instructions = endpoint_instructions(McpEndpointType::Developer, false, true)
2315 .expect("developer instructions must be present");
2316 assert!(
2317 instructions.contains("cluster-targeted introspection is unavailable")
2318 && !instructions.contains("through the query tool"),
2319 "instructions must not route to the hidden query tool: {instructions}",
2320 );
2321 }
2322
2323 #[mz_ore::test]
2326 fn test_tools_list_accepts_optional_params() {
2327 for body in [
2328 r#"{"jsonrpc":"2.0","id":1,"method":"tools/list"}"#,
2329 r#"{"jsonrpc":"2.0","id":1,"method":"tools/list","params":null}"#,
2330 r#"{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}"#,
2331 r#"{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"_meta":{"progressToken":0}}}"#,
2332 ] {
2333 let req: McpRequest = serde_json::from_str(body)
2334 .unwrap_or_else(|e| panic!("tools/list must deserialize: {body}: {e}"));
2335 assert!(
2336 matches!(req.method, McpMethod::ToolsList(_)),
2337 "expected ToolsList for {body}"
2338 );
2339 }
2340 }
2341
2342 #[mz_ore::test]
2346 fn test_named_methods_accept_optional_params() {
2347 for (body, want_ping) in [
2348 (r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#, true),
2349 (
2350 r#"{"jsonrpc":"2.0","id":1,"method":"ping","params":{}}"#,
2351 true,
2352 ),
2353 (
2354 r#"{"jsonrpc":"2.0","id":1,"method":"ping","params":{"_meta":{"progressToken":0}}}"#,
2355 true,
2356 ),
2357 (
2358 r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#,
2359 false,
2360 ),
2361 (
2362 r#"{"jsonrpc":"2.0","method":"notifications/initialized","params":{"_meta":{"x":1}}}"#,
2363 false,
2364 ),
2365 ] {
2366 let req: McpRequest = serde_json::from_str(body)
2367 .unwrap_or_else(|e| panic!("must deserialize: {body}: {e}"));
2368 if want_ping {
2369 assert!(matches!(req.method, McpMethod::Ping(_)), "for {body}");
2370 } else {
2371 assert!(
2372 matches!(req.method, McpMethod::NotificationsInitialized(_)),
2373 "for {body}"
2374 );
2375 }
2376 }
2377 }
2378
2379 #[mz_ore::test]
2380 fn test_tools_call_arguments_optional() {
2381 for body in [
2382 r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"get_data_products"}}"#,
2383 r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"get_data_products","arguments":{}}}"#,
2384 r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"get_data_products","arguments":null}}"#,
2385 ] {
2386 let req: McpRequest = serde_json::from_str(body)
2387 .unwrap_or_else(|e| panic!("must deserialize: {body}: {e}"));
2388 assert!(
2389 matches!(
2390 req.method,
2391 McpMethod::ToolsCall(ToolsCallParams::GetDataProducts(NoArguments {}))
2392 ),
2393 "for {body}"
2394 );
2395 }
2396
2397 let err = serde_json::from_str::<McpRequest>(
2399 r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"query_system_catalog"}}"#,
2400 )
2401 .expect_err("query_system_catalog requires sql_query");
2402 assert!(err.to_string().contains("sql_query"), "{err}");
2403 }
2404
2405 #[mz_ore::test]
2408 fn test_format_rows_response_within_limit() {
2409 let rows = vec![vec![json!("a"), json!(1)], vec![json!("b"), json!(2)]];
2410 let result = format_rows_response(raw_rows(rows), 1_000_000).unwrap();
2411 let McpResult::ToolContent(content) = result else {
2412 panic!("Expected ToolContent");
2413 };
2414 assert_eq!(content.content.len(), 1);
2415 assert!(content.content[0].text.contains("\"a\""));
2416 assert!(content.content[0].text.contains("\"b\""));
2417 }
2418
2419 #[mz_ore::test]
2420 fn test_format_rows_response_errors_when_over_limit() {
2421 let rows: Vec<Vec<serde_json::Value>> = (0..100)
2422 .map(|i| vec![json!(format!("row_{}", i)), json!(i)])
2423 .collect();
2424 let err = format_rows_response(raw_rows(rows), 500).unwrap_err();
2425 let msg = err.to_string();
2426 assert!(
2427 msg.contains("exceeds the 500 byte limit"),
2428 "Error should mention the size limit, got: {msg}"
2429 );
2430 assert!(
2431 msg.contains("Use LIMIT or WHERE"),
2432 "Error should suggest narrowing the query, got: {msg}"
2433 );
2434 }
2435
2436 #[mz_ore::test]
2437 fn test_format_rows_response_empty_rows() {
2438 let rows: Vec<Vec<serde_json::Value>> = vec![];
2439 let result = format_rows_response(raw_rows(rows), 1000).unwrap();
2440 let McpResult::ToolContent(content) = result else {
2441 panic!("Expected ToolContent");
2442 };
2443 assert_eq!(content.content.len(), 1);
2444 assert_eq!(content.content[0].text, "[]");
2445 }
2446
2447 #[mz_ore::test]
2450 fn test_safe_data_product_name_valid() {
2451 assert_eq!(
2453 safe_data_product_name(r#""materialize"."public"."my_view""#).unwrap(),
2454 r#""materialize"."public"."my_view""#
2455 );
2456 assert_eq!(
2458 safe_data_product_name(r#""public"."my_view""#).unwrap(),
2459 r#""public"."my_view""#
2460 );
2461 assert_eq!(safe_data_product_name("my_view").unwrap(), r#""my_view""#);
2463 }
2464
2465 #[mz_ore::test]
2466 fn test_safe_data_product_name_rejects_empty() {
2467 assert!(safe_data_product_name("").is_err());
2468 assert!(safe_data_product_name(" ").is_err());
2469 }
2470
2471 #[mz_ore::test]
2475 fn test_safe_data_product_name_enforces_size_limit() {
2476 use mz_sql_parser::parser::MAX_STATEMENT_BATCH_SIZE;
2477 let oversized: String = "(".repeat(MAX_STATEMENT_BATCH_SIZE + 1);
2478 let err = safe_data_product_name(&oversized).expect_err("should be rejected");
2479 let msg = err.to_string();
2480 assert!(
2481 msg.contains("statement batch size cannot exceed"),
2482 "expected size-guard error, got: {msg}"
2483 );
2484 }
2485
2486 #[mz_ore::test]
2487 fn test_safe_data_product_name_rejects_sql_injection() {
2488 assert!(safe_data_product_name("my_view; DROP TABLE users").is_err());
2490 assert!(safe_data_product_name("my_view UNION SELECT * FROM secrets").is_err());
2492 assert!(safe_data_product_name("my_view, secrets").is_err());
2494 assert!(safe_data_product_name("my_view WHERE 1=1 --").is_err());
2496 }
2497
2498 #[mz_ore::test]
2500 fn test_read_only_txn_comment_cannot_swallow_commit() {
2501 let sql = read_only_txn("SET CLUSTER = 'c'", "SELECT 1 --");
2502 assert!(
2503 sql.contains("\n; COMMIT;"),
2504 "COMMIT must sit on its own line: {sql}",
2505 );
2506 }
2507
2508 #[mz_ore::test]
2512 fn test_query_set_clause_without_replica() {
2513 let clause = query_set_clause("prod_cluster", None);
2514 assert_eq!(clause, "SET CLUSTER = 'prod_cluster'");
2515 }
2516
2517 #[mz_ore::test]
2521 fn test_query_set_clause_with_replica() {
2522 let clause = query_set_clause("prod_cluster", Some("r1"));
2523 assert_eq!(
2524 clause,
2525 "SET CLUSTER = 'prod_cluster'; SET CLUSTER_REPLICA = 'r1'"
2526 );
2527 }
2528
2529 #[mz_ore::test]
2533 fn test_query_set_clause_escapes_replica_name() {
2534 let clause = query_set_clause("c", Some("evil'; DROP TABLE secrets; --"));
2535 assert_eq!(
2536 clause,
2537 "SET CLUSTER = 'c'; SET CLUSTER_REPLICA = 'evil''; DROP TABLE secrets; --'"
2538 );
2539 }
2540
2541 #[mz_ore::test]
2544 fn test_validate_cluster_replica_accepts_none_and_names() {
2545 assert!(validate_cluster_replica(None).is_ok());
2546 assert!(validate_cluster_replica(Some("r1")).is_ok());
2547 }
2548
2549 #[mz_ore::test]
2553 fn test_validate_cluster_replica_rejects_empty() {
2554 for name in ["", " ", "\t\n"] {
2555 assert!(
2556 matches!(
2557 validate_cluster_replica(Some(name)),
2558 Err(McpRequestError::QueryValidationFailed(_))
2559 ),
2560 "expected validation error for {name:?}",
2561 );
2562 }
2563 }
2564
2565 #[mz_ore::test]
2571 fn test_build_read_query_with_cluster() {
2572 let sql = build_read_query("\"db\".\"sch\".\"v\"", 50, Some("prod_cluster"));
2573 assert!(sql.contains("BEGIN READ ONLY"), "{sql}");
2574 assert!(sql.contains("SET CLUSTER = 'prod_cluster'"), "{sql}");
2575 assert!(
2576 sql.contains("SELECT * FROM \"db\".\"sch\".\"v\" LIMIT 50"),
2577 "{sql}",
2578 );
2579 assert!(sql.contains("COMMIT"), "{sql}");
2580 }
2581
2582 #[mz_ore::test]
2585 fn test_build_read_query_without_cluster() {
2586 let sql = build_read_query("\"db\".\"sch\".\"v\"", 50, None);
2587 assert!(sql.contains("BEGIN READ ONLY"), "{sql}");
2588 assert!(!sql.contains("SET CLUSTER"), "{sql}");
2589 assert!(
2590 sql.contains("SELECT * FROM \"db\".\"sch\".\"v\" LIMIT 50"),
2591 "{sql}",
2592 );
2593 assert!(sql.contains("COMMIT"), "{sql}");
2594 }
2595
2596 #[mz_ore::test]
2601 fn test_build_read_query_escapes_cluster_name() {
2602 let sql = build_read_query(
2603 "\"db\".\"sch\".\"v\"",
2604 10,
2605 Some("evil'; DROP TABLE secrets; --"),
2606 );
2607 assert!(
2609 sql.contains("SET CLUSTER = 'evil''; DROP TABLE secrets; --'"),
2610 "single quote should be doubled inside the literal: {sql}",
2611 );
2612 assert_eq!(
2615 sql.matches("SET CLUSTER").count(),
2616 1,
2617 "exactly one SET CLUSTER statement: {sql}",
2618 );
2619 assert_eq!(
2620 sql.matches("DROP TABLE").count(),
2621 1,
2622 "DROP TABLE should appear once, inside the quoted literal: {sql}",
2623 );
2624 }
2625
2626 #[mz_ore::test]
2627 fn test_mcp_error_codes() {
2628 assert_eq!(
2629 McpRequestError::InvalidJsonRpcVersion.error_code(),
2630 error_codes::INVALID_REQUEST
2631 );
2632 assert_eq!(
2633 McpRequestError::MethodNotFound("test".to_string()).error_code(),
2634 error_codes::METHOD_NOT_FOUND
2635 );
2636 assert_eq!(
2637 McpRequestError::QueryExecutionFailed("test".to_string()).error_code(),
2638 error_codes::INTERNAL_ERROR
2639 );
2640 }
2641}