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