1use crate::types::{RiskLevel, ValidationError};
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6use std::collections::HashSet;
7
8pub fn resolve_server_id_from_env() -> Option<String> {
20 let candidate = std::env::var("PMCP_SERVER_ID")
21 .ok()
22 .or_else(|| std::env::var("AWS_LAMBDA_FUNCTION_NAME").ok())?;
23 if candidate.is_empty() {
24 None
25 } else {
26 Some(candidate)
27 }
28}
29
30#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct OperationEntry {
34 pub id: String,
37
38 pub category: String,
41
42 #[serde(default)]
44 pub description: String,
45
46 #[serde(default)]
54 pub path: Option<String>,
55}
56
57#[derive(Debug, Clone, Default)]
60pub struct OperationRegistry {
61 entries: Vec<RegisteredOperation>,
62}
63
64#[derive(Debug, Clone)]
65struct RegisteredOperation {
66 entry: OperationEntry,
67 method: Option<String>,
68 path: String,
69}
70
71impl RegisteredOperation {
72 fn match_rank(&self, method: Option<&str>, path: &str) -> Option<(bool, usize, bool)> {
75 if let (Some(own), Some(call)) = (&self.method, method) {
76 if !own.eq_ignore_ascii_case(call) {
77 return None;
78 }
79 }
80 let has_method = self.method.is_some();
81 if self.path == path {
82 return Some((true, usize::MAX, has_method));
83 }
84 let own: Vec<&str> = self.path.split('/').collect();
85 let call: Vec<&str> = path.split('/').collect();
86 if own.len() != call.len() {
87 return None;
88 }
89 let mut literals = 0;
90 for (o, c) in own.iter().zip(&call) {
91 if is_template_param(o) {
92 if c.is_empty() {
93 return None;
94 }
95 } else if o == c && !c.contains('{') {
96 literals += 1;
97 } else {
98 return None;
99 }
100 }
101 Some((false, literals, has_method))
102 }
103}
104
105fn is_template_param(segment: &str) -> bool {
106 segment.len() > 2 && segment.starts_with('{') && segment.ends_with('}')
107}
108
109fn split_entry_path(raw: &str) -> (Option<String>, String) {
111 let trimmed = raw.trim();
112 if let Some(idx) = trimmed.find([' ', ':']) {
113 let (head, rest) = trimmed.split_at(idx);
114 let rest = rest[1..].trim();
115 if rest.starts_with('/')
116 && matches!(
117 head.to_ascii_uppercase().as_str(),
118 "GET" | "HEAD" | "OPTIONS" | "POST" | "PUT" | "PATCH" | "DELETE"
119 )
120 {
121 return (Some(head.to_ascii_uppercase()), rest.to_string());
122 }
123 }
124 (None, trimmed.to_string())
125}
126
127fn category_rank(category: &str) -> u8 {
131 match category.trim().to_ascii_lowercase().as_str() {
132 "" => 0,
133 "read" => 1,
134 "write" => 2,
135 "delete" => 3,
136 "admin" => 4,
137 _ => 5,
138 }
139}
140
141impl OperationRegistry {
142 pub fn from_entries(entries: &[OperationEntry]) -> Self {
143 let entries = entries
144 .iter()
145 .filter_map(|entry| {
146 let (method, path) = split_entry_path(entry.path.as_deref()?);
147 Some(RegisteredOperation {
148 entry: entry.clone(),
149 method,
150 path,
151 })
152 })
153 .collect();
154 Self { entries }
155 }
156
157 pub fn lookup_entry(&self, method: Option<&str>, path: &str) -> Option<&OperationEntry> {
166 self.entries
167 .iter()
168 .filter_map(|op| op.match_rank(method, path).map(|rank| (rank, op)))
169 .max_by(|(a_rank, a), (b_rank, b)| {
170 a_rank.cmp(b_rank).then_with(|| {
171 category_rank(&a.entry.category).cmp(&category_rank(&b.entry.category))
172 })
173 })
174 .map(|(_, op)| &op.entry)
175 }
176
177 pub fn lookup(&self, path: &str) -> Option<&str> {
178 self.lookup_entry(None, path).map(|e| e.id.as_str())
179 }
180
181 pub fn lookup_category(&self, path: &str) -> Option<&str> {
184 self.lookup_entry(None, path)
185 .map(|e| e.category.as_str())
186 .filter(|c| !c.is_empty())
187 }
188
189 pub fn is_empty(&self) -> bool {
190 self.entries.is_empty()
191 }
192}
193
194#[derive(Debug, Clone, Serialize, Deserialize)]
196pub struct CodeModeConfig {
197 #[serde(default)]
199 pub enabled: bool,
200
201 #[serde(default)]
206 pub allow_mutations: bool,
207
208 #[serde(default)]
210 pub allowed_mutations: HashSet<String>,
211
212 #[serde(default)]
214 pub blocked_mutations: HashSet<String>,
215
216 #[serde(default)]
218 pub allow_introspection: bool,
219
220 #[serde(default)]
222 pub blocked_fields: HashSet<String>,
223
224 #[serde(default)]
226 pub allowed_queries: HashSet<String>,
227
228 #[serde(default)]
230 pub blocked_queries: HashSet<String>,
231
232 #[serde(default = "default_true")]
237 pub openapi_reads_enabled: bool,
238
239 #[serde(default)]
241 pub openapi_allow_writes: bool,
242
243 #[serde(default)]
245 pub openapi_allowed_writes: HashSet<String>,
246
247 #[serde(default)]
249 pub openapi_blocked_writes: HashSet<String>,
250
251 #[serde(default)]
253 pub openapi_allow_deletes: bool,
254
255 #[serde(default)]
257 pub openapi_allowed_deletes: HashSet<String>,
258
259 #[serde(default)]
261 pub openapi_blocked_paths: HashSet<String>,
262
263 #[serde(default)]
265 pub openapi_internal_blocked_fields: HashSet<String>,
266
267 #[serde(default)]
269 pub openapi_output_blocked_fields: HashSet<String>,
270
271 #[serde(default)]
273 pub openapi_require_output_declaration: bool,
274
275 #[serde(default = "default_true", alias = "reads_enabled")]
291 pub sql_reads_enabled: bool,
292
293 #[serde(default, alias = "allow_writes")]
295 pub sql_allow_writes: bool,
296
297 #[serde(default, alias = "allow_deletes")]
299 pub sql_allow_deletes: bool,
300
301 #[serde(default, alias = "allow_ddl")]
304 pub sql_allow_ddl: bool,
305
306 #[serde(default, alias = "allowed_statements")]
309 pub sql_allowed_statements: HashSet<String>,
310
311 #[serde(default, alias = "blocked_statements")]
313 pub sql_blocked_statements: HashSet<String>,
314
315 #[serde(default, alias = "blocked_tables")]
317 pub sql_blocked_tables: HashSet<String>,
318
319 #[serde(default, alias = "allowed_tables")]
321 pub sql_allowed_tables: HashSet<String>,
322
323 #[serde(default, alias = "blocked_columns")]
325 pub sql_blocked_columns: HashSet<String>,
326
327 #[serde(default = "default_sql_max_rows", alias = "max_rows")]
329 pub sql_max_rows: u64,
330
331 #[serde(default = "default_sql_max_joins", alias = "max_joins")]
333 pub sql_max_joins: u32,
334
335 #[serde(default = "default_true", alias = "require_where_on_writes")]
337 pub sql_require_where_on_writes: bool,
338
339 #[serde(default, alias = "require_limit")]
343 pub sql_require_limit: bool,
344
345 #[serde(default)]
350 pub action_tags: HashMap<String, String>,
351
352 #[serde(default = "default_max_depth")]
354 pub max_depth: u32,
355
356 #[serde(default = "default_max_field_count")]
358 pub max_field_count: u32,
359
360 #[serde(default = "default_max_cost")]
362 pub max_cost: u32,
363
364 #[serde(default)]
366 pub allowed_sensitive_categories: HashSet<String>,
367
368 #[serde(default = "default_token_ttl")]
370 pub token_ttl_seconds: i64,
371
372 #[serde(default = "default_auto_approve_levels")]
374 pub auto_approve_levels: Vec<RiskLevel>,
375
376 #[serde(default = "default_max_query_length")]
378 pub max_query_length: usize,
379
380 #[serde(default = "default_max_result_rows")]
382 pub max_result_rows: usize,
383
384 #[serde(default = "default_query_timeout")]
386 pub query_timeout_seconds: u32,
387
388 #[serde(default)]
390 pub server_id: Option<String>,
391
392 #[serde(default)]
399 pub sdk_operations: HashSet<String>,
400
401 #[serde(default)]
406 pub operations: Vec<OperationEntry>,
407}
408
409impl Default for CodeModeConfig {
410 fn default() -> Self {
411 Self {
412 enabled: false,
413 allow_mutations: false,
415 allowed_mutations: HashSet::new(),
416 blocked_mutations: HashSet::new(),
417 allow_introspection: false,
418 blocked_fields: HashSet::new(),
419 allowed_queries: HashSet::new(),
420 blocked_queries: HashSet::new(),
421 openapi_reads_enabled: true,
423 openapi_allow_writes: false,
424 openapi_allowed_writes: HashSet::new(),
425 openapi_blocked_writes: HashSet::new(),
426 openapi_allow_deletes: false,
427 openapi_allowed_deletes: HashSet::new(),
428 openapi_blocked_paths: HashSet::new(),
429 openapi_internal_blocked_fields: HashSet::new(),
430 openapi_output_blocked_fields: HashSet::new(),
431 openapi_require_output_declaration: false,
432 sql_reads_enabled: true,
434 sql_allow_writes: false,
435 sql_allow_deletes: false,
436 sql_allow_ddl: false,
437 sql_allowed_statements: HashSet::new(),
438 sql_blocked_statements: HashSet::new(),
439 sql_blocked_tables: HashSet::new(),
440 sql_allowed_tables: HashSet::new(),
441 sql_blocked_columns: HashSet::new(),
442 sql_max_rows: default_sql_max_rows(),
443 sql_max_joins: default_sql_max_joins(),
444 sql_require_where_on_writes: true,
445 sql_require_limit: false,
446 action_tags: HashMap::new(),
448 max_depth: default_max_depth(),
449 max_field_count: default_max_field_count(),
450 max_cost: default_max_cost(),
451 allowed_sensitive_categories: HashSet::new(),
452 token_ttl_seconds: default_token_ttl(),
453 auto_approve_levels: default_auto_approve_levels(),
454 max_query_length: default_max_query_length(),
455 max_result_rows: default_max_result_rows(),
456 query_timeout_seconds: default_query_timeout(),
457 server_id: None,
458 sdk_operations: HashSet::new(),
460 operations: Vec::new(),
461 }
462 }
463}
464
465#[derive(Deserialize)]
468struct TomlWrapper {
469 #[serde(default)]
470 code_mode: CodeModeConfig,
471}
472
473impl CodeModeConfig {
474 pub fn from_toml(toml_str: &str) -> Result<Self, toml::de::Error> {
489 let wrapper: TomlWrapper = toml::from_str(toml_str)?;
490 Ok(wrapper.code_mode)
491 }
492
493 pub fn enabled() -> Self {
495 Self {
496 enabled: true,
497 ..Default::default()
498 }
499 }
500
501 pub fn is_sdk_mode(&self) -> bool {
503 !self.sdk_operations.is_empty()
504 }
505
506 pub fn should_auto_approve(&self, risk_level: RiskLevel) -> bool {
508 self.auto_approve_levels.contains(&risk_level)
509 }
510
511 pub fn server_id(&self) -> &str {
518 self.server_id.as_deref().unwrap_or("unknown")
519 }
520
521 pub fn resolve_server_id(&mut self) {
532 if self.server_id.is_some() {
533 return;
534 }
535 self.server_id = resolve_server_id_from_env();
536 }
537
538 pub fn require_server_id(&self) -> Result<&str, ValidationError> {
544 self.server_id.as_deref().ok_or_else(|| {
545 ValidationError::ConfigError(
546 "server_id is not set. Set it in config.toml, PMCP_SERVER_ID env var, \
547 or AWS_LAMBDA_FUNCTION_NAME (Lambda). Without it, AVP authorization \
548 will default-deny silently."
549 .into(),
550 )
551 })
552 }
553
554 pub fn to_server_config_entity(&self) -> crate::policy::ServerConfigEntity {
556 crate::policy::ServerConfigEntity {
557 server_id: self.server_id().to_string(),
558 server_type: "graphql".to_string(),
559 allow_write: self.allow_mutations,
560 allow_delete: self.allow_mutations,
561 allow_admin: self.allow_introspection,
562 allowed_operations: self.allowed_mutations.clone(),
563 blocked_operations: self.blocked_mutations.clone(),
564 max_depth: self.max_depth,
565 max_field_count: self.max_field_count,
566 max_cost: self.max_cost,
567 max_api_calls: 50,
568 blocked_fields: self.blocked_fields.clone(),
569 allowed_sensitive_categories: self.allowed_sensitive_categories.clone(),
570 }
571 }
572
573 #[cfg(feature = "openapi-code-mode")]
575 pub fn to_openapi_server_entity(&self) -> crate::policy::OpenAPIServerEntity {
576 let mut allowed_operations = self.openapi_allowed_writes.clone();
577 allowed_operations.extend(self.openapi_allowed_deletes.clone());
578
579 let write_mode = if !self.openapi_allow_writes {
580 "deny_all"
581 } else if !self.openapi_allowed_writes.is_empty() {
582 "allowlist"
583 } else if !self.openapi_blocked_writes.is_empty() {
584 "blocklist"
585 } else {
586 "allow_all"
587 };
588
589 crate::policy::OpenAPIServerEntity {
590 server_id: self.server_id().to_string(),
591 server_type: "openapi".to_string(),
592 allow_write: self.openapi_allow_writes,
593 allow_delete: self.openapi_allow_deletes,
594 allow_admin: false,
595 write_mode: write_mode.to_string(),
596 max_depth: self.max_depth,
597 max_cost: self.max_cost,
598 max_api_calls: 50,
599 max_loop_iterations: 100,
600 max_script_length: self.max_query_length as u32,
601 max_nesting_depth: self.max_depth,
602 execution_timeout_seconds: self.query_timeout_seconds,
603 allowed_operations,
604 blocked_operations: self.openapi_blocked_writes.clone(),
605 allowed_methods: HashSet::new(),
606 blocked_methods: HashSet::new(),
607 allowed_path_patterns: HashSet::new(),
608 blocked_path_patterns: self.openapi_blocked_paths.clone(),
609 sensitive_path_patterns: self.openapi_blocked_paths.clone(),
610 auto_approve_read_only: self.openapi_reads_enabled,
611 max_api_calls_for_auto_approve: 10,
612 internal_blocked_fields: self.openapi_internal_blocked_fields.clone(),
613 output_blocked_fields: self.openapi_output_blocked_fields.clone(),
614 require_output_declaration: self.openapi_require_output_declaration,
615 }
616 }
617
618 #[cfg(feature = "sql-code-mode")]
620 pub fn to_sql_server_entity(&self) -> crate::policy::SqlServerEntity {
621 crate::policy::SqlServerEntity {
622 server_id: self.server_id().to_string(),
623 server_type: "sql".to_string(),
624 allow_write: self.sql_allow_writes,
625 allow_delete: self.sql_allow_deletes,
626 allow_admin: self.sql_allow_ddl,
627 max_rows: self.sql_max_rows,
628 max_joins: self.sql_max_joins,
629 allowed_operations: self.sql_allowed_statements.clone(),
630 blocked_operations: self.sql_blocked_statements.clone(),
631 blocked_tables: self.sql_blocked_tables.clone(),
632 blocked_columns: self.sql_blocked_columns.clone(),
633 allowed_tables: self.sql_allowed_tables.clone(),
634 }
635 }
636}
637
638fn default_true() -> bool {
639 true
640}
641
642fn default_token_ttl() -> i64 {
643 300 }
645
646fn default_auto_approve_levels() -> Vec<RiskLevel> {
647 vec![RiskLevel::Low]
648}
649
650fn default_max_query_length() -> usize {
651 10000
652}
653
654fn default_max_result_rows() -> usize {
655 10000
656}
657
658fn default_query_timeout() -> u32 {
659 30
660}
661
662fn default_max_depth() -> u32 {
663 10
664}
665
666fn default_max_field_count() -> u32 {
667 100
668}
669
670fn default_max_cost() -> u32 {
671 1000
672}
673
674fn default_sql_max_rows() -> u64 {
675 10_000
676}
677
678fn default_sql_max_joins() -> u32 {
679 5
680}
681
682#[cfg(test)]
683mod tests {
684 use super::*;
685
686 #[test]
687 fn test_default_config() {
688 let config = CodeModeConfig::default();
689 assert!(!config.enabled);
690 assert!(!config.allow_mutations);
691 assert_eq!(config.token_ttl_seconds, 300);
692 assert_eq!(config.auto_approve_levels, vec![RiskLevel::Low]);
693 }
694
695 #[test]
696 fn test_enabled_config() {
697 let config = CodeModeConfig::enabled();
698 assert!(config.enabled);
699 }
700
701 #[test]
702 fn test_auto_approve() {
703 let config = CodeModeConfig::default();
704 assert!(config.should_auto_approve(RiskLevel::Low));
705 assert!(!config.should_auto_approve(RiskLevel::Medium));
706 assert!(!config.should_auto_approve(RiskLevel::High));
707 assert!(!config.should_auto_approve(RiskLevel::Critical));
708 }
709
710 #[test]
711 fn test_operation_registry_from_entries() {
712 let entries = vec![
713 OperationEntry {
714 id: "getCostAnomalies".to_string(),
715 category: "read".to_string(),
716 description: "Get cost anomalies".to_string(),
717 path: Some("/getCostAnomalies".to_string()),
718 },
719 OperationEntry {
720 id: "listInstances".to_string(),
721 category: "read".to_string(),
722 description: "List EC2 instances".to_string(),
723 path: Some("/listInstances".to_string()),
724 },
725 ];
726 let registry = OperationRegistry::from_entries(&entries);
727 assert_eq!(
728 registry.lookup("/getCostAnomalies"),
729 Some("getCostAnomalies")
730 );
731 assert_eq!(registry.lookup("/listInstances"), Some("listInstances"));
732 }
733
734 #[test]
735 fn test_operation_registry_lookup_unregistered() {
736 let entries = vec![OperationEntry {
737 id: "getCostAnomalies".to_string(),
738 category: "read".to_string(),
739 description: String::new(),
740 path: Some("/getCostAnomalies".to_string()),
741 }];
742 let registry = OperationRegistry::from_entries(&entries);
743 assert_eq!(registry.lookup("/unknownPath"), None);
744 assert_eq!(registry.lookup(""), None);
745 }
746
747 #[test]
748 fn test_operation_registry_lookup_category() {
749 let entries = vec![
750 OperationEntry {
751 id: "getCostAnomalies".to_string(),
752 category: "read".to_string(),
753 description: String::new(),
754 path: Some("/getCostAnomalies".to_string()),
755 },
756 OperationEntry {
757 id: "deleteReservation".to_string(),
758 category: "delete".to_string(),
759 description: String::new(),
760 path: Some("/deleteReservation".to_string()),
761 },
762 OperationEntry {
763 id: "updateBudget".to_string(),
764 category: "write".to_string(),
765 description: String::new(),
766 path: Some("/updateBudget".to_string()),
767 },
768 ];
769 let registry = OperationRegistry::from_entries(&entries);
770 assert_eq!(registry.lookup_category("/getCostAnomalies"), Some("read"));
771 assert_eq!(
772 registry.lookup_category("/deleteReservation"),
773 Some("delete")
774 );
775 assert_eq!(registry.lookup_category("/updateBudget"), Some("write"));
776 assert_eq!(registry.lookup_category("/unknownPath"), None);
777 }
778
779 #[test]
780 fn test_operation_registry_empty_category_excluded() {
781 let entries = vec![OperationEntry {
782 id: "legacyOp".to_string(),
783 category: String::new(), description: String::new(),
785 path: Some("/legacyOp".to_string()),
786 }];
787 let registry = OperationRegistry::from_entries(&entries);
788 assert_eq!(registry.lookup("/legacyOp"), Some("legacyOp"));
790 assert_eq!(registry.lookup_category("/legacyOp"), None);
792 }
793
794 #[test]
795 fn test_operation_registry_is_empty() {
796 let empty_registry = OperationRegistry::from_entries(&[]);
797 assert!(empty_registry.is_empty());
798
799 let entries = vec![OperationEntry {
800 id: "op1".to_string(),
801 category: "read".to_string(),
802 description: String::new(),
803 path: Some("/op1".to_string()),
804 }];
805 let registry = OperationRegistry::from_entries(&entries);
806 assert!(!registry.is_empty());
807 }
808
809 fn op(id: &str, category: &str, path: &str) -> OperationEntry {
810 OperationEntry {
811 id: id.to_string(),
812 category: category.to_string(),
813 description: String::new(),
814 path: Some(path.to_string()),
815 }
816 }
817
818 #[test]
819 fn test_operation_registry_matches_templates() {
820 let registry = OperationRegistry::from_entries(&[
821 op("getItem", "read", "/items/{id}"),
822 op("searchItems", "read", "/items/search"),
823 ]);
824 assert_eq!(registry.lookup("/items/42"), Some("getItem"));
825 assert_eq!(registry.lookup("/items/{id}"), Some("getItem"));
826 assert_eq!(registry.lookup("/items/search"), Some("searchItems"));
828 assert_eq!(registry.lookup("/items/42/owner"), None);
830 assert_eq!(registry.lookup("/items/"), None);
831 let literal_only = OperationRegistry::from_entries(&[op("s", "read", "/items/search")]);
833 assert_eq!(literal_only.lookup("/items/{...}"), None);
834 }
835
836 #[test]
837 fn test_operation_registry_method_prefix() {
838 let registry = OperationRegistry::from_entries(&[
839 op("listItems", "read", "GET /items"),
840 op("createItem", "write", "POST:/items"),
841 ]);
842 let id = |m, p| registry.lookup_entry(Some(m), p).map(|e| e.id.as_str());
843 assert_eq!(id("GET", "/items"), Some("listItems"));
844 assert_eq!(id("post", "/items"), Some("createItem"));
845 assert_eq!(id("DELETE", "/items"), None);
846 }
847
848 #[test]
849 fn test_operation_registry_ties_resolve_to_the_strictest_category() {
850 let registry = OperationRegistry::from_entries(&[
851 op("a", "read", "/x/{id}"),
852 op("b", "admin", "/x/{key}"),
853 ]);
854 assert_eq!(registry.lookup_category("/x/1"), Some("admin"));
855 }
856
857 #[test]
858 fn test_operation_entry_deserialization() {
859 let toml_str = r#"
860id = "getCostAnomalies"
861category = "read"
862description = "Get cost anomalies"
863path = "/getCostAnomalies"
864"#;
865 let entry: OperationEntry =
866 toml::from_str(toml_str).expect("Failed to deserialize OperationEntry");
867 assert_eq!(entry.id, "getCostAnomalies");
868 assert_eq!(entry.category, "read");
869 assert_eq!(entry.description, "Get cost anomalies");
870 assert_eq!(entry.path, Some("/getCostAnomalies".to_string()));
871 }
872
873 #[test]
874 fn test_code_mode_config_with_operations() {
875 let toml_str = r#"
876enabled = true
877
878[[operations]]
879id = "getCostAnomalies"
880category = "read"
881description = "Get cost anomalies"
882path = "/getCostAnomalies"
883
884[[operations]]
885id = "listInstances"
886category = "read"
887path = "/listInstances"
888"#;
889 let config: CodeModeConfig = toml::from_str(toml_str).expect("Failed to deserialize");
890 assert!(config.enabled);
891 assert_eq!(config.operations.len(), 2);
892 assert_eq!(config.operations[0].id, "getCostAnomalies");
893 assert_eq!(config.operations[1].id, "listInstances");
894 }
895
896 #[test]
897 fn test_code_mode_config_without_operations_defaults_to_empty() {
898 let toml_str = r#"
899enabled = true
900"#;
901 let config: CodeModeConfig = toml::from_str(toml_str).expect("Failed to deserialize");
902 assert!(config.enabled);
903 assert!(config.operations.is_empty());
904 }
905
906 #[test]
907 fn test_from_toml_extracts_code_mode_section() {
908 let toml_str = r#"
909[server]
910name = "cost-coach"
911type = "openapi-api"
912
913[code_mode]
914enabled = true
915token_ttl_seconds = 600
916server_id = "cost-coach"
917
918[[code_mode.operations]]
919id = "getCostAndUsage"
920category = "read"
921description = "Historical cost and usage data"
922path = "/getCostAndUsage"
923
924[[code_mode.operations]]
925id = "getCostAnomalies"
926category = "read"
927description = "Cost anomalies detected by AWS"
928path = "/getCostAnomalies"
929
930[[tools]]
931name = "some_tool"
932"#;
933 let config = CodeModeConfig::from_toml(toml_str).expect("Failed to parse");
934 assert!(config.enabled);
935 assert_eq!(config.token_ttl_seconds, 600);
936 assert_eq!(config.server_id, Some("cost-coach".to_string()));
937 assert_eq!(config.operations.len(), 2);
938 assert_eq!(config.operations[0].id, "getCostAndUsage");
939 assert_eq!(config.operations[1].id, "getCostAnomalies");
940 assert_eq!(
941 config.operations[0].path,
942 Some("/getCostAndUsage".to_string())
943 );
944 }
945
946 #[test]
947 fn test_from_toml_missing_code_mode_returns_default() {
948 let toml_str = r#"
949[server]
950name = "some-server"
951"#;
952 let config = CodeModeConfig::from_toml(toml_str).expect("Failed to parse");
953 assert!(!config.enabled);
954 assert!(config.operations.is_empty());
955 assert_eq!(config.token_ttl_seconds, 300); }
957
958 use std::sync::Mutex;
967 static ENV_LOCK: Mutex<()> = Mutex::new(());
968
969 struct EnvGuard {
970 _lock: std::sync::MutexGuard<'static, ()>,
971 }
972
973 impl EnvGuard {
974 fn acquire() -> Self {
975 let lock = ENV_LOCK
976 .lock()
977 .unwrap_or_else(|poisoned| poisoned.into_inner());
978 std::env::remove_var("PMCP_SERVER_ID");
979 std::env::remove_var("AWS_LAMBDA_FUNCTION_NAME");
980 Self { _lock: lock }
981 }
982 }
983
984 impl Drop for EnvGuard {
985 fn drop(&mut self) {
986 std::env::remove_var("PMCP_SERVER_ID");
987 std::env::remove_var("AWS_LAMBDA_FUNCTION_NAME");
988 }
989 }
990
991 #[test]
992 fn resolve_server_id_from_explicit_config_takes_precedence() {
993 let _g = EnvGuard::acquire();
994 std::env::set_var("PMCP_SERVER_ID", "from-env");
995
996 let mut config = CodeModeConfig {
997 server_id: Some("from-config".to_string()),
998 ..Default::default()
999 };
1000 config.resolve_server_id();
1001
1002 assert_eq!(config.server_id.as_deref(), Some("from-config"));
1003 }
1004
1005 #[test]
1006 fn resolve_server_id_from_pmcp_env() {
1007 let _g = EnvGuard::acquire();
1008 std::env::set_var("PMCP_SERVER_ID", "my-server");
1009
1010 let mut config = CodeModeConfig::default();
1011 config.resolve_server_id();
1012
1013 assert_eq!(config.server_id.as_deref(), Some("my-server"));
1014 }
1015
1016 #[test]
1017 fn resolve_server_id_from_lambda_env() {
1018 let _g = EnvGuard::acquire();
1019 std::env::set_var("AWS_LAMBDA_FUNCTION_NAME", "my-lambda-fn");
1020
1021 let mut config = CodeModeConfig::default();
1022 config.resolve_server_id();
1023
1024 assert_eq!(config.server_id.as_deref(), Some("my-lambda-fn"));
1025 }
1026
1027 #[test]
1028 fn resolve_server_id_pmcp_wins_over_lambda() {
1029 let _g = EnvGuard::acquire();
1030 std::env::set_var("PMCP_SERVER_ID", "explicit");
1031 std::env::set_var("AWS_LAMBDA_FUNCTION_NAME", "lambda-fn");
1032
1033 let mut config = CodeModeConfig::default();
1034 config.resolve_server_id();
1035
1036 assert_eq!(config.server_id.as_deref(), Some("explicit"));
1037 }
1038
1039 #[test]
1040 fn resolve_server_id_leaves_none_when_unset() {
1041 let _g = EnvGuard::acquire();
1042 let mut config = CodeModeConfig::default();
1043 config.resolve_server_id();
1044 assert!(config.server_id.is_none());
1045 }
1046
1047 #[test]
1048 fn require_server_id_errors_when_unset() {
1049 let config = CodeModeConfig::default();
1050 let result = config.require_server_id();
1051 assert!(matches!(result, Err(ValidationError::ConfigError(_))));
1052 }
1053
1054 #[test]
1055 fn require_server_id_returns_value_when_set() {
1056 let config = CodeModeConfig {
1057 server_id: Some("my-server".to_string()),
1058 ..Default::default()
1059 };
1060 assert_eq!(config.require_server_id().unwrap(), "my-server");
1061 }
1062
1063 #[test]
1064 fn resolve_server_id_from_env_free_fn_treats_empty_as_unset() {
1065 let _g = EnvGuard::acquire();
1066 std::env::set_var("PMCP_SERVER_ID", "");
1067 assert_eq!(resolve_server_id_from_env(), None);
1068 }
1069
1070 #[test]
1075 fn sql_config_accepts_unprefixed_toml_names() {
1076 let toml_str = r#"
1077enabled = true
1078allow_writes = true
1079allow_deletes = true
1080allow_ddl = true
1081allowed_tables = ["users", "orders"]
1082blocked_tables = ["secrets"]
1083blocked_columns = ["password", "ssn"]
1084max_rows = 5000
1085max_joins = 3
1086require_where_on_writes = false
1087"#;
1088 let config: CodeModeConfig =
1089 toml::from_str(toml_str).expect("Failed to deserialize with unprefixed aliases");
1090
1091 assert!(config.enabled);
1092 assert!(config.sql_allow_writes);
1093 assert!(config.sql_allow_deletes);
1094 assert!(config.sql_allow_ddl);
1095 assert!(config.sql_allowed_tables.contains("users"));
1096 assert!(config.sql_allowed_tables.contains("orders"));
1097 assert!(config.sql_blocked_tables.contains("secrets"));
1098 assert!(config.sql_blocked_columns.contains("password"));
1099 assert_eq!(config.sql_max_rows, 5000);
1100 assert_eq!(config.sql_max_joins, 3);
1101 assert!(!config.sql_require_where_on_writes);
1102 }
1103
1104 #[test]
1105 fn sql_config_accepts_prefixed_toml_names() {
1106 let toml_str = r#"
1107enabled = true
1108sql_allow_writes = true
1109sql_blocked_tables = ["secrets"]
1110sql_max_rows = 5000
1111"#;
1112 let config: CodeModeConfig =
1113 toml::from_str(toml_str).expect("Failed to deserialize with prefixed names");
1114
1115 assert!(config.sql_allow_writes);
1116 assert!(config.sql_blocked_tables.contains("secrets"));
1117 assert_eq!(config.sql_max_rows, 5000);
1118 }
1119}