Skip to main content

pmcp_code_mode/
config.rs

1//! Code Mode configuration.
2
3use crate::types::{RiskLevel, ValidationError};
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6use std::collections::HashSet;
7
8/// Resolve a `server_id` from environment variables.
9///
10/// Checks, in order:
11/// 1. `PMCP_SERVER_ID`
12/// 2. `AWS_LAMBDA_FUNCTION_NAME` (Lambda runtime)
13///
14/// Returns `None` if neither is set. Empty strings are treated as unset.
15///
16/// This is the same resolution chain used by
17/// [`CodeModeConfig::resolve_server_id`] — exposed as a free function so tests
18/// and non-pipeline code can share it.
19pub 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/// A single declared operation in Code Mode configuration.
31/// Maps a raw API path to a canonical plain-name ID for Cedar policies.
32#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct OperationEntry {
34    /// Canonical operation ID (plain name, no method prefix).
35    /// This is what appears in Cedar policy calledOperations.
36    pub id: String,
37
38    /// Action category for AVP action routing.
39    /// Values: "read", "write", "delete", "admin"
40    pub category: String,
41
42    /// Human-readable description for admin UI and LLM context.
43    #[serde(default)]
44    pub description: String,
45
46    /// API path this ID maps to (e.g., "/getCostAnomalies"), matched against
47    /// the `api.*()` calls of a script.
48    ///
49    /// May start with an HTTP method (`"GET /items/{id}"` or
50    /// `"GET:/items/{id}"`), in which case only calls with that method match.
51    /// A `{param}` segment matches any one segment. See
52    /// [`OperationRegistry::lookup_entry`].
53    #[serde(default)]
54    pub path: Option<String>,
55}
56
57/// Registry built from [[code_mode.operations]] config entries.
58/// Maps API calls to canonical operation IDs and categories.
59#[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    /// How specifically this entry matches a call: `None` when it does not
73    /// match, otherwise (exact path, literal segment count, has a method).
74    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
109/// Split an optional leading HTTP method off an entry path.
110fn 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
127/// Rank used to break ties between equally specific entries: the stricter
128/// declared category wins, and an unrecognised one counts as strictest so it
129/// is refused rather than silently passed over.
130fn 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    /// The entry a call to `path` (with `method`, when known) maps to.
158    ///
159    /// The most specific match wins: an exact path over a template, then the
160    /// template with more literal segments, then an entry that names the
161    /// method. A `{param}` segment in an entry matches any one non-empty
162    /// segment; a segment of `path` that is itself a placeholder (contains
163    /// `{`) matches only a `{param}` segment. Among equally specific entries
164    /// the one with the strictest category is returned.
165    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    /// Look up the declared category for a path (e.g., "read", "write", "delete", "admin").
182    /// Returns `None` if the path has no registry entry or no category declared.
183    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/// Configuration for Code Mode.
195#[derive(Debug, Clone, Serialize, Deserialize)]
196pub struct CodeModeConfig {
197    /// Whether Code Mode is enabled for this server
198    #[serde(default)]
199    pub enabled: bool,
200
201    // ========================================================================
202    // GraphQL-specific settings
203    // ========================================================================
204    /// Whether to allow mutations (MVP: false)
205    #[serde(default)]
206    pub allow_mutations: bool,
207
208    /// Allowed mutation names (whitelist). If empty and allow_mutations=true, all are allowed.
209    #[serde(default)]
210    pub allowed_mutations: HashSet<String>,
211
212    /// Blocked mutation names (blacklist). Always blocked even if allow_mutations=true.
213    #[serde(default)]
214    pub blocked_mutations: HashSet<String>,
215
216    /// Whether to allow introspection queries
217    #[serde(default)]
218    pub allow_introspection: bool,
219
220    /// Fields that should never be returned (Type.field format) - GraphQL
221    #[serde(default)]
222    pub blocked_fields: HashSet<String>,
223
224    /// Allowed query names (whitelist). If empty and mode is allowlist, none are allowed.
225    #[serde(default)]
226    pub allowed_queries: HashSet<String>,
227
228    /// Blocked query names (blocklist). Always blocked even if reads enabled.
229    #[serde(default)]
230    pub blocked_queries: HashSet<String>,
231
232    // ========================================================================
233    // OpenAPI-specific settings
234    // ========================================================================
235    /// Whether read operations (GET) are enabled (default: true)
236    #[serde(default = "default_true")]
237    pub openapi_reads_enabled: bool,
238
239    /// Whether write operations (POST, PUT, PATCH) are allowed globally
240    #[serde(default)]
241    pub openapi_allow_writes: bool,
242
243    /// Allowed write operations (operationId or "METHOD /path")
244    #[serde(default)]
245    pub openapi_allowed_writes: HashSet<String>,
246
247    /// Blocked write operations
248    #[serde(default)]
249    pub openapi_blocked_writes: HashSet<String>,
250
251    /// Whether delete operations (DELETE) are allowed globally
252    #[serde(default)]
253    pub openapi_allow_deletes: bool,
254
255    /// Allowed delete operations (operationId or "METHOD /path")
256    #[serde(default)]
257    pub openapi_allowed_deletes: HashSet<String>,
258
259    /// Blocked paths (glob patterns like "/admin/*")
260    #[serde(default)]
261    pub openapi_blocked_paths: HashSet<String>,
262
263    /// Fields that are stripped from API responses entirely (no access)
264    #[serde(default)]
265    pub openapi_internal_blocked_fields: HashSet<String>,
266
267    /// Fields that can be used internally but not in script output
268    #[serde(default)]
269    pub openapi_output_blocked_fields: HashSet<String>,
270
271    /// Whether scripts must declare their return type with @returns
272    #[serde(default)]
273    pub openapi_require_output_declaration: bool,
274
275    // ========================================================================
276    // SQL-specific settings
277    //
278    // SQL fields accept both their prefixed name (`sql_allow_writes`) and the
279    // unprefixed natural form (`allow_writes`). Downstream SQL servers can use
280    // the unprefixed names in their `[code_mode]` block without a manual
281    // conversion layer:
282    //
283    //     [code_mode]
284    //     reads_enabled = true    # same as sql_reads_enabled
285    //     allow_writes = false    # same as sql_allow_writes
286    //     blocked_tables = ["secrets"]
287    //     max_rows = 5000
288    // ========================================================================
289    /// Whether SELECT statements are enabled (default: true).
290    #[serde(default = "default_true", alias = "reads_enabled")]
291    pub sql_reads_enabled: bool,
292
293    /// Whether INSERT/UPDATE/MERGE statements are allowed globally.
294    #[serde(default, alias = "allow_writes")]
295    pub sql_allow_writes: bool,
296
297    /// Whether DELETE/TRUNCATE statements are allowed globally.
298    #[serde(default, alias = "allow_deletes")]
299    pub sql_allow_deletes: bool,
300
301    /// Whether DDL (CREATE/ALTER/DROP/GRANT/REVOKE) is allowed globally.
302    /// Default is `false` — DDL is almost never appropriate for LLM-generated code.
303    #[serde(default, alias = "allow_ddl")]
304    pub sql_allow_ddl: bool,
305
306    /// Allowed statement types ("SELECT"/"INSERT"/"UPDATE"/"DELETE"/"DDL").
307    /// If non-empty, only statement types in this set are allowed.
308    #[serde(default, alias = "allowed_statements")]
309    pub sql_allowed_statements: HashSet<String>,
310
311    /// Blocked statement types. Always blocked even if globally allowed.
312    #[serde(default, alias = "blocked_statements")]
313    pub sql_blocked_statements: HashSet<String>,
314
315    /// Tables that are always forbidden (blocklist mode).
316    #[serde(default, alias = "blocked_tables")]
317    pub sql_blocked_tables: HashSet<String>,
318
319    /// If non-empty, only these tables can be accessed (allowlist mode).
320    #[serde(default, alias = "allowed_tables")]
321    pub sql_allowed_tables: HashSet<String>,
322
323    /// Columns that may not be referenced in any statement (e.g., `password`, `ssn`).
324    #[serde(default, alias = "blocked_columns")]
325    pub sql_blocked_columns: HashSet<String>,
326
327    /// Maximum row-count estimate allowed (based on LIMIT or default estimate).
328    #[serde(default = "default_sql_max_rows", alias = "max_rows")]
329    pub sql_max_rows: u64,
330
331    /// Maximum number of JOINs in a single statement.
332    #[serde(default = "default_sql_max_joins", alias = "max_joins")]
333    pub sql_max_joins: u32,
334
335    /// Whether to require a WHERE clause for UPDATE/DELETE statements.
336    #[serde(default = "default_true", alias = "require_where_on_writes")]
337    pub sql_require_where_on_writes: bool,
338
339    /// Whether read-only (SELECT-class) statements MUST declare a LIMIT.
340    /// Opt-in safety guard; default false (no behavior change for configs
341    /// that omit it). Enforced in check_sql_config_authorization.
342    #[serde(default, alias = "require_limit")]
343    pub sql_require_limit: bool,
344
345    // ========================================================================
346    // Common settings
347    // ========================================================================
348    /// Action tags to override inferred actions for specific operations.
349    #[serde(default)]
350    pub action_tags: HashMap<String, String>,
351
352    /// Maximum query depth
353    #[serde(default = "default_max_depth")]
354    pub max_depth: u32,
355
356    /// Maximum field count per query
357    #[serde(default = "default_max_field_count")]
358    pub max_field_count: u32,
359
360    /// Maximum estimated query cost
361    #[serde(default = "default_max_cost")]
362    pub max_cost: u32,
363
364    /// Allowed sensitive data categories
365    #[serde(default)]
366    pub allowed_sensitive_categories: HashSet<String>,
367
368    /// Token time-to-live in seconds
369    #[serde(default = "default_token_ttl")]
370    pub token_ttl_seconds: i64,
371
372    /// Risk levels that can be auto-approved without human confirmation
373    #[serde(default = "default_auto_approve_levels")]
374    pub auto_approve_levels: Vec<RiskLevel>,
375
376    /// Maximum query length in characters
377    #[serde(default = "default_max_query_length")]
378    pub max_query_length: usize,
379
380    /// Maximum result rows to return
381    #[serde(default = "default_max_result_rows")]
382    pub max_result_rows: usize,
383
384    /// Query execution timeout in seconds
385    #[serde(default = "default_query_timeout")]
386    pub query_timeout_seconds: u32,
387
388    /// Server ID for token generation
389    #[serde(default)]
390    pub server_id: Option<String>,
391
392    // ========================================================================
393    // SDK-backed settings
394    // ========================================================================
395    /// Allowed SDK operation names for SDK-backed Code Mode.
396    /// When non-empty, Code Mode uses SDK dispatch instead of HTTP.
397    /// Operations are validated at compile time — unlisted names are rejected.
398    #[serde(default)]
399    pub sdk_operations: HashSet<String>,
400
401    /// Declared operations for plain-name ID mapping in Cedar entities.
402    /// Parsed from [[code_mode.operations]] TOML sections.
403    /// When non-empty, ScriptEntity calledOperations uses IDs from the registry
404    /// built from these entries. Unregistered paths fall back to METHOD:/path.
405    #[serde(default)]
406    pub operations: Vec<OperationEntry>,
407}
408
409impl Default for CodeModeConfig {
410    fn default() -> Self {
411        Self {
412            enabled: false,
413            // GraphQL
414            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
422            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
433            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            // Common
447            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
459            sdk_operations: HashSet::new(),
460            operations: Vec::new(),
461        }
462    }
463}
464
465/// Wrapper for deserializing the `[code_mode]` section from a full TOML config file.
466/// The file may contain other sections (`[server]`, `[[tools]]`, etc.) which are ignored.
467#[derive(Deserialize)]
468struct TomlWrapper {
469    #[serde(default)]
470    code_mode: CodeModeConfig,
471}
472
473impl CodeModeConfig {
474    /// Parse `CodeModeConfig` from a full TOML config string.
475    ///
476    /// Extracts the `[code_mode]` section (including `[[code_mode.operations]]`)
477    /// and ignores all other sections. This is the recommended way for external
478    /// servers to build their config from `config.toml`:
479    ///
480    /// ```rust,ignore
481    /// const CONFIG_TOML: &str = include_str!("../../config.toml");
482    ///
483    /// let config = CodeModeConfig::from_toml(CONFIG_TOML)
484    ///     .expect("Invalid code_mode section in config.toml");
485    /// ```
486    ///
487    /// If the TOML has no `[code_mode]` section, returns `CodeModeConfig::default()`.
488    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    /// Create a new config with Code Mode enabled.
494    pub fn enabled() -> Self {
495        Self {
496            enabled: true,
497            ..Default::default()
498        }
499    }
500
501    /// Returns true if this config enables SDK-backed Code Mode.
502    pub fn is_sdk_mode(&self) -> bool {
503        !self.sdk_operations.is_empty()
504    }
505
506    /// Check if a risk level should be auto-approved.
507    pub fn should_auto_approve(&self, risk_level: RiskLevel) -> bool {
508        self.auto_approve_levels.contains(&risk_level)
509    }
510
511    /// Get the server ID, falling back to a default.
512    ///
513    /// **Note:** The `"unknown"` fallback produces silent AVP default-deny failures
514    /// (no Cedar policy matches a server_id of `"unknown"`). Prefer
515    /// [`resolve_server_id`](Self::resolve_server_id) to auto-fill from environment,
516    /// or [`require_server_id`](Self::require_server_id) to fail fast.
517    pub fn server_id(&self) -> &str {
518        self.server_id.as_deref().unwrap_or("unknown")
519    }
520
521    /// Auto-resolve `server_id` from environment if not already set.
522    ///
523    /// Resolution order:
524    /// 1. `self.server_id` (if already set, e.g., from TOML) — no change
525    /// 2. `PMCP_SERVER_ID` env var
526    /// 3. `AWS_LAMBDA_FUNCTION_NAME` env var (Lambda runtime)
527    /// 4. Left as `None` — caller is responsible for handling
528    ///
529    /// [`ValidationPipeline`](crate::ValidationPipeline) constructors call this
530    /// automatically, so wrappers rarely need to invoke it directly.
531    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    /// Return the `server_id`, or an error if not resolved.
539    ///
540    /// Use this in production code paths that require AVP authorization —
541    /// it fails fast with a clear message instead of letting `"unknown"`
542    /// reach AVP and produce a silent default-deny.
543    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    /// Convert to ServerConfigEntity for policy evaluation.
555    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    /// Convert to OpenAPIServerEntity for policy evaluation (OpenAPI Code Mode).
574    #[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    /// Convert to `SqlServerEntity` for policy evaluation (SQL Code Mode).
619    #[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 // 5 minutes
644}
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(), // empty = not declared
784            description: String::new(),
785            path: Some("/legacyOp".to_string()),
786        }];
787        let registry = OperationRegistry::from_entries(&entries);
788        // ID lookup still works
789        assert_eq!(registry.lookup("/legacyOp"), Some("legacyOp"));
790        // Category lookup returns None for empty category
791        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        // A literal segment beats a template segment.
827        assert_eq!(registry.lookup("/items/search"), Some("searchItems"));
828        // Segment counts must agree, and a param needs a non-empty segment.
829        assert_eq!(registry.lookup("/items/42/owner"), None);
830        assert_eq!(registry.lookup("/items/"), None);
831        // A placeholder in the call matches only a param, never a literal.
832        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); // default
956    }
957
958    // =========================================================================
959    // server_id resolution tests
960    //
961    // These tests mutate process-wide env vars. Cargo parallelizes tests across
962    // threads in the same process, so a shared Mutex serializes them — without
963    // this, set_var/remove_var in one test would race with another.
964    // =========================================================================
965
966    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    // =========================================================================
1071    // SQL TOML DX tests (serde aliases)
1072    // =========================================================================
1073
1074    #[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}