use std::collections::{HashMap, HashSet};
use std::fmt;
use std::hash::{DefaultHasher, Hash, Hasher};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use async_trait::async_trait;
use serde_json::Value;
use crate::config::ServerConfig;
#[non_exhaustive]
#[derive(Debug, Clone, Copy)]
pub struct OutboundRequest<'a> {
pub tool: &'a str,
pub method: &'a str,
pub path: &'a str,
pub query: &'a [(String, String)],
pub body: Option<&'a Value>,
pub call_id: &'a str,
pub phase: RequestPhase,
}
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum RequestPhase {
#[default]
Execute,
Validate,
}
impl<'a> OutboundRequest<'a> {
#[must_use]
pub fn new(
tool: &'a str,
method: &'a str,
path: &'a str,
query: &'a [(String, String)],
body: Option<&'a Value>,
) -> Self {
Self {
tool,
method,
path,
query,
body,
call_id: "",
phase: RequestPhase::Execute,
}
}
#[must_use]
pub fn with_phase(mut self, phase: RequestPhase) -> Self {
self.phase = phase;
self
}
#[must_use]
pub fn with_call_id(mut self, call_id: &'a str) -> Self {
self.call_id = call_id;
self
}
}
pub(crate) fn next_call_id() -> String {
static PREFIX: OnceLock<u64> = OnceLock::new();
static COUNTER: AtomicU64 = AtomicU64::new(0);
let prefix = *PREFIX.get_or_init(|| {
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_or(0, |d| d.as_nanos() as u64);
nanos ^ (u64::from(std::process::id()) << 32)
});
format!("{prefix:x}-{:x}", COUNTER.fetch_add(1, Ordering::Relaxed))
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PolicyRefusal {
message: String,
}
impl PolicyRefusal {
#[must_use]
pub fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
#[must_use]
pub fn message(&self) -> &str {
&self.message
}
}
impl fmt::Display for PolicyRefusal {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.message)
}
}
impl std::error::Error for PolicyRefusal {}
#[async_trait]
pub trait RequestPolicy: Send + Sync {
async fn check(&self, req: &OutboundRequest<'_>) -> Result<(), PolicyRefusal>;
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ArgumentRefusal {
message: String,
}
impl ArgumentRefusal {
#[must_use]
pub fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
#[must_use]
pub fn message(&self) -> &str {
&self.message
}
}
impl fmt::Display for ArgumentRefusal {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.message)
}
}
impl std::error::Error for ArgumentRefusal {}
pub trait ArgumentValidator: Send + Sync {
fn validate(&self, args: &Value) -> Result<(), ArgumentRefusal>;
}
#[derive(Clone, Default)]
pub(crate) struct ArgumentValidators {
map: HashMap<String, Arc<dyn ArgumentValidator>>,
}
impl ArgumentValidators {
pub fn insert(&mut self, tool: impl Into<String>, validator: Arc<dyn ArgumentValidator>) {
let tool = tool.into();
if self.map.contains_key(&tool) {
tracing::warn!(
target: "pmcp_server_toolkit::policy",
tool = %tool,
"an ArgumentValidator was already registered for this tool — the earlier one is \
REPLACED and will never run"
);
}
self.map.insert(tool, validator);
}
#[must_use]
pub fn get(&self, tool: &str) -> Option<Arc<dyn ArgumentValidator>> {
self.map.get(tool).map(Arc::clone)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.map.is_empty()
}
#[must_use]
pub fn names(&self) -> Vec<&str> {
let mut names: Vec<&str> = self.map.keys().map(String::as_str).collect();
names.sort_unstable();
names
}
}
impl fmt::Debug for ArgumentValidators {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ArgumentValidators")
.field("tools", &self.names())
.finish()
}
}
#[derive(Clone, Default)]
pub struct ToolkitHooks {
policy: Option<Arc<dyn RequestPolicy>>,
validators: ArgumentValidators,
}
impl ToolkitHooks {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_request_policy(mut self, policy: Arc<dyn RequestPolicy>) -> Self {
if self.policy.is_some() {
tracing::warn!(
target: "pmcp_server_toolkit::policy",
"a RequestPolicy was already registered — the earlier one is REPLACED and will \
never run"
);
}
self.policy = Some(policy);
self
}
#[must_use]
pub fn with_argument_validator(
mut self,
tool: impl Into<String>,
validator: Arc<dyn ArgumentValidator>,
) -> Self {
self.validators.insert(tool, validator);
self
}
#[must_use]
pub fn request_policy(&self) -> Option<Arc<dyn RequestPolicy>> {
self.policy.as_ref().map(Arc::clone)
}
#[must_use]
pub fn argument_validator_for(&self, tool: &str) -> Option<Arc<dyn ArgumentValidator>> {
self.validators.get(tool)
}
#[must_use]
pub fn validator_names(&self) -> Vec<&str> {
self.validators.names()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.policy.is_none() && self.validators.is_empty()
}
}
impl fmt::Debug for ToolkitHooks {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ToolkitHooks")
.field("request_policy", &self.policy.is_some())
.field("validators", &self.validators)
.finish()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReportLevel {
Info,
Warn,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReportLine {
pub level: ReportLevel,
pub text: String,
}
#[must_use]
pub fn render_validation_report(config: &ServerConfig, hooks: &ToolkitHooks) -> Vec<ReportLine> {
let report = config.validation_report();
let mut out = Vec::new();
let schema_check = if cfg!(feature = "input-validation") {
if report.enforce_input_schema {
"ON"
} else {
"OFF"
}
} else {
"OFF (feature)"
};
out.push(ReportLine {
level: if schema_check == "ON" {
ReportLevel::Info
} else {
ReportLevel::Warn
},
text: format!(
"input validation: schema_check={schema_check} default_max_length={} \
additional_properties={} strict={} tools={}",
report.default_max_length,
report.additional_properties,
report.strict,
report.tools.len()
),
});
if !cfg!(feature = "input-validation") {
out.push(ReportLine {
level: ReportLevel::Warn,
text: "input validation: the `input-validation` feature is OFF, so NO tool's \
arguments are checked against its declared inputSchema. It is in the \
toolkit's default feature set — an unenforced build is an explicit opt-out."
.to_string(),
});
}
for tool in &report.tools {
let rules = if tool.rules.is_empty() {
"(none declared; only the always-on path-placeholder character floor and \
length cap apply)"
.to_string()
} else {
tool.rules.join("; ")
};
out.push(ReportLine {
level: ReportLevel::Info,
text: format!("input validation: tool '{}' enforces {rules}", tool.tool),
});
}
if report.opt_outs.is_empty() {
out.push(ReportLine {
level: ReportLevel::Info,
text: "input validation: no [server.validation] opt-out is active — every rule \
this config can enforce is enforced."
.to_string(),
});
} else {
for opt_out in &report.opt_outs {
out.push(ReportLine {
level: ReportLevel::Warn,
text: format!("input validation: [server.validation] opt-out ACTIVE — {opt_out}"),
});
}
}
render_hooks_lines(config, hooks, &mut out);
out
}
fn render_hooks_lines(config: &ServerConfig, hooks: &ToolkitHooks, out: &mut Vec<ReportLine>) {
out.push(ReportLine {
level: ReportLevel::Info,
text: format!(
"input validation: E1 RequestPolicy registered={}",
hooks.request_policy().is_some()
),
});
let names = hooks.validator_names();
if names.is_empty() {
out.push(ReportLine {
level: ReportLevel::Info,
text: "input validation: no E2 ArgumentValidator is registered".to_string(),
});
return;
}
out.push(ReportLine {
level: ReportLevel::Info,
text: format!(
"input validation: E2 ArgumentValidator registered for {}",
names.join(", ")
),
});
for name in names {
if !config.tools.iter().any(|t| t.name == name) {
out.push(ReportLine {
level: ReportLevel::Warn,
text: format!(
"input validation: an ArgumentValidator is registered for '{name}', which \
this config declares no [[tools]] entry for — it will never run"
),
});
}
}
}
pub fn emit_validation_report(config: &ServerConfig, hooks: &ToolkitHooks) {
let lines = render_validation_report(config, hooks);
if !claim_report_emission(&config.server.name, &config.server.version, &lines) {
return;
}
for line in lines {
match line.level {
ReportLevel::Info => {
tracing::info!(target: "pmcp_server_toolkit::policy", "{}", line.text);
},
ReportLevel::Warn => {
tracing::warn!(target: "pmcp_server_toolkit::policy", "{}", line.text);
},
}
}
}
fn claim_report_emission(name: &str, version: &str, lines: &[ReportLine]) -> bool {
static EMITTED: OnceLock<Mutex<HashSet<u64>>> = OnceLock::new();
let mut hasher = DefaultHasher::new();
name.hash(&mut hasher);
version.hash(&mut hasher);
for line in lines {
line.text.hash(&mut hasher);
}
let key = hasher.finish();
EMITTED
.get_or_init(|| Mutex::new(HashSet::new()))
.lock()
.map_or(true, |mut seen| seen.insert(key))
}
#[cfg(test)]
mod tests {
use super::{
next_call_id, ArgumentRefusal, ArgumentValidator, ArgumentValidators, OutboundRequest,
PolicyRefusal, RequestPolicy, ToolkitHooks,
};
use serde_json::{json, Value};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
struct Refuse(&'static str);
#[async_trait::async_trait]
impl RequestPolicy for Refuse {
async fn check(&self, _req: &OutboundRequest<'_>) -> Result<(), PolicyRefusal> {
Err(PolicyRefusal::new(self.0))
}
}
struct CountingValidator(Arc<AtomicUsize>);
impl ArgumentValidator for CountingValidator {
fn validate(&self, _args: &Value) -> Result<(), ArgumentRefusal> {
self.0.fetch_add(1, Ordering::SeqCst);
Ok(())
}
}
#[test]
fn outbound_request_constructor_exposes_every_field() {
let query = vec![("q".to_string(), "x".to_string())];
let body = json!({ "note": "n" });
let req = OutboundRequest::new("t", "GET", "https://h/p/1", &query, Some(&body));
assert_eq!(req.tool, "t");
assert_eq!(req.method, "GET");
assert_eq!(req.path, "https://h/p/1");
assert_eq!(req.query.len(), 1);
assert_eq!(req.body, Some(&body));
}
#[test]
fn refusals_display_their_own_message_verbatim() {
assert_eq!(PolicyRefusal::new("nope").to_string(), "nope");
assert_eq!(ArgumentRefusal::new("bad combo").to_string(), "bad combo");
assert_eq!(PolicyRefusal::new("nope").message(), "nope");
assert_eq!(ArgumentRefusal::new("bad combo").message(), "bad combo");
}
#[tokio::test]
async fn a_policy_refusal_carries_the_policy_message() {
let policy = Refuse("blocked by test policy");
let empty: Vec<(String, String)> = Vec::new();
let req = OutboundRequest::new("t", "GET", "https://h/p", &empty, None);
let err = policy.check(&req).await.expect_err("refuses");
assert_eq!(err.message(), "blocked by test policy");
}
#[test]
fn validator_registration_is_last_one_wins() {
let first = Arc::new(AtomicUsize::new(0));
let second = Arc::new(AtomicUsize::new(0));
let mut reg = ArgumentValidators::default();
assert!(reg.is_empty());
reg.insert("t", Arc::new(CountingValidator(Arc::clone(&first))));
reg.insert("t", Arc::new(CountingValidator(Arc::clone(&second))));
assert_eq!(
reg.names(),
vec!["t"],
"the second registration replaced the first"
);
reg.get("t")
.expect("registered")
.validate(&json!({}))
.expect("allows");
assert_eq!(
first.load(Ordering::SeqCst),
0,
"the replaced validator ran"
);
assert_eq!(second.load(Ordering::SeqCst), 1);
}
#[test]
fn default_hooks_register_nothing() {
let hooks = ToolkitHooks::default();
assert!(hooks.is_empty());
assert!(hooks.request_policy().is_none());
assert!(hooks.argument_validator_for("anything").is_none());
assert!(hooks.validator_names().is_empty());
}
#[test]
fn hooks_builder_records_both_kinds() {
let hooks = ToolkitHooks::new()
.with_request_policy(Arc::new(Refuse("x")))
.with_argument_validator(
"b",
Arc::new(CountingValidator(Arc::new(AtomicUsize::new(0)))),
)
.with_argument_validator(
"a",
Arc::new(CountingValidator(Arc::new(AtomicUsize::new(0)))),
);
assert!(!hooks.is_empty());
assert!(hooks.request_policy().is_some());
assert_eq!(hooks.validator_names(), vec!["a", "b"]);
}
#[test]
fn hooks_debug_never_renders_a_policy_body() {
let hooks = ToolkitHooks::new().with_request_policy(Arc::new(Refuse("secret-ish")));
let rendered = format!("{hooks:?}");
assert!(rendered.contains("request_policy: true"));
assert!(!rendered.contains("secret-ish"));
}
#[test]
fn call_id_defaults_to_unattributed_and_the_builder_sets_it() {
let req = OutboundRequest::new("t", "GET", "/p", &[], None);
assert_eq!(req.call_id, "", "new() must not invent an id");
assert_eq!(req.with_call_id("abc").call_id, "abc");
}
#[test]
fn next_call_id_is_non_empty_and_never_repeats() {
let ids: std::collections::HashSet<String> = (0..2000).map(|_| next_call_id()).collect();
assert_eq!(ids.len(), 2000, "every minted id must be distinct");
assert!(ids.iter().all(|id| !id.is_empty()));
}
#[test]
fn a_request_defaults_to_the_execute_phase_and_can_be_marked_validate() {
let req = OutboundRequest::new("t", "GET", "/x", &[], None);
assert_eq!(req.phase, super::RequestPhase::Execute);
assert_eq!(
req.with_phase(super::RequestPhase::Validate).phase,
super::RequestPhase::Validate
);
}
}