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 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 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 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 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 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 pub(super) async fn check_tool_refresh(&mut self) {
315 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 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 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 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 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 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 pub async fn init_semantic_index(&mut self) {
469 self.rebuild_semantic_index().await;
470 }
471
472 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 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 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 _ => {
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 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 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 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
691fn 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
706fn 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 allow_untrusted_without_allowlist: false,
728 expected_tools: Vec::new(),
729 roots: Vec::new(),
730 tool_metadata: std::collections::HashMap::new(),
731 elicitation_enabled: false,
732 elicitation_timeout_secs: 120,
733 env_isolation: false,
734 media_passthrough: false,
735 }
736}
737
738fn build_elicitation_fields(
740 schema: &rmcp::model::ElicitationSchema,
741) -> Vec<crate::channel::ElicitationField> {
742 use crate::channel::{ElicitationField, ElicitationFieldType};
743 use rmcp::model::PrimitiveSchemaDefinition;
744
745 schema
746 .properties
747 .iter()
748 .map(|(name, prop)| {
749 let json = serde_json::to_value(prop).unwrap_or_default();
754 let description = json
755 .get("description")
756 .and_then(|v| v.as_str())
757 .map(sanitize_elicitation_message);
758
759 let field_type = match prop {
760 PrimitiveSchemaDefinition::Boolean(_) => ElicitationFieldType::Boolean,
761 PrimitiveSchemaDefinition::Integer(_) => ElicitationFieldType::Integer,
762 PrimitiveSchemaDefinition::Number(_) => ElicitationFieldType::Number,
763 PrimitiveSchemaDefinition::Enum(_) => {
764 let vals = json
767 .get("enum")
768 .and_then(|v| v.as_array())
769 .map(|arr| {
770 arr.iter()
771 .filter_map(|v| v.as_str())
772 .map(sanitize_elicitation_message)
773 .collect::<Vec<_>>()
774 })
775 .unwrap_or_default();
776 ElicitationFieldType::Enum(vals)
777 }
778 PrimitiveSchemaDefinition::String(_) => ElicitationFieldType::String,
779 _ => {
782 tracing::debug!(
783 "unknown PrimitiveSchemaDefinition variant, defaulting to String"
784 );
785 ElicitationFieldType::String
786 }
787 };
788 let required = schema.required.as_deref().is_some_and(|r| r.contains(name));
789 ElicitationField {
790 name: name.clone(),
793 description,
794 field_type,
795 required,
796 }
797 })
798 .collect()
799}
800
801const SENSITIVE_FIELD_PATTERNS: &[&str] = &[
803 "password",
804 "passwd",
805 "token",
806 "secret",
807 "key",
808 "credential",
809 "apikey",
810 "api_key",
811 "auth",
812 "authorization",
813 "private",
814 "passphrase",
815 "pin",
816];
817
818fn is_sensitive_field(field_name: &str) -> bool {
820 let lower = field_name.to_lowercase();
821 SENSITIVE_FIELD_PATTERNS
822 .iter()
823 .any(|pattern| lower.contains(pattern))
824}
825
826fn sanitize_elicitation_message(message: &str) -> String {
828 const MAX_CHARS: usize = 500;
829 message
831 .chars()
832 .filter(|c| !c.is_control() || *c == '\n' || *c == '\t')
833 .take(MAX_CHARS)
834 .collect()
835}
836
837impl<C: Channel + Send + 'static> zeph_commands::McpAccess for Agent<C> {
838 fn handle_mcp<'a>(
841 &'a mut self,
842 args: &'a str,
843 ) -> std::pin::Pin<
844 Box<
845 dyn std::future::Future<Output = Result<String, zeph_commands::CommandError>>
846 + Send
847 + 'a,
848 >,
849 > {
850 let args_owned = args.to_owned();
853 let parts: Vec<String> = args_owned.split_whitespace().map(str::to_owned).collect();
854 let sub = parts.first().cloned().unwrap_or_default();
855
856 match sub.as_str() {
857 "list" => {
858 let manager = self.services.mcp.manager.clone();
860 let tools_snapshot: Vec<(String, String)> = self
861 .services
862 .mcp
863 .tools
864 .iter()
865 .map(|t| (t.server_id.clone(), t.name.clone()))
866 .collect();
867 Box::pin(async move {
868 use std::fmt::Write;
869 let Some(manager) = manager else {
870 return Ok("MCP is not enabled.".to_owned());
871 };
872 let server_ids = manager.list_servers().await;
873 if server_ids.is_empty() {
874 return Ok("No MCP servers connected.".to_owned());
875 }
876 let mut output = String::from("Connected MCP servers:\n");
877 let mut total = 0usize;
878 for id in &server_ids {
879 let count = tools_snapshot.iter().filter(|(sid, _)| sid == id).count();
880 total += count;
881 let _ = writeln!(output, "- {id} ({count} tools)");
882 }
883 let _ = write!(output, "Total: {total} tool(s)");
884 Ok(output)
885 })
886 }
887 "tools" => {
888 let server_id = parts.get(1).cloned();
890 let owned_tools: Vec<(String, String)> = if let Some(ref sid) = server_id {
891 self.services
892 .mcp
893 .tools
894 .iter()
895 .filter(|t| &t.server_id == sid)
896 .map(|t| (t.name.clone(), t.description.clone()))
897 .collect()
898 } else {
899 Vec::new()
900 };
901 Box::pin(async move {
902 use std::fmt::Write;
903 let Some(server_id) = server_id else {
904 return Ok("Usage: /mcp tools <server_id>".to_owned());
905 };
906 if owned_tools.is_empty() {
907 return Ok(format!("No tools found for server '{server_id}'."));
908 }
909 let mut output =
910 format!("Tools for '{server_id}' ({} total):\n", owned_tools.len());
911 for (name, desc) in &owned_tools {
912 if desc.is_empty() {
913 let _ = writeln!(output, "- {name}");
914 } else {
915 let _ = writeln!(output, "- {name} — {desc}");
916 }
917 }
918 Ok(output)
919 })
920 }
921 _ => Box::pin(async move {
928 self.handle_mcp_command(&args_owned)
929 .await
930 .map_err(|e| zeph_commands::CommandError::new(e.to_string()))
931 }),
932 }
933 }
934}
935
936#[cfg(test)]
937mod tests {
938 use super::super::agent_tests::{
939 MockChannel, MockToolExecutor, create_test_registry, mock_provider,
940 };
941 use super::*;
942 use std::assert_matches;
943
944 #[tokio::test]
945 async fn handle_mcp_command_unknown_subcommand_shows_usage() {
946 let provider = mock_provider(vec![]);
947 let channel = MockChannel::new(vec![]);
948 let registry = create_test_registry();
949 let executor = MockToolExecutor::no_tools();
950 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
951
952 let result = agent.handle_mcp_command("unknown").await.unwrap();
953 assert!(
954 result.contains("Usage: /mcp"),
955 "expected usage message, got: {result:?}"
956 );
957 }
958
959 #[tokio::test]
960 async fn handle_mcp_list_no_manager_shows_disabled() {
961 let provider = mock_provider(vec![]);
962 let channel = MockChannel::new(vec![]);
963 let registry = create_test_registry();
964 let executor = MockToolExecutor::no_tools();
965 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
966
967 let result = agent.handle_mcp_command("list").await.unwrap();
968 assert!(
969 result.contains("MCP is not enabled"),
970 "expected not-enabled message, got: {result:?}"
971 );
972 }
973
974 #[tokio::test]
975 async fn handle_mcp_tools_no_server_id_shows_usage() {
976 let provider = mock_provider(vec![]);
977 let channel = MockChannel::new(vec![]);
978 let registry = create_test_registry();
979 let executor = MockToolExecutor::no_tools();
980 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
981
982 let result = agent.handle_mcp_command("tools").await.unwrap();
983 assert!(
984 result.contains("Usage: /mcp tools"),
985 "expected tools usage message, got: {result:?}"
986 );
987 }
988
989 #[tokio::test]
990 async fn handle_mcp_remove_no_server_id_shows_usage() {
991 let provider = mock_provider(vec![]);
992 let channel = MockChannel::new(vec![]);
993 let registry = create_test_registry();
994 let executor = MockToolExecutor::no_tools();
995 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
996
997 let result = agent.handle_mcp_command("remove").await.unwrap();
998 assert!(
999 result.contains("Usage: /mcp remove"),
1000 "expected remove usage message, got: {result:?}"
1001 );
1002 }
1003
1004 #[tokio::test]
1005 async fn handle_mcp_remove_no_manager_shows_disabled() {
1006 let provider = mock_provider(vec![]);
1007 let channel = MockChannel::new(vec![]);
1008 let registry = create_test_registry();
1009 let executor = MockToolExecutor::no_tools();
1010 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1011
1012 let result = agent.handle_mcp_command("remove my-server").await.unwrap();
1013 assert!(
1014 result.contains("MCP is not enabled"),
1015 "expected not-enabled message, got: {result:?}"
1016 );
1017 }
1018
1019 #[tokio::test]
1020 async fn handle_mcp_add_insufficient_args_shows_usage() {
1021 let provider = mock_provider(vec![]);
1022 let channel = MockChannel::new(vec![]);
1023 let registry = create_test_registry();
1024 let executor = MockToolExecutor::no_tools();
1025 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1026
1027 let result = agent.handle_mcp_command("add server-id").await.unwrap();
1029 assert!(
1030 result.contains("Usage: /mcp add"),
1031 "expected add usage message, got: {result:?}"
1032 );
1033 }
1034
1035 #[tokio::test]
1036 async fn handle_mcp_tools_with_unknown_server_shows_no_tools() {
1037 let provider = mock_provider(vec![]);
1038 let channel = MockChannel::new(vec![]);
1039 let registry = create_test_registry();
1040 let executor = MockToolExecutor::no_tools();
1041 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1042
1043 let result = agent
1045 .handle_mcp_command("tools nonexistent-server")
1046 .await
1047 .unwrap();
1048 assert!(
1049 result.contains("No tools found"),
1050 "expected no-tools message, got: {result:?}"
1051 );
1052 }
1053
1054 #[tokio::test]
1055 async fn mcp_tool_count_starts_at_zero() {
1056 let provider = mock_provider(vec![]);
1057 let channel = MockChannel::new(vec![]);
1058 let registry = create_test_registry();
1059 let executor = MockToolExecutor::no_tools();
1060 let agent = Agent::new(provider, channel, registry, None, 5, executor);
1061
1062 assert_eq!(agent.services.mcp.tool_count(), 0);
1063 }
1064
1065 fn test_mcp_tool(
1066 server_id: &str,
1067 name: &str,
1068 input_schema: serde_json::Value,
1069 ) -> zeph_mcp::McpTool {
1070 zeph_mcp::McpTool {
1071 server_id: server_id.to_owned(),
1072 name: name.to_owned(),
1073 description: format!("{name} description"),
1074 input_schema,
1075 output_schema: None,
1076 security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1077 }
1078 }
1079
1080 #[tokio::test]
1084 async fn rehydrate_mcp_tools_replaces_stub_with_live_schema() {
1085 let provider = mock_provider(vec![]);
1086 let channel = MockChannel::new(vec![]);
1087 let registry = create_test_registry();
1088 let executor = MockToolExecutor::no_tools();
1089 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1090
1091 let real_schema =
1092 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}});
1093 agent.services.mcp.tools = vec![test_mcp_tool("fs", "read_file", real_schema.clone())];
1094
1095 let stub = test_mcp_tool("fs", "read_file", serde_json::json!({}));
1098
1099 let rehydrated = agent.rehydrate_mcp_tools(vec![stub]);
1100
1101 assert_eq!(rehydrated.len(), 1);
1102 assert_eq!(rehydrated[0].input_schema, real_schema);
1103 }
1104
1105 #[tokio::test]
1109 async fn rehydrate_mcp_tools_drops_hit_with_no_live_match() {
1110 let provider = mock_provider(vec![]);
1111 let channel = MockChannel::new(vec![]);
1112 let registry = create_test_registry();
1113 let executor = MockToolExecutor::no_tools();
1114 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1115
1116 agent.services.mcp.tools = vec![test_mcp_tool("fs", "other_tool", serde_json::json!({}))];
1117
1118 let stub = test_mcp_tool("fs", "read_file", serde_json::json!({}));
1119
1120 let rehydrated = agent.rehydrate_mcp_tools(vec![stub]);
1121
1122 assert!(
1123 rehydrated.is_empty(),
1124 "stale hit with no live match must be dropped, not passed through with an empty schema"
1125 );
1126 }
1127
1128 #[tokio::test]
1129 async fn rehydrate_mcp_tools_mixed_batch_keeps_only_matches() {
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 schema_a =
1137 serde_json::json!({"type": "object", "properties": {"a": {"type": "string"}}});
1138 let schema_c =
1139 serde_json::json!({"type": "object", "properties": {"c": {"type": "number"}}});
1140 agent.services.mcp.tools = vec![
1141 test_mcp_tool("srv1", "tool_a", schema_a.clone()),
1142 test_mcp_tool("srv2", "tool_c", schema_c.clone()),
1143 ];
1144
1145 let hits = vec![
1146 test_mcp_tool("srv1", "tool_a", serde_json::json!({})),
1147 test_mcp_tool("srv1", "tool_b", serde_json::json!({})), test_mcp_tool("srv2", "tool_c", serde_json::json!({})),
1149 ];
1150
1151 let rehydrated = agent.rehydrate_mcp_tools(hits);
1152
1153 assert_eq!(rehydrated.len(), 2);
1154 assert_eq!(rehydrated[0].name, "tool_a");
1155 assert_eq!(rehydrated[0].input_schema, schema_a);
1156 assert_eq!(rehydrated[1].name, "tool_c");
1157 assert_eq!(rehydrated[1].input_schema, schema_c);
1158 }
1159
1160 #[tokio::test]
1164 async fn rehydrated_tool_schema_reaches_llm_prompt() {
1165 let provider = mock_provider(vec![]);
1166 let channel = MockChannel::new(vec![]);
1167 let registry = create_test_registry();
1168 let executor = MockToolExecutor::no_tools();
1169 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1170
1171 let real_schema = serde_json::json!({
1172 "type": "object",
1173 "properties": {"query": {"type": "string"}},
1174 "required": ["query"]
1175 });
1176 agent.services.mcp.tools = vec![test_mcp_tool("search", "web_search", real_schema)];
1177
1178 let stub = test_mcp_tool("search", "web_search", serde_json::json!({}));
1179 let rehydrated = agent.rehydrate_mcp_tools(vec![stub]);
1180
1181 let prompt = zeph_mcp::format_mcp_tools_prompt(&rehydrated);
1182
1183 assert!(
1184 !prompt.contains("<parameters>{}</parameters>"),
1185 "expected real schema in prompt, got empty parameters block: {prompt}"
1186 );
1187 assert!(
1188 prompt.contains("\"query\""),
1189 "expected real schema fields in prompt: {prompt}"
1190 );
1191 }
1192
1193 #[tokio::test]
1194 async fn check_tool_refresh_no_rx_is_noop() {
1195 let provider = mock_provider(vec![]);
1196 let channel = MockChannel::new(vec![]);
1197 let registry = create_test_registry();
1198 let executor = MockToolExecutor::no_tools();
1199 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1200 agent.check_tool_refresh().await;
1202 assert_eq!(agent.services.mcp.tool_count(), 0);
1203 }
1204
1205 #[tokio::test]
1206 async fn check_tool_refresh_no_change_is_noop() {
1207 let provider = mock_provider(vec![]);
1208 let channel = MockChannel::new(vec![]);
1209 let registry = create_test_registry();
1210 let executor = MockToolExecutor::no_tools();
1211 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1212
1213 let (tx, rx) = tokio::sync::watch::channel(Vec::new());
1214 agent.services.mcp.tool_rx = Some(rx);
1215 agent.check_tool_refresh().await;
1217 assert_eq!(agent.services.mcp.tool_count(), 0);
1218 drop(tx);
1219 }
1220
1221 #[tokio::test]
1222 async fn check_tool_refresh_with_empty_initial_value_does_not_replace_tools() {
1223 let provider = mock_provider(vec![]);
1224 let channel = MockChannel::new(vec![]);
1225 let registry = create_test_registry();
1226 let executor = MockToolExecutor::no_tools();
1227 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1228 agent.services.mcp.tools = vec![zeph_mcp::McpTool {
1229 server_id: "srv".into(),
1230 name: "existing_tool".into(),
1231 description: String::new(),
1232 input_schema: serde_json::json!({}),
1233 output_schema: None,
1234 security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1235 }];
1236
1237 let (_tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1238 agent.services.mcp.tool_rx = Some(rx);
1239 agent.check_tool_refresh().await;
1241 assert_eq!(agent.services.mcp.tool_count(), 1);
1242 }
1243
1244 #[tokio::test]
1245 async fn check_tool_refresh_applies_update() {
1246 let provider = mock_provider(vec![]);
1247 let channel = MockChannel::new(vec![]);
1248 let registry = create_test_registry();
1249 let executor = MockToolExecutor::no_tools();
1250 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1251
1252 let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1253 agent.services.mcp.tool_rx = Some(rx);
1254
1255 let new_tools = vec![zeph_mcp::McpTool {
1256 server_id: "srv".into(),
1257 name: "refreshed_tool".into(),
1258 description: String::new(),
1259 input_schema: serde_json::json!({}),
1260 output_schema: None,
1261 security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1262 }];
1263 tx.send(new_tools).unwrap();
1264
1265 agent.check_tool_refresh().await;
1266 assert_eq!(agent.services.mcp.tool_count(), 1);
1267 assert_eq!(agent.services.mcp.tools[0].name, "refreshed_tool");
1268 }
1269
1270 #[tokio::test]
1275 async fn check_tool_refresh_updates_shadow_sentinel_mcp_tool_ids() {
1276 use crate::agent::shadow_sentinel::{
1277 ProbeVerdict, SafetyProbe, ShadowEventStore, ShadowSentinel,
1278 };
1279
1280 struct NoopProbe;
1281 impl SafetyProbe for NoopProbe {
1282 fn evaluate<'a>(
1283 &'a self,
1284 _: &'a str,
1285 _: &'a serde_json::Value,
1286 _: &'a [crate::agent::shadow_sentinel::SentinelEvent],
1287 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = ProbeVerdict> + Send + 'a>>
1288 {
1289 Box::pin(async { ProbeVerdict::Allow })
1290 }
1291 }
1292
1293 let provider = mock_provider(vec![]);
1294 let channel = MockChannel::new(vec![]);
1295 let registry = create_test_registry();
1296 let executor = MockToolExecutor::no_tools();
1297 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1298
1299 let pool = zeph_db::DbConfig {
1300 url: ":memory:".to_owned(),
1301 ..Default::default()
1302 }
1303 .connect()
1304 .await
1305 .expect("connect + migrate in-memory sqlite pool");
1306 let store = ShadowEventStore::new(pool);
1307 let sentinel = std::sync::Arc::new(ShadowSentinel::new(
1308 store,
1309 Box::new(NoopProbe),
1310 zeph_config::ShadowSentinelConfig::default(),
1311 "test-session",
1312 ));
1313 agent.services.security.shadow_sentinel = Some(std::sync::Arc::clone(&sentinel));
1314
1315 let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1316 agent.services.mcp.tool_rx = Some(rx);
1317
1318 let new_tool = zeph_mcp::McpTool {
1319 server_id: "srv".into(),
1320 name: "refreshed_tool".into(),
1321 description: String::new(),
1322 input_schema: serde_json::json!({}),
1323 output_schema: None,
1324 security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1325 };
1326 let expected_id = new_tool.sanitized_id();
1327 tx.send(vec![new_tool]).unwrap();
1328
1329 agent.check_tool_refresh().await;
1330
1331 assert!(
1332 sentinel.mcp_tool_ids_handle().read().contains(&expected_id),
1333 "ShadowSentinel's mcp_tool_ids must be refreshed after a tools/list_changed event"
1334 );
1335 }
1336
1337 #[tokio::test]
1338 async fn check_tool_refresh_without_mcp_tool_ids_handle_does_not_panic() {
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 assert!(agent.services.security.mcp_tool_ids.is_none());
1347
1348 let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1349 agent.services.mcp.tool_rx = Some(rx);
1350 let new_tools = vec![zeph_mcp::McpTool {
1351 server_id: "srv".into(),
1352 name: "refreshed_tool".into(),
1353 description: String::new(),
1354 input_schema: serde_json::json!({}),
1355 output_schema: None,
1356 security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1357 }];
1358 tx.send(new_tools).unwrap();
1359
1360 agent.check_tool_refresh().await;
1361 assert_eq!(agent.services.mcp.tool_count(), 1);
1362 }
1363
1364 #[tokio::test]
1365 async fn check_tool_refresh_updates_attached_mcp_tool_ids_handle() {
1366 let provider = mock_provider(vec![]);
1367 let channel = MockChannel::new(vec![]);
1368 let registry = create_test_registry();
1369 let executor = MockToolExecutor::no_tools();
1370 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1371
1372 let handle =
1373 std::sync::Arc::new(parking_lot::RwLock::new(std::collections::HashSet::new()));
1374 agent.services.security.mcp_tool_ids = Some(std::sync::Arc::clone(&handle));
1375
1376 let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1377 agent.services.mcp.tool_rx = Some(rx);
1378 let new_tools = vec![zeph_mcp::McpTool {
1379 server_id: "srv".into(),
1380 name: "refreshed_tool".into(),
1381 description: String::new(),
1382 input_schema: serde_json::json!({}),
1383 output_schema: None,
1384 security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1385 }];
1386 tx.send(new_tools).unwrap();
1387
1388 agent.check_tool_refresh().await;
1389
1390 assert!(
1391 handle.read().contains("srv_refreshed_tool"),
1392 "expected the sanitized id of the newly-connected tool in the handle, got: {:?}",
1393 *handle.read()
1394 );
1395 }
1396
1397 #[tokio::test]
1398 async fn check_tool_refresh_drops_disconnected_tool_from_mcp_tool_ids_handle() {
1399 let provider = mock_provider(vec![]);
1400 let channel = MockChannel::new(vec![]);
1401 let registry = create_test_registry();
1402 let executor = MockToolExecutor::no_tools();
1403 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1404
1405 let handle =
1408 std::sync::Arc::new(parking_lot::RwLock::new(std::collections::HashSet::from([
1409 "stale_server_old_tool".to_owned(),
1410 ])));
1411 agent.services.security.mcp_tool_ids = Some(std::sync::Arc::clone(&handle));
1412
1413 let (tx, rx) = tokio::sync::watch::channel(Vec::<zeph_mcp::McpTool>::new());
1414 agent.services.mcp.tool_rx = Some(rx);
1415 let new_tools = vec![zeph_mcp::McpTool {
1416 server_id: "srv".into(),
1417 name: "refreshed_tool".into(),
1418 description: String::new(),
1419 input_schema: serde_json::json!({}),
1420 output_schema: None,
1421 security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1422 }];
1423 tx.send(new_tools).unwrap();
1424
1425 agent.check_tool_refresh().await;
1426
1427 let ids = handle.read();
1428 assert!(
1429 !ids.contains("stale_server_old_tool"),
1430 "disconnected server's tool id must be dropped (replace, not union), got: {ids:?}"
1431 );
1432 assert!(ids.contains("srv_refreshed_tool"));
1433 }
1434
1435 #[tokio::test]
1439 async fn check_tool_refresh_updates_mcp_tool_ids_handle_via_pending_semantic_rebuild() {
1440 let provider = mock_provider(vec![]);
1441 let channel = MockChannel::new(vec![]);
1442 let registry = create_test_registry();
1443 let executor = MockToolExecutor::no_tools();
1444 let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
1445
1446 let handle =
1447 std::sync::Arc::new(parking_lot::RwLock::new(std::collections::HashSet::new()));
1448 agent.services.security.mcp_tool_ids = Some(std::sync::Arc::clone(&handle));
1449
1450 agent.services.mcp.tools = vec![zeph_mcp::McpTool {
1453 server_id: "srv".into(),
1454 name: "added_tool".into(),
1455 description: String::new(),
1456 input_schema: serde_json::json!({}),
1457 output_schema: None,
1458 security_meta: zeph_config::mcp_security::ToolSecurityMeta::default(),
1459 }];
1460 agent.services.mcp.pending_semantic_rebuild = true;
1461
1462 agent.check_tool_refresh().await;
1463
1464 assert!(
1465 handle.read().contains("srv_added_tool"),
1466 "expected the sanitized id of the /mcp add-connected tool in the handle, got: {:?}",
1467 *handle.read()
1468 );
1469 assert!(!agent.services.mcp.pending_semantic_rebuild);
1470 }
1471
1472 #[test]
1473 fn sanitize_elicitation_message_strips_control_chars() {
1474 let input = "hello\x01world\x1b[31mred\x1b[0m";
1475 let output = sanitize_elicitation_message(input);
1476 assert!(!output.contains('\x01'));
1477 assert!(!output.contains('\x1b'));
1478 assert!(output.contains("hello"));
1479 assert!(output.contains("world"));
1480 }
1481
1482 #[test]
1483 fn sanitize_elicitation_message_preserves_newline_and_tab() {
1484 let input = "line1\nline2\ttabbed";
1485 let output = sanitize_elicitation_message(input);
1486 assert_eq!(output, "line1\nline2\ttabbed");
1487 }
1488
1489 #[test]
1490 fn sanitize_elicitation_message_caps_at_500_chars() {
1491 let input: String = "a".repeat(600);
1493 let output = sanitize_elicitation_message(&input);
1494 assert_eq!(output.chars().count(), 500);
1495 }
1496
1497 #[test]
1498 fn sanitize_elicitation_message_handles_multibyte_boundary() {
1499 let input: String = "é".repeat(300); let output = sanitize_elicitation_message(&input);
1502 assert_eq!(output.chars().count(), 300);
1504 }
1505
1506 #[test]
1507 fn build_elicitation_fields_maps_primitive_types() {
1508 use crate::channel::ElicitationFieldType;
1509 use rmcp::model::{
1510 BooleanSchema, ElicitationSchema, IntegerSchema, NumberSchema,
1511 PrimitiveSchemaDefinition, StringSchema,
1512 };
1513 use std::collections::BTreeMap;
1514
1515 let mut props = BTreeMap::new();
1516 props.insert(
1517 "flag".to_owned(),
1518 PrimitiveSchemaDefinition::Boolean(BooleanSchema::new()),
1519 );
1520 props.insert(
1521 "count".to_owned(),
1522 PrimitiveSchemaDefinition::Integer(IntegerSchema::new()),
1523 );
1524 props.insert(
1525 "ratio".to_owned(),
1526 PrimitiveSchemaDefinition::Number(NumberSchema::new()),
1527 );
1528 props.insert(
1529 "name".to_owned(),
1530 PrimitiveSchemaDefinition::String(StringSchema::new()),
1531 );
1532
1533 let schema = ElicitationSchema::new(props);
1534 let fields = build_elicitation_fields(&schema);
1535
1536 let get = |n: &str| fields.iter().find(|f| f.name == n).unwrap();
1537 assert_matches!(get("flag").field_type, ElicitationFieldType::Boolean);
1538 assert_matches!(get("count").field_type, ElicitationFieldType::Integer);
1539 assert_matches!(get("ratio").field_type, ElicitationFieldType::Number);
1540 assert_matches!(get("name").field_type, ElicitationFieldType::String);
1541 }
1542
1543 #[test]
1544 fn build_elicitation_fields_required_flag() {
1545 use rmcp::model::{ElicitationSchema, PrimitiveSchemaDefinition, StringSchema};
1546 use std::collections::BTreeMap;
1547
1548 let mut props = BTreeMap::new();
1549 props.insert(
1550 "req".to_owned(),
1551 PrimitiveSchemaDefinition::String(StringSchema::new()),
1552 );
1553 props.insert(
1554 "opt".to_owned(),
1555 PrimitiveSchemaDefinition::String(StringSchema::new()),
1556 );
1557
1558 let mut schema = ElicitationSchema::new(props);
1559 schema.required = Some(vec!["req".to_owned()]);
1560
1561 let fields = build_elicitation_fields(&schema);
1562 let req = fields.iter().find(|f| f.name == "req").unwrap();
1563 let opt = fields.iter().find(|f| f.name == "opt").unwrap();
1564 assert!(req.required);
1565 assert!(!opt.required);
1566 }
1567
1568 #[test]
1569 fn is_sensitive_field_detects_common_patterns() {
1570 assert!(is_sensitive_field("password"));
1571 assert!(is_sensitive_field("PASSWORD"));
1572 assert!(is_sensitive_field("user_password"));
1573 assert!(is_sensitive_field("api_token"));
1574 assert!(is_sensitive_field("SECRET_KEY"));
1575 assert!(is_sensitive_field("auth_header"));
1576 assert!(is_sensitive_field("private_key"));
1577 }
1578
1579 #[test]
1580 fn is_sensitive_field_allows_non_sensitive_names() {
1581 assert!(!is_sensitive_field("username"));
1582 assert!(!is_sensitive_field("email"));
1583 assert!(!is_sensitive_field("message"));
1584 assert!(!is_sensitive_field("description"));
1585 assert!(!is_sensitive_field("subject"));
1586 }
1587}