1use rmcp::model::{ElicitResult, ElicitationAction};
5
6use super::{Agent, Channel, LlmProvider};
7
8impl<C: Channel> Agent<C> {
9 #[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 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 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 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 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 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 pub(super) async fn check_tool_refresh(&mut self) {
276 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 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 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 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 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 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 pub async fn init_semantic_index(&mut self) {
430 self.rebuild_semantic_index().await;
431 }
432
433 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 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 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 _ => {
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 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 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 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
651fn 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
666fn 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
696fn 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 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 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 _ => {
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 name: name.clone(),
751 description,
752 field_type,
753 required,
754 }
755 })
756 .collect()
757}
758
759const 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
776fn 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
784fn sanitize_elicitation_message(message: &str) -> String {
786 const MAX_CHARS: usize = 500;
787 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 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 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 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 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 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 #[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 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 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 #[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 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 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 let input: String = "é".repeat(300); let output = sanitize_elicitation_message(&input);
1233 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}