1use pmcp::types::{ToolAnnotations, ToolInfo};
12use serde::{Deserialize, Serialize};
13use serde_json::{json, Value};
14
15use crate::types::{
16 PolicyViolation, RiskLevel, UnifiedAction, ValidationMetadata, ValidationResult,
17};
18
19#[derive(Debug, Clone, Serialize, Deserialize)]
25pub struct ValidationResponse {
26 #[serde(flatten)]
28 pub result: ValidationResult,
29
30 pub auto_approved: bool,
32
33 pub action: Option<UnifiedAction>,
35
36 #[serde(skip_serializing_if = "Option::is_none")]
38 pub validated_code_hash: Option<String>,
39}
40
41impl ValidationResponse {
42 pub fn success(
44 explanation: String,
45 risk_level: RiskLevel,
46 approval_token: String,
47 metadata: ValidationMetadata,
48 ) -> Self {
49 Self {
50 result: ValidationResult::success(explanation, risk_level, approval_token, metadata),
51 auto_approved: false,
52 action: None,
53 validated_code_hash: None,
54 }
55 }
56
57 pub fn failure(violations: Vec<PolicyViolation>, metadata: ValidationMetadata) -> Self {
59 Self {
60 result: ValidationResult::failure(violations, metadata),
61 auto_approved: false,
62 action: None,
63 validated_code_hash: None,
64 }
65 }
66
67 pub fn from_result(result: ValidationResult) -> Self {
69 Self {
70 result,
71 auto_approved: false,
72 action: None,
73 validated_code_hash: None,
74 }
75 }
76
77 pub fn with_code_hash(mut self, hash: String) -> Self {
79 self.validated_code_hash = Some(hash);
80 self
81 }
82
83 pub fn with_action(mut self, action: UnifiedAction) -> Self {
85 self.action = Some(action);
86 self
87 }
88
89 pub fn with_auto_approved(mut self, auto_approved: bool) -> Self {
91 self.auto_approved = auto_approved;
92 self
93 }
94
95 pub fn with_warnings(mut self, warnings: Vec<String>) -> Self {
97 self.result.warnings = warnings;
98 self
99 }
100
101 pub fn to_json_response(&self) -> (Value, bool) {
116 let is_valid = self.result.is_valid;
117 let explanation = if is_valid {
118 self.result.explanation.as_str()
119 } else {
120 ""
121 };
122 let (accessed_types, accessed_fields) = if is_valid {
123 (
124 self.result.metadata.accessed_types.clone(),
125 self.result.metadata.accessed_fields.clone(),
126 )
127 } else {
128 (Default::default(), Default::default())
129 };
130 let response = json!({
131 "valid": self.result.is_valid,
132 "explanation": explanation,
133 "risk_level": format!("{}", self.result.risk_level),
134 "approval_token": self.result.approval_token,
135 "action": self.action.as_ref().map(|a| a.to_string()),
136 "auto_approved": self.auto_approved,
137 "warnings": self.result.warnings,
138 "violations": self.result.violations.iter().map(|v| json!({
139 "policy": v.policy_name,
140 "rule": v.rule,
141 "message": v.message,
142 "suggestion": v.suggestion
143 })).collect::<Vec<_>>(),
144 "validated_code_hash": self.validated_code_hash,
145 "metadata": {
146 "is_read_only": self.result.metadata.is_read_only,
147 "accessed_types": accessed_types,
148 "accessed_fields": accessed_fields,
149 "validation_time_ms": self.result.metadata.validation_time_ms
150 }
151 });
152
153 (response, !self.result.is_valid)
154 }
155}
156
157#[async_trait::async_trait]
159pub trait CodeModeHandler: Send + Sync {
160 fn server_name(&self) -> &str;
162
163 fn is_enabled(&self) -> bool;
165
166 fn code_format(&self) -> &str;
168
169 async fn validate_code_impl(
171 &self,
172 code: &str,
173 variables: Option<&Value>,
174 dry_run: bool,
175 user_id: &str,
176 session_id: &str,
177 ) -> Result<ValidationResponse, String>;
178
179 async fn execute_code_impl(
181 &self,
182 code: &str,
183 approval_token: &str,
184 variables: Option<&Value>,
185 ) -> Result<Value, String>;
186
187 fn is_policy_configured(&self) -> bool {
193 false
194 }
195
196 fn is_avp_configured(&self) -> bool {
198 self.is_policy_configured()
199 }
200
201 async fn pre_handle_hook(&self) -> Result<Option<(Value, bool)>, String> {
207 Ok(None)
208 }
209
210 fn is_code_mode_tool(&self, name: &str) -> bool {
216 name == "validate_code" || name == "execute_code"
217 }
218
219 fn get_tools(&self) -> Vec<ToolInfo> {
221 if !self.is_enabled() {
222 return vec![];
223 }
224
225 CodeModeToolBuilder::new(self.code_format()).build_tools()
226 }
227
228 async fn handle_tool(
230 &self,
231 name: &str,
232 arguments: Value,
233 user_id: &str,
234 session_id: &str,
235 ) -> Result<(Value, bool), String> {
236 if !self.is_policy_configured() {
238 return Ok((
239 json!({
240 "error": "Code Mode requires a policy evaluator to be configured. \
241 Configure AVP, local Cedar, or another policy backend.",
242 "valid": false
243 }),
244 true,
245 ));
246 }
247
248 if let Some(response) = self.pre_handle_hook().await? {
250 return Ok(response);
251 }
252
253 match name {
254 "validate_code" => {
255 self.handle_validate_code(arguments, user_id, session_id)
256 .await
257 },
258 "execute_code" => self.handle_execute_code(arguments).await,
259 _ => Err(format!("Unknown Code Mode tool: {}", name)),
260 }
261 }
262
263 async fn handle_validate_code(
265 &self,
266 arguments: Value,
267 user_id: &str,
268 session_id: &str,
269 ) -> Result<(Value, bool), String> {
270 let mut input: ValidateCodeInput =
271 serde_json::from_value(arguments).map_err(|e| format!("Invalid arguments: {}", e))?;
272
273 input.code = input.code.trim().to_string();
274
275 let response = self
276 .validate_code_impl(
277 &input.code,
278 input.variables.as_ref(),
279 input.dry_run.unwrap_or(false),
280 user_id,
281 session_id,
282 )
283 .await?;
284
285 Ok(response.to_json_response())
286 }
287
288 async fn handle_execute_code(&self, arguments: Value) -> Result<(Value, bool), String> {
290 let mut input: ExecuteCodeInput =
291 serde_json::from_value(arguments).map_err(|e| format!("Invalid arguments: {}", e))?;
292
293 input.code = input.code.trim().to_string();
294
295 let result = self
296 .execute_code_impl(&input.code, &input.approval_token, input.variables.as_ref())
297 .await?;
298
299 Ok((result, false))
300 }
301}
302
303#[derive(Debug, Deserialize)]
305pub struct ValidateCodeInput {
306 pub code: String,
307 #[serde(default)]
308 pub variables: Option<Value>,
309 #[serde(default)]
310 pub format: Option<String>,
311 #[serde(default)]
312 pub dry_run: Option<bool>,
313}
314
315#[derive(Debug, Deserialize)]
317pub struct ExecuteCodeInput {
318 pub code: String,
319 pub approval_token: String,
320 #[serde(default)]
321 pub variables: Option<Value>,
322}
323
324pub struct CodeModeToolBuilder {
326 code_format: String,
327}
328
329impl CodeModeToolBuilder {
330 pub fn new(code_format: &str) -> Self {
332 Self {
333 code_format: code_format.to_string(),
334 }
335 }
336
337 pub fn build_tools(&self) -> Vec<ToolInfo> {
339 vec![self.build_validate_tool(), self.build_execute_tool()]
340 }
341
342 fn safe_read_annotations() -> ToolAnnotations {
357 ToolAnnotations::new()
358 .with_read_only(true)
359 .with_destructive(false)
360 .with_open_world(false)
361 .with_idempotent(true)
362 }
363
364 pub fn build_validate_tool(&self) -> ToolInfo {
365 ToolInfo::with_annotations(
366 "validate_code",
367 Some(
368 "Validates code and returns a business-language explanation with an approval token. \
369 The code is analyzed for security, complexity, and data access patterns. \
370 You MUST call this before execute_code."
371 .to_string(),
372 ),
373 json!({
374 "type": "object",
375 "properties": {
376 "code": {
377 "type": "string",
378 "description": "The code to validate"
379 },
380 "variables": {
381 "type": "object",
382 "description": "Optional variables for the query"
383 },
384 "format": {
385 "type": "string",
386 "enum": [&self.code_format],
387 "description": format!("Code format. Defaults to '{}' for this server.", self.code_format)
388 },
389 "dry_run": {
390 "type": "boolean",
391 "description": "If true, validate without generating approval token"
392 }
393 },
394 "required": ["code"]
395 }),
396 Self::safe_read_annotations(),
397 )
398 }
399
400 pub fn build_execute_tool(&self) -> ToolInfo {
417 ToolInfo::new(
418 "execute_code",
419 Some(
420 "Executes validated code using an approval token. \
421 The token must be obtained from validate_code and the code must match exactly."
422 .into(),
423 ),
424 json!({
425 "type": "object",
426 "properties": {
427 "code": {
428 "type": "string",
429 "description": "The code to execute (must match validated code)"
430 },
431 "approval_token": {
432 "type": "string",
433 "description": "The approval token from validate_code"
434 },
435 "variables": {
436 "type": "object",
437 "description": "Optional variables for the query"
438 }
439 },
440 "required": ["code", "approval_token"]
441 }),
442 )
443 }
444}
445
446pub fn format_error_response(error: &str) -> (Value, bool) {
448 (
449 json!({
450 "error": error,
451 "valid": false
452 }),
453 true,
454 )
455}
456
457pub fn format_execution_error(error: &str) -> (Value, bool) {
459 (
460 json!({
461 "error": error
462 }),
463 true,
464 )
465}
466
467#[cfg(test)]
468mod tests {
469 use super::*;
470
471 #[test]
472 fn test_validation_response_to_json() {
473 let response = ValidationResponse::success(
474 "Test explanation".into(),
475 RiskLevel::Low,
476 "token123".into(),
477 ValidationMetadata::default(),
478 )
479 .with_action(UnifiedAction::Read)
480 .with_auto_approved(true);
481
482 let (json, is_error) = response.to_json_response();
483
484 assert!(!is_error);
485 assert_eq!(json["valid"], true);
486 assert_eq!(json["explanation"], "Test explanation");
487 assert_eq!(json["risk_level"], "LOW");
488 assert_eq!(json["approval_token"], "token123");
489 assert_eq!(json["action"], "Read");
490 assert_eq!(json["auto_approved"], true);
491 }
492
493 #[test]
494 fn test_validation_response_failure() {
495 let violations = vec![PolicyViolation::new("policy", "rule", "message")];
496 let response = ValidationResponse::failure(violations, ValidationMetadata::default());
497
498 let (json, is_error) = response.to_json_response();
499
500 assert!(is_error);
501 assert_eq!(json["valid"], false);
502 }
503
504 #[test]
507 fn test_rejected_response_does_not_echo_the_code() {
508 let metadata = ValidationMetadata {
509 accessed_types: vec!["/secret/path?q=synthetic".into()],
510 accessed_fields: vec!["GET".into()],
511 ..ValidationMetadata::default()
512 };
513 let mut response = ValidationResponse::failure(
514 vec![PolicyViolation::new("policy", "rule", "message")],
515 metadata,
516 );
517 response.result.explanation = "API calls: Get /secret/path?q=synthetic".into();
518
519 let (json, is_error) = response.to_json_response();
520 assert!(is_error);
521 assert!(
522 !json.to_string().contains("secret/path"),
523 "a rejection must not echo the path: {json}"
524 );
525 assert_eq!(json["explanation"], "");
526 assert_eq!(json["metadata"]["accessed_types"], serde_json::json!([]));
527 assert_eq!(json["metadata"]["accessed_fields"], serde_json::json!([]));
528 assert_eq!(json["violations"][0]["rule"], "rule", "violations survive");
529 }
530
531 #[test]
532 fn test_tool_builder() {
533 let builder = CodeModeToolBuilder::new("graphql");
534 let tools = builder.build_tools();
535
536 assert_eq!(tools.len(), 2);
537 assert_eq!(tools[0].name, "validate_code");
538 assert_eq!(tools[1].name, "execute_code");
539 }
540}