Skip to main content

zeph_core/agent/
mcp.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4use rmcp::model::{ElicitResult, ElicitationAction};
5
6use super::{Agent, Channel, LlmProvider};
7
8impl<C: Channel> Agent<C> {
9    /// Dispatch a `/mcp` subcommand, returning the output as a `String`.
10    ///
11    /// All output is collected into the returned string; no channel sends are
12    /// performed.  This makes the future `Send`-compatible for use in
13    /// `AgentAccess::handle_mcp`.
14    #[tracing::instrument(skip_all, name = "core.agent.handle_mcp_command")]
15    pub(super) async fn handle_mcp_command(
16        &mut self,
17        args: &str,
18    ) -> Result<String, super::error::AgentError> {
19        let parts: Vec<&str> = args.split_whitespace().collect();
20        match parts.first().copied() {
21            Some("add") => self.handle_mcp_add(&parts[1..]).await,
22            Some("list") => self.handle_mcp_list().await,
23            Some("tools") => Ok(self.handle_mcp_tools(parts.get(1).copied())),
24            Some("remove") => self.handle_mcp_remove(parts.get(1).copied()).await,
25            _ => Ok("Usage: /mcp add|list|tools|remove".to_owned()),
26        }
27    }
28
29    async fn handle_mcp_add(&mut self, args: &[&str]) -> Result<String, super::error::AgentError> {
30        if args.len() < 2 {
31            return Ok("Usage: /mcp add <id> <command> [args...] | /mcp add <id> <url>".to_owned());
32        }
33
34        // Clone the Arc so no borrow of self.services.mcp.manager is held across .await.
35        let Some(manager) = self.services.mcp.manager.clone() else {
36            return Ok("MCP is not enabled.".to_owned());
37        };
38
39        let target = args[1];
40        if let Some(err) = validate_mcp_command(target, &self.services.mcp.allowed_commands) {
41            return Ok(err);
42        }
43
44        // SEC-MCP-03: enforce server limit
45        let current_count = manager.list_servers().await.len();
46        if current_count >= self.services.mcp.max_dynamic {
47            return Ok(format!(
48                "Server limit reached ({}/{}).",
49                current_count, self.services.mcp.max_dynamic
50            ));
51        }
52
53        let entry = build_server_entry(args[0], target, &args[2..]);
54
55        match manager.add_server(&entry).await {
56            Ok(tools) => {
57                let count = tools.len();
58                self.services
59                    .mcp
60                    .server_outcomes
61                    .push(zeph_mcp::ServerConnectOutcome {
62                        id: entry.id.clone(),
63                        connected: true,
64                        tool_count: count,
65                        error: String::new(),
66                    });
67                self.services.mcp.tools.extend(tools);
68                self.services.mcp.sync_executor_tools();
69                self.services.mcp.pruning_cache.reset();
70                // Defer rebuild to check_tool_refresh (next turn) so this method
71                // stays Send-compatible for use in AgentAccess::handle_mcp.
72                self.services.mcp.pending_semantic_rebuild = true;
73                self.update_mcp_metrics();
74                Ok(format!(
75                    "Connected MCP server '{}' ({count} tool(s))",
76                    entry.id
77                ))
78            }
79            Err(e) => {
80                tracing::warn!(server_id = entry.id, "MCP add failed: {e:#}");
81                Ok(format!("Failed to connect server '{}': {e}", entry.id))
82            }
83        }
84    }
85
86    async fn handle_mcp_list(&mut self) -> Result<String, super::error::AgentError> {
87        use std::fmt::Write;
88
89        let Some(manager) = self.services.mcp.manager.clone() else {
90            return Ok("MCP is not enabled.".to_owned());
91        };
92
93        let server_ids = manager.list_servers().await;
94        if server_ids.is_empty() {
95            return Ok("No MCP servers connected.".to_owned());
96        }
97
98        let mut output = String::from("Connected MCP servers:\n");
99        let mut total = 0usize;
100        for id in &server_ids {
101            let count = self
102                .services
103                .mcp
104                .tools
105                .iter()
106                .filter(|t| t.server_id == *id)
107                .count();
108            total += count;
109            let _ = writeln!(output, "- {id} ({count} tools)");
110        }
111        let _ = write!(output, "Total: {total} tool(s)");
112
113        Ok(output)
114    }
115
116    fn handle_mcp_tools(&mut self, server_id: Option<&str>) -> String {
117        use std::fmt::Write;
118
119        let Some(server_id) = server_id else {
120            return "Usage: /mcp tools <server_id>".to_owned();
121        };
122
123        let tools: Vec<_> = self
124            .services
125            .mcp
126            .tools
127            .iter()
128            .filter(|t| t.server_id == server_id)
129            .collect();
130
131        if tools.is_empty() {
132            return format!("No tools found for server '{server_id}'.");
133        }
134
135        let mut output = format!("Tools for '{server_id}' ({} total):\n", tools.len());
136        for t in &tools {
137            if t.description.is_empty() {
138                let _ = writeln!(output, "- {}", t.name);
139            } else {
140                let _ = writeln!(output, "- {} — {}", t.name, t.description);
141            }
142        }
143        output
144    }
145
146    async fn handle_mcp_remove(
147        &mut self,
148        server_id: Option<&str>,
149    ) -> Result<String, super::error::AgentError> {
150        let Some(server_id) = server_id else {
151            return Ok("Usage: /mcp remove <id>".to_owned());
152        };
153
154        // Clone the Arc so no borrow of self.services.mcp.manager is held across .await.
155        let Some(manager) = self.services.mcp.manager.clone() else {
156            return Ok("MCP is not enabled.".to_owned());
157        };
158
159        match manager.remove_server(server_id).await {
160            Ok(()) => {
161                let before = self.services.mcp.tools.len();
162                self.services.mcp.tools.retain(|t| t.server_id != server_id);
163                let removed = before - self.services.mcp.tools.len();
164                self.services
165                    .mcp
166                    .server_outcomes
167                    .retain(|o| o.id != server_id);
168                self.services.mcp.sync_executor_tools();
169                self.services.mcp.pruning_cache.reset();
170                // Defer rebuild to check_tool_refresh (next turn) so this method
171                // stays Send-compatible for use in AgentAccess::handle_mcp.
172                self.services.mcp.pending_semantic_rebuild = true;
173                self.update_mcp_metrics();
174                let sid = server_id.to_owned();
175                self.update_metrics(|m| {
176                    m.active_mcp_tools
177                        .retain(|name| !name.starts_with(&format!("{sid}:")));
178                });
179                Ok(format!(
180                    "Disconnected MCP server '{server_id}' (removed {removed} tools)"
181                ))
182            }
183            Err(e) => {
184                tracing::warn!(server_id, "MCP remove failed: {e:#}");
185                Ok(format!("Failed to remove server '{server_id}': {e}"))
186            }
187        }
188    }
189
190    pub(super) async fn append_mcp_prompt(&mut self, query: &str, system_prompt: &mut String) {
191        let matched_tools = self.match_mcp_tools(query).await;
192        let active_mcp: Vec<String> = matched_tools
193            .iter()
194            .map(zeph_mcp::McpTool::qualified_name)
195            .collect();
196        let mcp_total = self.services.mcp.tools.len();
197        let (mcp_server_count, mcp_connected_count) =
198            if self.services.mcp.server_outcomes.is_empty() {
199                let connected = self
200                    .services
201                    .mcp
202                    .tools
203                    .iter()
204                    .map(|t| &t.server_id)
205                    .collect::<std::collections::HashSet<_>>()
206                    .len();
207                (connected, connected)
208            } else {
209                let total = self.services.mcp.server_outcomes.len();
210                let connected = self
211                    .services
212                    .mcp
213                    .server_outcomes
214                    .iter()
215                    .filter(|o| o.connected)
216                    .count();
217                (total, connected)
218            };
219        self.update_metrics(|m| {
220            m.active_mcp_tools = active_mcp;
221            m.mcp_tool_count = mcp_total;
222            m.mcp_server_count = mcp_server_count;
223            m.mcp_connected_count = mcp_connected_count;
224        });
225        if let Some(ref manager) = self.services.mcp.manager {
226            let instructions = manager.all_server_instructions().await;
227            if !instructions.is_empty() {
228                system_prompt.push_str("\n\n");
229                system_prompt.push_str(&instructions);
230            }
231        }
232        if !matched_tools.is_empty() {
233            let tool_names: Vec<&str> = matched_tools.iter().map(|t| t.name.as_str()).collect();
234            tracing::debug!(
235                skills = ?self.services.skill.active_skill_names,
236                mcp_tools = ?tool_names,
237                "matched items"
238            );
239            let tools_prompt = zeph_mcp::format_mcp_tools_prompt(&matched_tools);
240            if !tools_prompt.is_empty() {
241                system_prompt.push_str("\n\n");
242                system_prompt.push_str(&tools_prompt);
243            }
244        }
245    }
246
247    async fn match_mcp_tools(&self, query: &str) -> Vec<zeph_mcp::McpTool> {
248        let Some(ref registry) = self.services.mcp.registry else {
249            return self.services.mcp.tools.clone();
250        };
251        let provider = self.embedding_provider.clone();
252        registry
253            .search(query, self.services.skill.max_active_skills, |text| {
254                let owned = text.to_owned();
255                let p = provider.clone();
256                Box::pin(async move { p.embed(&owned).await })
257            })
258            .await
259    }
260
261    /// Poll the watch receiver for tool list updates from `tools/list_changed` notifications,
262    /// and process any deferred semantic index rebuild requests.
263    ///
264    /// Called once per agent turn, before processing user input.  Two triggers cause a rebuild:
265    /// - A `tools/list_changed` notification from an MCP server (via `tool_rx`).
266    /// - `pending_semantic_rebuild == true`, set by `/mcp add` or `/mcp remove` when dispatched
267    ///   via `AgentAccess::handle_mcp` (which cannot call `rebuild_semantic_index` directly
268    ///   because the future would be `!Send`).
269    ///
270    /// Both branches also call [`refresh_mcp_tool_ids`](Self::refresh_mcp_tool_ids) so
271    /// `TrustGateExecutor`'s Quarantine-deny set stays current with MCP servers connected
272    /// after startup (#5747) — otherwise it is only ever populated once, at agent construction.
273    ///
274    /// If neither trigger fires, this is a no-op.
275    pub(super) async fn check_tool_refresh(&mut self) {
276        // Handle deferred rebuild from /mcp add|remove via AgentAccess.
277        if self.services.mcp.pending_semantic_rebuild {
278            self.services.mcp.pending_semantic_rebuild = false;
279            self.refresh_mcp_tool_ids();
280            self.rebuild_semantic_index().await;
281            self.sync_mcp_registry().await;
282            self.refresh_shadow_sentinel_mcp_tool_ids();
283            let mcp_total = self.services.mcp.tools.len();
284            let mcp_servers = self
285                .services
286                .mcp
287                .tools
288                .iter()
289                .map(|t| &t.server_id)
290                .collect::<std::collections::HashSet<_>>()
291                .len();
292            self.update_metrics(|m| {
293                m.mcp_tool_count = mcp_total;
294                m.mcp_server_count = mcp_servers;
295            });
296        }
297
298        let Some(ref mut rx) = self.services.mcp.tool_rx else {
299            return;
300        };
301        if !rx.has_changed().unwrap_or(false) {
302            return;
303        }
304        let new_tools = rx.borrow_and_update().clone();
305        if new_tools.is_empty() {
306            // Guard against replacing a non-empty initial tool list with the watch's empty
307            // initial value. The watch is only updated after a real tools/list_changed event.
308            //
309            // This early return also means `refresh_mcp_tool_ids` below is skipped, so
310            // `mcp_tool_ids` can retain stale ids past this point — but only fail-safe (over-deny
311            // of ids that no longer map to a live tool), never fail-open. In practice this branch
312            // is unreachable with an empty list anyway: the `tools/list_changed` producer never
313            // sends an empty vec, and removing the last server goes through the
314            // `pending_semantic_rebuild` branch above, which has no such guard.
315            return;
316        }
317        tracing::info!(
318            tools = new_tools.len(),
319            "tools/list_changed: agent tool list refreshed"
320        );
321        self.services.mcp.tools = new_tools;
322        self.services.mcp.sync_executor_tools();
323        self.services.mcp.pruning_cache.reset();
324        self.refresh_mcp_tool_ids();
325        self.rebuild_semantic_index().await;
326        self.sync_mcp_registry().await;
327        self.refresh_shadow_sentinel_mcp_tool_ids();
328        let mcp_total = self.services.mcp.tools.len();
329        let mcp_servers = self
330            .services
331            .mcp
332            .tools
333            .iter()
334            .map(|t| &t.server_id)
335            .collect::<std::collections::HashSet<_>>()
336            .len();
337        self.update_metrics(|m| {
338            m.mcp_tool_count = mcp_total;
339            m.mcp_server_count = mcp_servers;
340        });
341    }
342
343    /// Refreshes `ShadowSentinel`'s registered-MCP-tool-id set from the current
344    /// `services.mcp.tools` list, so a server connected after startup (via `/mcp add` or a live
345    /// `tools/list_changed` notification) is reflected without waiting for a process restart.
346    ///
347    /// Only covers `ShadowSentinel`'s own tool-id set, used for risk classification.
348    /// `TrustGateExecutor`'s separate Quarantine-deny set is refreshed independently by
349    /// [`refresh_mcp_tool_ids`](Self::refresh_mcp_tool_ids) (#5747).
350    fn refresh_shadow_sentinel_mcp_tool_ids(&self) {
351        let Some(ref sentinel) = self.services.security.shadow_sentinel else {
352            return;
353        };
354        let ids: std::collections::HashSet<String> = self
355            .services
356            .mcp
357            .tools
358            .iter()
359            .map(zeph_mcp::McpTool::sanitized_id)
360            .collect();
361        *sentinel.mcp_tool_ids_handle().write() = ids;
362    }
363
364    /// Rebuilds `TrustGateExecutor`'s MCP tool-id registry (`self.services.security.mcp_tool_ids`)
365    /// from the current `self.services.mcp.tools`, using the same id derivation
366    /// (`McpTool::sanitized_id`) and replace-not-union semantics as the startup-time
367    /// `register_mcp_tool_ids` (binary crate `agent_setup.rs`). Replace semantics correctly
368    /// drop ids for servers that disconnected since the last refresh. A no-op when no handle
369    /// was attached via `AgentBuilder::with_mcp_tool_ids_handle`.
370    fn refresh_mcp_tool_ids(&self) {
371        let Some(ref handle) = self.services.security.mcp_tool_ids else {
372            return;
373        };
374        let ids: std::collections::HashSet<String> = self
375            .services
376            .mcp
377            .tools
378            .iter()
379            .map(zeph_mcp::McpTool::sanitized_id)
380            .collect();
381        *handle.write() = ids;
382    }
383
384    pub(super) async fn sync_mcp_registry(&mut self) {
385        if self.services.mcp.registry.is_none() {
386            return;
387        }
388        if !self.embedding_provider.supports_embeddings() {
389            return;
390        }
391        // Clone tools before .await to avoid holding &self.services.mcp.tools across an await point.
392        let tools = self.services.mcp.tools.clone();
393        let provider = self.embedding_provider.clone();
394        let embedding_model = self.services.skill.embedding_model.clone();
395        let embed_timeout =
396            std::time::Duration::from_secs(self.runtime.config.timeouts.embedding_seconds);
397        let embed_fn = move |text: &str| -> zeph_mcp::registry::EmbedFuture {
398            let owned = text.to_owned();
399            let p = provider.clone();
400            Box::pin(async move {
401                if let Ok(result) = tokio::time::timeout(embed_timeout, p.embed(&owned)).await {
402                    result
403                } else {
404                    tracing::warn!(
405                        timeout_secs = embed_timeout.as_secs(),
406                        "MCP registry: embedding timed out"
407                    );
408                    Err(zeph_llm::LlmError::Timeout)
409                }
410            })
411        };
412        // Take registry out of self to avoid holding &mut self.services.mcp.registry across .await.
413        // No early returns between take() and put-back — the await is the only yield point here.
414        let Some(mut registry) = self.services.mcp.registry.take() else {
415            return;
416        };
417        if let Err(e) = registry.sync(&tools, &embedding_model, embed_fn).await {
418            tracing::warn!("failed to sync MCP tool registry: {e:#}");
419        }
420        self.services.mcp.registry = Some(registry);
421    }
422
423    /// Build (or rebuild) the in-memory semantic tool index for embedding-based discovery.
424    /// Build the initial semantic tool index after agent construction.
425    ///
426    /// Must be called once after `with_mcp` and `with_mcp_discovery` are applied,
427    /// before the first user turn.  Subsequent rebuilds happen automatically on
428    /// tool list change events (`check_tool_refresh`, `/mcp add`, `/mcp remove`).
429    pub async fn init_semantic_index(&mut self) {
430        self.rebuild_semantic_index().await;
431    }
432
433    /// Drain and process all pending elicitation requests without blocking.
434    ///
435    /// Call this at the start of each turn and between tool calls to prevent
436    /// elicitation events from accumulating while the agent loop is busy.
437    pub(super) async fn process_pending_elicitations(&mut self) {
438        loop {
439            let Some(ref mut rx) = self.services.mcp.elicitation_rx else {
440                return;
441            };
442            match rx.try_recv() {
443                Ok(event) => {
444                    self.handle_elicitation_event(event).await;
445                }
446                Err(tokio::sync::mpsc::error::TryRecvError::Empty) => return,
447                Err(tokio::sync::mpsc::error::TryRecvError::Disconnected) => {
448                    self.services.mcp.elicitation_rx = None;
449                    return;
450                }
451            }
452        }
453    }
454
455    /// Handle a single elicitation event by routing it to the active channel.
456    pub(super) async fn handle_elicitation_event(&mut self, event: zeph_mcp::ElicitationEvent) {
457        use crate::channel::{ElicitationRequest, ElicitationResponse};
458
459        let decline = ElicitResult::new(ElicitationAction::Decline);
460
461        let channel_request = match &event.request {
462            rmcp::model::ElicitRequestParams::FormElicitationParams {
463                message,
464                requested_schema,
465                ..
466            } => {
467                let fields = build_elicitation_fields(requested_schema);
468                ElicitationRequest {
469                    server_name: event.server_id.clone(),
470                    message: sanitize_elicitation_message(message),
471                    fields,
472                }
473            }
474            rmcp::model::ElicitRequestParams::UrlElicitationParams { .. } => {
475                // URL elicitation not supported in phase 1 — decline.
476                tracing::debug!(
477                    server_id = event.server_id,
478                    "URL elicitation not supported, declining"
479                );
480                let _ = event.response_tx.send(decline);
481                return;
482            }
483            // ElicitRequestParams is #[non_exhaustive] — decline unknown future variants.
484            _ => {
485                tracing::debug!(
486                    server_id = event.server_id,
487                    "unknown elicitation request variant, declining"
488                );
489                let _ = event.response_tx.send(decline);
490                return;
491            }
492        };
493
494        if self.services.mcp.elicitation_warn_sensitive_fields {
495            let sensitive: Vec<&str> = channel_request
496                .fields
497                .iter()
498                .filter(|f| is_sensitive_field(&f.name))
499                .map(|f| f.name.as_str())
500                .collect();
501            if !sensitive.is_empty() {
502                let fields_list = sensitive.join(", ");
503                let warning = format!(
504                    "Warning: [{}] is requesting sensitive information (field: {}). \
505                     Only proceed if you trust this server.",
506                    channel_request.server_name, fields_list,
507                );
508                tracing::warn!(
509                    server_id = event.server_id,
510                    fields = %fields_list,
511                    "elicitation requests sensitive fields"
512                );
513                let _ = self.channel.send(&warning).await;
514            }
515        }
516
517        let _ = self
518            .channel
519            .send_status("MCP server requesting input…")
520            .await;
521        let response = match self.channel.elicit(channel_request).await {
522            Ok(r) => r,
523            Err(e) => {
524                tracing::warn!(
525                    server_id = event.server_id,
526                    "elicitation channel error: {e:#}"
527                );
528                let _ = self.channel.send_status("").await;
529                let _ = event.response_tx.send(decline);
530                return;
531            }
532        };
533        let _ = self.channel.send_status("").await;
534
535        let result = match response {
536            ElicitationResponse::Accepted(value) => {
537                ElicitResult::new(ElicitationAction::Accept).with_content(value)
538            }
539            ElicitationResponse::Declined => ElicitResult::new(ElicitationAction::Decline),
540            ElicitationResponse::Cancelled => ElicitResult::new(ElicitationAction::Cancel),
541        };
542
543        if event.response_tx.send(result).is_err() {
544            tracing::warn!(
545                server_id = event.server_id,
546                "elicitation response dropped — handler disconnected"
547            );
548        }
549    }
550
551    fn update_mcp_metrics(&mut self) {
552        let mcp_total = self.services.mcp.tools.len();
553        let mcp_server_count = self.services.mcp.server_outcomes.len();
554        let mcp_connected_count = self
555            .services
556            .mcp
557            .server_outcomes
558            .iter()
559            .filter(|o| o.connected)
560            .count();
561        let mcp_servers: Vec<crate::metrics::McpServerStatus> = self
562            .services
563            .mcp
564            .server_outcomes
565            .iter()
566            .map(|o| crate::metrics::McpServerStatus {
567                id: o.id.clone(),
568                status: if o.connected {
569                    crate::metrics::McpServerConnectionStatus::Connected
570                } else {
571                    crate::metrics::McpServerConnectionStatus::Failed
572                },
573                tool_count: o.tool_count,
574                error: o.error.clone(),
575            })
576            .collect();
577        self.update_metrics(|m| {
578            m.mcp_tool_count = mcp_total;
579            m.mcp_server_count = mcp_server_count;
580            m.mcp_connected_count = mcp_connected_count;
581            m.mcp_servers = mcp_servers;
582        });
583    }
584
585    /// Rebuild the in-memory semantic tool index.
586    ///
587    /// Only runs when `discovery_strategy == Embedding`.  On failure (all embeddings fail),
588    /// sets `semantic_index = None` and logs at WARN — the caller falls back to all tools.
589    ///
590    /// Called at:
591    /// - initial setup via `init_semantic_index()`
592    /// - `tools/list_changed` notification
593    /// - `/mcp add` and `/mcp remove`
594    pub(in crate::agent) async fn rebuild_semantic_index(&mut self) {
595        if self.services.mcp.discovery_strategy != zeph_mcp::ToolDiscoveryStrategy::Embedding {
596            return;
597        }
598
599        if self.services.mcp.tools.is_empty() {
600            self.services.mcp.semantic_index = None;
601            return;
602        }
603
604        // Resolve embedding provider: dedicated discovery provider → primary embedding provider.
605        let provider = self
606            .services
607            .mcp
608            .discovery_provider
609            .clone()
610            .unwrap_or_else(|| self.embedding_provider.clone());
611
612        let inner_embed = provider.embed_fn();
613        let embed_timeout =
614            std::time::Duration::from_secs(self.runtime.config.timeouts.embedding_seconds);
615        let embed_fn = move |text: &str| -> zeph_llm::provider::EmbedFuture {
616            let fut = inner_embed(text);
617            Box::pin(async move {
618                if let Ok(result) = tokio::time::timeout(embed_timeout, fut).await {
619                    result
620                } else {
621                    tracing::warn!(
622                        timeout_secs = embed_timeout.as_secs(),
623                        "semantic index: embedding probe timed out"
624                    );
625                    Err(zeph_llm::LlmError::Timeout)
626                }
627            })
628        };
629
630        // Clone tools before .await to avoid holding &self.services.mcp.tools across an await point.
631        let tools = self.services.mcp.tools.clone();
632        match zeph_mcp::SemanticToolIndex::build(&tools, &embed_fn).await {
633            Ok(idx) => {
634                tracing::info!(
635                    indexed = idx.len(),
636                    total = self.services.mcp.tools.len(),
637                    "semantic tool index built"
638                );
639                self.services.mcp.semantic_index = Some(idx);
640            }
641            Err(e) => {
642                tracing::warn!(
643                    "semantic tool index build failed, falling back to all tools: {e:#}"
644                );
645                self.services.mcp.semantic_index = None;
646            }
647        }
648    }
649}
650
651/// SEC-MCP-01: validate that a stdio command target is on the allowlist.
652///
653/// Returns `Some(error_message)` when the command is blocked, `None` when it is allowed.
654fn validate_mcp_command(target: &str, allowed_commands: &[String]) -> Option<String> {
655    let is_url = target.starts_with("http://") || target.starts_with("https://");
656    if !is_url && !allowed_commands.is_empty() && !allowed_commands.iter().any(|c| c == target) {
657        Some(format!(
658            "Command '{target}' is not allowed. Permitted: {}",
659            allowed_commands.join(", ")
660        ))
661    } else {
662        None
663    }
664}
665
666/// Build a `ServerEntry` for a newly added MCP server from parsed `/mcp add` arguments.
667fn build_server_entry(id: &str, target: &str, extra_args: &[&str]) -> zeph_mcp::ServerEntry {
668    let is_url = target.starts_with("http://") || target.starts_with("https://");
669    let transport = if is_url {
670        zeph_mcp::McpTransport::Http {
671            url: target.to_owned(),
672            headers: std::collections::HashMap::new(),
673        }
674    } else {
675        zeph_mcp::McpTransport::Stdio {
676            command: target.to_owned(),
677            args: extra_args.iter().map(|&s| s.to_owned()).collect(),
678            env: std::collections::HashMap::new(),
679        }
680    };
681    zeph_mcp::ServerEntry {
682        id: id.to_owned(),
683        transport,
684        timeout: std::time::Duration::from_secs(30),
685        trust_level: zeph_config::McpTrustLevel::Untrusted,
686        tool_allowlist: None,
687        expected_tools: Vec::new(),
688        roots: Vec::new(),
689        tool_metadata: std::collections::HashMap::new(),
690        elicitation_enabled: false,
691        elicitation_timeout_secs: 120,
692        env_isolation: false,
693    }
694}
695
696/// Convert an rmcp `ElicitationSchema` into channel-agnostic `ElicitationField` list.
697fn build_elicitation_fields(
698    schema: &rmcp::model::ElicitationSchema,
699) -> Vec<crate::channel::ElicitationField> {
700    use crate::channel::{ElicitationField, ElicitationFieldType};
701    use rmcp::model::PrimitiveSchemaDefinition;
702
703    schema
704        .properties
705        .iter()
706        .map(|(name, prop)| {
707            // Extract field type and description by serializing the PrimitiveSchemaDefinition
708            // to JSON and reading the discriminator field.  This avoids deep-matching the
709            // nested EnumSchema / StringSchema / … variants of rmcp's type-safe schema
710            // hierarchy.
711            let json = serde_json::to_value(prop).unwrap_or_default();
712            let description = json
713                .get("description")
714                .and_then(|v| v.as_str())
715                .map(sanitize_elicitation_message);
716
717            let field_type = match prop {
718                PrimitiveSchemaDefinition::Boolean(_) => ElicitationFieldType::Boolean,
719                PrimitiveSchemaDefinition::Integer(_) => ElicitationFieldType::Integer,
720                PrimitiveSchemaDefinition::Number(_) => ElicitationFieldType::Number,
721                PrimitiveSchemaDefinition::Enum(_) => {
722                    // Extract enum values from the serialized form.  All EnumSchema variants
723                    // serialise their allowed values under "enum" or inside "items.enum".
724                    let vals = json
725                        .get("enum")
726                        .and_then(|v| v.as_array())
727                        .map(|arr| {
728                            arr.iter()
729                                .filter_map(|v| v.as_str())
730                                .map(sanitize_elicitation_message)
731                                .collect::<Vec<_>>()
732                        })
733                        .unwrap_or_default();
734                    ElicitationFieldType::Enum(vals)
735                }
736                PrimitiveSchemaDefinition::String(_) => ElicitationFieldType::String,
737                // Any future `#[non_exhaustive]` variant falls back to
738                // `ElicitationFieldType::String` rather than panicking.
739                _ => {
740                    tracing::debug!(
741                        "unknown PrimitiveSchemaDefinition variant, defaulting to String"
742                    );
743                    ElicitationFieldType::String
744                }
745            };
746            let required = schema.required.as_deref().is_some_and(|r| r.contains(name));
747            ElicitationField {
748                // Keep the raw schema key intact — it is used verbatim as the response map key.
749                // Channels sanitize it at display time only (see build_field_prompt / build_telegram_field_prompt).
750                name: name.clone(),
751                description,
752                field_type,
753                required,
754            }
755        })
756        .collect()
757}
758
759/// Sensitive field name patterns (case-insensitive substring match).
760const SENSITIVE_FIELD_PATTERNS: &[&str] = &[
761    "password",
762    "passwd",
763    "token",
764    "secret",
765    "key",
766    "credential",
767    "apikey",
768    "api_key",
769    "auth",
770    "authorization",
771    "private",
772    "passphrase",
773    "pin",
774];
775
776/// Returns `true` when `field_name` matches any sensitive pattern (case-insensitive).
777fn is_sensitive_field(field_name: &str) -> bool {
778    let lower = field_name.to_lowercase();
779    SENSITIVE_FIELD_PATTERNS
780        .iter()
781        .any(|pattern| lower.contains(pattern))
782}
783
784/// Sanitize an elicitation message: cap length (in chars, not bytes) and strip control chars.
785fn sanitize_elicitation_message(message: &str) -> String {
786    const MAX_CHARS: usize = 500;
787    // Collect up to MAX_CHARS chars, filtering control characters that could manipulate terminals.
788    message
789        .chars()
790        .filter(|c| !c.is_control() || *c == '\n' || *c == '\t')
791        .take(MAX_CHARS)
792        .collect()
793}
794
795#[cfg(test)]
796mod tests {
797    use super::super::agent_tests::{
798        MockChannel, MockToolExecutor, create_test_registry, mock_provider,
799    };
800    use super::*;
801    use std::assert_matches;
802
803    #[tokio::test]
804    async fn handle_mcp_command_unknown_subcommand_shows_usage() {
805        let provider = mock_provider(vec![]);
806        let channel = MockChannel::new(vec![]);
807        let registry = create_test_registry();
808        let executor = MockToolExecutor::no_tools();
809        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
810
811        let result = agent.handle_mcp_command("unknown").await.unwrap();
812        assert!(
813            result.contains("Usage: /mcp"),
814            "expected usage message, got: {result:?}"
815        );
816    }
817
818    #[tokio::test]
819    async fn handle_mcp_list_no_manager_shows_disabled() {
820        let provider = mock_provider(vec![]);
821        let channel = MockChannel::new(vec![]);
822        let registry = create_test_registry();
823        let executor = MockToolExecutor::no_tools();
824        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
825
826        let result = agent.handle_mcp_command("list").await.unwrap();
827        assert!(
828            result.contains("MCP is not enabled"),
829            "expected not-enabled message, got: {result:?}"
830        );
831    }
832
833    #[tokio::test]
834    async fn handle_mcp_tools_no_server_id_shows_usage() {
835        let provider = mock_provider(vec![]);
836        let channel = MockChannel::new(vec![]);
837        let registry = create_test_registry();
838        let executor = MockToolExecutor::no_tools();
839        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
840
841        let result = agent.handle_mcp_command("tools").await.unwrap();
842        assert!(
843            result.contains("Usage: /mcp tools"),
844            "expected tools usage message, got: {result:?}"
845        );
846    }
847
848    #[tokio::test]
849    async fn handle_mcp_remove_no_server_id_shows_usage() {
850        let provider = mock_provider(vec![]);
851        let channel = MockChannel::new(vec![]);
852        let registry = create_test_registry();
853        let executor = MockToolExecutor::no_tools();
854        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
855
856        let result = agent.handle_mcp_command("remove").await.unwrap();
857        assert!(
858            result.contains("Usage: /mcp remove"),
859            "expected remove usage message, got: {result:?}"
860        );
861    }
862
863    #[tokio::test]
864    async fn handle_mcp_remove_no_manager_shows_disabled() {
865        let provider = mock_provider(vec![]);
866        let channel = MockChannel::new(vec![]);
867        let registry = create_test_registry();
868        let executor = MockToolExecutor::no_tools();
869        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
870
871        let result = agent.handle_mcp_command("remove my-server").await.unwrap();
872        assert!(
873            result.contains("MCP is not enabled"),
874            "expected not-enabled message, got: {result:?}"
875        );
876    }
877
878    #[tokio::test]
879    async fn handle_mcp_add_insufficient_args_shows_usage() {
880        let provider = mock_provider(vec![]);
881        let channel = MockChannel::new(vec![]);
882        let registry = create_test_registry();
883        let executor = MockToolExecutor::no_tools();
884        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
885
886        // "add" with only 1 arg (needs at least 2: id + command)
887        let result = agent.handle_mcp_command("add server-id").await.unwrap();
888        assert!(
889            result.contains("Usage: /mcp add"),
890            "expected add usage message, got: {result:?}"
891        );
892    }
893
894    #[tokio::test]
895    async fn handle_mcp_tools_with_unknown_server_shows_no_tools() {
896        let provider = mock_provider(vec![]);
897        let channel = MockChannel::new(vec![]);
898        let registry = create_test_registry();
899        let executor = MockToolExecutor::no_tools();
900        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
901
902        // mcp.tools is empty, so any server will have no tools
903        let result = agent
904            .handle_mcp_command("tools nonexistent-server")
905            .await
906            .unwrap();
907        assert!(
908            result.contains("No tools found"),
909            "expected no-tools message, got: {result:?}"
910        );
911    }
912
913    #[tokio::test]
914    async fn mcp_tool_count_starts_at_zero() {
915        let provider = mock_provider(vec![]);
916        let channel = MockChannel::new(vec![]);
917        let registry = create_test_registry();
918        let executor = MockToolExecutor::no_tools();
919        let agent = Agent::new(provider, channel, registry, None, 5, executor);
920
921        assert_eq!(agent.services.mcp.tool_count(), 0);
922    }
923
924    #[tokio::test]
925    async fn check_tool_refresh_no_rx_is_noop() {
926        let provider = mock_provider(vec![]);
927        let channel = MockChannel::new(vec![]);
928        let registry = create_test_registry();
929        let executor = MockToolExecutor::no_tools();
930        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
931        // No tool_rx set; check_tool_refresh should be a no-op.
932        agent.check_tool_refresh().await;
933        assert_eq!(agent.services.mcp.tool_count(), 0);
934    }
935
936    #[tokio::test]
937    async fn check_tool_refresh_no_change_is_noop() {
938        let provider = mock_provider(vec![]);
939        let channel = MockChannel::new(vec![]);
940        let registry = create_test_registry();
941        let executor = MockToolExecutor::no_tools();
942        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
943
944        let (tx, rx) = tokio::sync::watch::channel(Vec::new());
945        agent.services.mcp.tool_rx = Some(rx);
946        // No changes sent; has_changed() returns false.
947        agent.check_tool_refresh().await;
948        assert_eq!(agent.services.mcp.tool_count(), 0);
949        drop(tx);
950    }
951
952    #[tokio::test]
953    async fn check_tool_refresh_with_empty_initial_value_does_not_replace_tools() {
954        let provider = mock_provider(vec![]);
955        let channel = MockChannel::new(vec![]);
956        let registry = create_test_registry();
957        let executor = MockToolExecutor::no_tools();
958        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
959        agent.services.mcp.tools = vec![zeph_mcp::McpTool {
960            server_id: "srv".into(),
961            name: "existing_tool".into(),
962            description: String::new(),
963            input_schema: serde_json::json!({}),
964            output_schema: None,
965            security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
966        }];
967
968        let (_tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
969        agent.services.mcp.tool_rx = Some(rx);
970        // has_changed() is false for a fresh receiver; tools unchanged.
971        agent.check_tool_refresh().await;
972        assert_eq!(agent.services.mcp.tool_count(), 1);
973    }
974
975    #[tokio::test]
976    async fn check_tool_refresh_applies_update() {
977        let provider = mock_provider(vec![]);
978        let channel = MockChannel::new(vec![]);
979        let registry = create_test_registry();
980        let executor = MockToolExecutor::no_tools();
981        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
982
983        let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
984        agent.services.mcp.tool_rx = Some(rx);
985
986        let new_tools = vec![zeph_mcp::McpTool {
987            server_id: "srv".into(),
988            name: "refreshed_tool".into(),
989            description: String::new(),
990            input_schema: serde_json::json!({}),
991            output_schema: None,
992            security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
993        }];
994        tx.send(new_tools).unwrap();
995
996        agent.check_tool_refresh().await;
997        assert_eq!(agent.services.mcp.tool_count(), 1);
998        assert_eq!(agent.services.mcp.tools[0].name, "refreshed_tool");
999    }
1000
1001    /// #5736 follow-up (S1): a server connected after startup via a `tools/list_changed`
1002    /// notification must be reflected in `ShadowSentinel`'s registered-MCP-tool-id set, not just
1003    /// `services.mcp.tools` — otherwise the escalation fix in `classify_tool` stays blind to any
1004    /// MCP tool the agent didn't already know about at process start.
1005    #[tokio::test]
1006    async fn check_tool_refresh_updates_shadow_sentinel_mcp_tool_ids() {
1007        use crate::agent::shadow_sentinel::{
1008            ProbeVerdict, SafetyProbe, ShadowEventStore, ShadowSentinel,
1009        };
1010
1011        struct NoopProbe;
1012        impl SafetyProbe for NoopProbe {
1013            fn evaluate<'a>(
1014                &'a self,
1015                _: &'a str,
1016                _: &'a serde_json::Value,
1017                _: &'a [crate::agent::shadow_sentinel::SentinelEvent],
1018            ) -> std::pin::Pin<Box<dyn std::future::Future<Output = ProbeVerdict> + Send + 'a>>
1019            {
1020                Box::pin(async { ProbeVerdict::Allow })
1021            }
1022        }
1023
1024        let provider = mock_provider(vec![]);
1025        let channel = MockChannel::new(vec![]);
1026        let registry = create_test_registry();
1027        let executor = MockToolExecutor::no_tools();
1028        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1029
1030        let pool = zeph_db::DbConfig {
1031            url: ":memory:".to_owned(),
1032            ..Default::default()
1033        }
1034        .connect()
1035        .await
1036        .expect("connect + migrate in-memory sqlite pool");
1037        let store = ShadowEventStore::new(pool);
1038        let sentinel = std::sync::Arc::new(ShadowSentinel::new(
1039            store,
1040            Box::new(NoopProbe),
1041            zeph_config::ShadowSentinelConfig::default(),
1042            "test-session",
1043        ));
1044        agent.services.security.shadow_sentinel = Some(std::sync::Arc::clone(&sentinel));
1045
1046        let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1047        agent.services.mcp.tool_rx = Some(rx);
1048
1049        let new_tool = zeph_mcp::McpTool {
1050            server_id: "srv".into(),
1051            name: "refreshed_tool".into(),
1052            description: String::new(),
1053            input_schema: serde_json::json!({}),
1054            output_schema: None,
1055            security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1056        };
1057        let expected_id = new_tool.sanitized_id();
1058        tx.send(vec![new_tool]).unwrap();
1059
1060        agent.check_tool_refresh().await;
1061
1062        assert!(
1063            sentinel.mcp_tool_ids_handle().read().contains(&expected_id),
1064            "ShadowSentinel's mcp_tool_ids must be refreshed after a tools/list_changed event"
1065        );
1066    }
1067
1068    #[tokio::test]
1069    async fn check_tool_refresh_without_mcp_tool_ids_handle_does_not_panic() {
1070        let provider = mock_provider(vec![]);
1071        let channel = MockChannel::new(vec![]);
1072        let registry = create_test_registry();
1073        let executor = MockToolExecutor::no_tools();
1074        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1075        // No handle attached (services.security.mcp_tool_ids is None by default) — refresh
1076        // must be a no-op w.r.t. that field, and the tool list update must still apply.
1077        assert!(agent.services.security.mcp_tool_ids.is_none());
1078
1079        let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1080        agent.services.mcp.tool_rx = Some(rx);
1081        let new_tools = vec![zeph_mcp::McpTool {
1082            server_id: "srv".into(),
1083            name: "refreshed_tool".into(),
1084            description: String::new(),
1085            input_schema: serde_json::json!({}),
1086            output_schema: None,
1087            security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1088        }];
1089        tx.send(new_tools).unwrap();
1090
1091        agent.check_tool_refresh().await;
1092        assert_eq!(agent.services.mcp.tool_count(), 1);
1093    }
1094
1095    #[tokio::test]
1096    async fn check_tool_refresh_updates_attached_mcp_tool_ids_handle() {
1097        let provider = mock_provider(vec![]);
1098        let channel = MockChannel::new(vec![]);
1099        let registry = create_test_registry();
1100        let executor = MockToolExecutor::no_tools();
1101        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1102
1103        let handle =
1104            std::sync::Arc::new(parking_lot::RwLock::new(std::collections::HashSet::new()));
1105        agent.services.security.mcp_tool_ids = Some(std::sync::Arc::clone(&handle));
1106
1107        let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1108        agent.services.mcp.tool_rx = Some(rx);
1109        let new_tools = vec![zeph_mcp::McpTool {
1110            server_id: "srv".into(),
1111            name: "refreshed_tool".into(),
1112            description: String::new(),
1113            input_schema: serde_json::json!({}),
1114            output_schema: None,
1115            security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1116        }];
1117        tx.send(new_tools).unwrap();
1118
1119        agent.check_tool_refresh().await;
1120
1121        assert!(
1122            handle.read().contains("srv_refreshed_tool"),
1123            "expected the sanitized id of the newly-connected tool in the handle, got: {:?}",
1124            *handle.read()
1125        );
1126    }
1127
1128    #[tokio::test]
1129    async fn check_tool_refresh_drops_disconnected_tool_from_mcp_tool_ids_handle() {
1130        let provider = mock_provider(vec![]);
1131        let channel = MockChannel::new(vec![]);
1132        let registry = create_test_registry();
1133        let executor = MockToolExecutor::no_tools();
1134        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1135
1136        // Simulate a stale id left over from an earlier (startup-time) population, for a
1137        // server that has since disconnected.
1138        let handle =
1139            std::sync::Arc::new(parking_lot::RwLock::new(std::collections::HashSet::from([
1140                "stale_server_old_tool".to_owned(),
1141            ])));
1142        agent.services.security.mcp_tool_ids = Some(std::sync::Arc::clone(&handle));
1143
1144        let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1145        agent.services.mcp.tool_rx = Some(rx);
1146        let new_tools = vec![zeph_mcp::McpTool {
1147            server_id: "srv".into(),
1148            name: "refreshed_tool".into(),
1149            description: String::new(),
1150            input_schema: serde_json::json!({}),
1151            output_schema: None,
1152            security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1153        }];
1154        tx.send(new_tools).unwrap();
1155
1156        agent.check_tool_refresh().await;
1157
1158        let ids = handle.read();
1159        assert!(
1160            !ids.contains("stale_server_old_tool"),
1161            "disconnected server's tool id must be dropped (replace, not union), got: {ids:?}"
1162        );
1163        assert!(ids.contains("srv_refreshed_tool"));
1164    }
1165
1166    /// Covers the `pending_semantic_rebuild` trigger (set by `/mcp add`/`/mcp remove`), the
1167    /// second of the two `check_tool_refresh` branches that call `refresh_mcp_tool_ids` — the
1168    /// other tests above only exercise the `tools/list_changed` (`tool_rx`) branch.
1169    #[tokio::test]
1170    async fn check_tool_refresh_updates_mcp_tool_ids_handle_via_pending_semantic_rebuild() {
1171        let provider = mock_provider(vec![]);
1172        let channel = MockChannel::new(vec![]);
1173        let registry = create_test_registry();
1174        let executor = MockToolExecutor::no_tools();
1175        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1176
1177        let handle =
1178            std::sync::Arc::new(parking_lot::RwLock::new(std::collections::HashSet::new()));
1179        agent.services.security.mcp_tool_ids = Some(std::sync::Arc::clone(&handle));
1180
1181        // Mirrors what handle_mcp_add/handle_mcp_remove already do before setting the flag:
1182        // self.services.mcp.tools is updated first, then pending_semantic_rebuild = true.
1183        agent.services.mcp.tools = vec![zeph_mcp::McpTool {
1184            server_id: "srv".into(),
1185            name: "added_tool".into(),
1186            description: String::new(),
1187            input_schema: serde_json::json!({}),
1188            output_schema: None,
1189            security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1190        }];
1191        agent.services.mcp.pending_semantic_rebuild = true;
1192
1193        agent.check_tool_refresh().await;
1194
1195        assert!(
1196            handle.read().contains("srv_added_tool"),
1197            "expected the sanitized id of the /mcp add-connected tool in the handle, got: {:?}",
1198            *handle.read()
1199        );
1200        assert!(!agent.services.mcp.pending_semantic_rebuild);
1201    }
1202
1203    #[test]
1204    fn sanitize_elicitation_message_strips_control_chars() {
1205        let input = "hello\x01world\x1b[31mred\x1b[0m";
1206        let output = sanitize_elicitation_message(input);
1207        assert!(!output.contains('\x01'));
1208        assert!(!output.contains('\x1b'));
1209        assert!(output.contains("hello"));
1210        assert!(output.contains("world"));
1211    }
1212
1213    #[test]
1214    fn sanitize_elicitation_message_preserves_newline_and_tab() {
1215        let input = "line1\nline2\ttabbed";
1216        let output = sanitize_elicitation_message(input);
1217        assert_eq!(output, "line1\nline2\ttabbed");
1218    }
1219
1220    #[test]
1221    fn sanitize_elicitation_message_caps_at_500_chars() {
1222        // Build a 600-char ASCII string — no multi-byte boundary issue.
1223        let input: String = "a".repeat(600);
1224        let output = sanitize_elicitation_message(&input);
1225        assert_eq!(output.chars().count(), 500);
1226    }
1227
1228    #[test]
1229    fn sanitize_elicitation_message_handles_multibyte_boundary() {
1230        // "é" is 2 bytes.  Build a string where a naive &str[..500] would panic.
1231        let input: String = "é".repeat(300); // 300 chars = 600 bytes
1232        let output = sanitize_elicitation_message(&input);
1233        // Should truncate to exactly 500 chars without panic.
1234        assert_eq!(output.chars().count(), 300);
1235    }
1236
1237    #[test]
1238    fn build_elicitation_fields_maps_primitive_types() {
1239        use crate::channel::ElicitationFieldType;
1240        use rmcp::model::{
1241            BooleanSchema, ElicitationSchema, IntegerSchema, NumberSchema,
1242            PrimitiveSchemaDefinition, StringSchema,
1243        };
1244        use std::collections::BTreeMap;
1245
1246        let mut props = BTreeMap::new();
1247        props.insert(
1248            "flag".to_owned(),
1249            PrimitiveSchemaDefinition::Boolean(BooleanSchema::new()),
1250        );
1251        props.insert(
1252            "count".to_owned(),
1253            PrimitiveSchemaDefinition::Integer(IntegerSchema::new()),
1254        );
1255        props.insert(
1256            "ratio".to_owned(),
1257            PrimitiveSchemaDefinition::Number(NumberSchema::new()),
1258        );
1259        props.insert(
1260            "name".to_owned(),
1261            PrimitiveSchemaDefinition::String(StringSchema::new()),
1262        );
1263
1264        let schema = ElicitationSchema::new(props);
1265        let fields = build_elicitation_fields(&schema);
1266
1267        let get = |n: &str| fields.iter().find(|f| f.name == n).unwrap();
1268        assert_matches!(get("flag").field_type, ElicitationFieldType::Boolean);
1269        assert_matches!(get("count").field_type, ElicitationFieldType::Integer);
1270        assert_matches!(get("ratio").field_type, ElicitationFieldType::Number);
1271        assert_matches!(get("name").field_type, ElicitationFieldType::String);
1272    }
1273
1274    #[test]
1275    fn build_elicitation_fields_required_flag() {
1276        use rmcp::model::{ElicitationSchema, PrimitiveSchemaDefinition, StringSchema};
1277        use std::collections::BTreeMap;
1278
1279        let mut props = BTreeMap::new();
1280        props.insert(
1281            "req".to_owned(),
1282            PrimitiveSchemaDefinition::String(StringSchema::new()),
1283        );
1284        props.insert(
1285            "opt".to_owned(),
1286            PrimitiveSchemaDefinition::String(StringSchema::new()),
1287        );
1288
1289        let mut schema = ElicitationSchema::new(props);
1290        schema.required = Some(vec!["req".to_owned()]);
1291
1292        let fields = build_elicitation_fields(&schema);
1293        let req = fields.iter().find(|f| f.name == "req").unwrap();
1294        let opt = fields.iter().find(|f| f.name == "opt").unwrap();
1295        assert!(req.required);
1296        assert!(!opt.required);
1297    }
1298
1299    #[test]
1300    fn is_sensitive_field_detects_common_patterns() {
1301        assert!(is_sensitive_field("password"));
1302        assert!(is_sensitive_field("PASSWORD"));
1303        assert!(is_sensitive_field("user_password"));
1304        assert!(is_sensitive_field("api_token"));
1305        assert!(is_sensitive_field("SECRET_KEY"));
1306        assert!(is_sensitive_field("auth_header"));
1307        assert!(is_sensitive_field("private_key"));
1308    }
1309
1310    #[test]
1311    fn is_sensitive_field_allows_non_sensitive_names() {
1312        assert!(!is_sensitive_field("username"));
1313        assert!(!is_sensitive_field("email"));
1314        assert!(!is_sensitive_field("message"));
1315        assert!(!is_sensitive_field("description"));
1316        assert!(!is_sensitive_field("subject"));
1317    }
1318}