use std::collections::hash_map::DefaultHasher;
use std::collections::{BTreeMap, HashMap};
use std::hash::{Hash, Hasher};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::{Duration, Instant};
use super::client::{self, ClientService};
use super::config::ResolvedTransport;
const IDLE_TTL: Duration = Duration::from_mins(5);
struct Entry {
service: Arc<ClientService>,
last_used: Instant,
}
fn pool() -> &'static Mutex<HashMap<u64, Entry>> {
static POOL: OnceLock<Mutex<HashMap<u64, Entry>>> = OnceLock::new();
POOL.get_or_init(|| Mutex::new(HashMap::new()))
}
#[must_use]
pub fn key(transport: &ResolvedTransport) -> u64 {
let mut h = DefaultHasher::new();
pool_identity(transport).hash(&mut h);
h.finish()
}
#[must_use]
fn pool_identity(transport: &ResolvedTransport) -> String {
match transport {
ResolvedTransport::Stdio {
command,
args,
env,
binary_sha256,
capabilities,
} => format!(
"stdio|command={command:?}|args={args:?}|env={env:?}|binary_sha256={binary_sha256:?}|capabilities={capabilities:?}"
),
ResolvedTransport::Http {
url,
headers,
secret_fingerprints,
} => format!(
"http|url={url:?}|headers={:?}|secrets={:?}",
public_values_case_insensitive(headers, secret_fingerprints),
normalized_secret_fingerprints(secret_fingerprints)
),
}
}
fn public_values_case_insensitive(
values: &BTreeMap<String, String>,
secret_fingerprints: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
values
.iter()
.filter(|(name, _)| {
!secret_fingerprints
.keys()
.any(|secret| secret.eq_ignore_ascii_case(name))
})
.map(|(name, value)| (name.clone(), value.clone()))
.collect()
}
fn normalized_secret_fingerprints(
secret_fingerprints: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
secret_fingerprints
.iter()
.map(|(name, fingerprint)| (name.to_ascii_lowercase(), fingerprint.clone()))
.collect()
}
pub async fn acquire(
transport: &ResolvedTransport,
timeout: Duration,
) -> Result<Arc<ClientService>, String> {
let k = key(transport);
{
let mut map = lock();
let now = Instant::now();
map.retain(|_, e| now.duration_since(e.last_used) < IDLE_TTL && !e.service.is_closed());
if let Some(entry) = map.get_mut(&k) {
entry.last_used = now;
return Ok(entry.service.clone());
}
}
let service = Arc::new(client::open(transport, timeout).await?);
let mut map = lock();
map.insert(
k,
Entry {
service: service.clone(),
last_used: Instant::now(),
},
);
Ok(service)
}
pub fn evict(key: u64) {
lock().remove(&key);
}
pub fn clear() {
lock().clear();
}
#[must_use]
pub fn len() -> usize {
lock().len()
}
fn lock() -> std::sync::MutexGuard<'static, HashMap<u64, Entry>> {
pool()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::BTreeMap;
fn stdio(cmd: &str) -> ResolvedTransport {
ResolvedTransport::Stdio {
command: cmd.into(),
args: vec![],
env: BTreeMap::new(),
binary_sha256: String::new(),
capabilities: None,
}
}
#[test]
fn key_is_stable_and_wiring_sensitive() {
assert_eq!(key(&stdio("a")), key(&stdio("a")), "same wiring → same key");
assert_ne!(
key(&stdio("a")),
key(&stdio("b")),
"different command → different key"
);
}
#[test]
fn http_key_identity_uses_secret_fingerprint_not_raw_value() {
let mut headers = BTreeMap::new();
headers.insert("Authorization".into(), "Bearer raw-token-one".into());
let mut secrets = BTreeMap::new();
secrets.insert("Authorization".into(), "fp-one".into());
let first = ResolvedTransport::Http {
url: "https://gitlab.example/mcp".into(),
headers: headers.clone(),
secret_fingerprints: secrets.clone(),
};
assert!(!pool_identity(&first).contains("raw-token-one"));
headers.insert("Authorization".into(), "Bearer raw-token-two".into());
let same_fingerprint = ResolvedTransport::Http {
url: "https://gitlab.example/mcp".into(),
headers,
secret_fingerprints: secrets.clone(),
};
assert_eq!(key(&first), key(&same_fingerprint));
secrets.insert("Authorization".into(), "fp-two".into());
let rotated = ResolvedTransport::Http {
url: "https://gitlab.example/mcp".into(),
headers: BTreeMap::from([("Authorization".into(), "Bearer raw-token-two".into())]),
secret_fingerprints: secrets,
};
assert_ne!(key(&first), key(&rotated));
}
#[test]
fn http_key_treats_secret_header_names_case_insensitively() {
let upper = ResolvedTransport::Http {
url: "https://example.com/mcp".into(),
headers: BTreeMap::from([("Authorization".into(), "private-token".into())]),
secret_fingerprints: BTreeMap::from([("Authorization".into(), "fingerprint".into())]),
};
let lower_fingerprint = ResolvedTransport::Http {
url: "https://example.com/mcp".into(),
headers: BTreeMap::from([("Authorization".into(), "private-token".into())]),
secret_fingerprints: BTreeMap::from([("authorization".into(), "fingerprint".into())]),
};
assert!(!pool_identity(&lower_fingerprint).contains("private-token"));
assert_eq!(key(&upper), key(&lower_fingerprint));
}
#[test]
fn key_changes_when_secret_format_changes() {
use super::super::config::{GatewayServer, SecretMementoRef, TransportKind};
use super::super::memento::SecretMementoStore;
let store = SecretMementoStore::global();
store.put("mcp/test/pool-format", "token");
let mut server = GatewayServer {
name: "remote".into(),
transport: TransportKind::Http,
url: "https://example.com/mcp".into(),
secret_headers: BTreeMap::from([(
"Authorization".into(),
SecretMementoRef {
id: "mcp/test/pool-format".into(),
format: String::new(),
},
)]),
..Default::default()
};
let raw = server.resolve().expect("resolve raw secret");
server
.secret_headers
.get_mut("Authorization")
.expect("secret header")
.format = "Bearer {secret}".into();
let bearer = server.resolve().expect("resolve formatted secret");
store.remove("mcp/test/pool-format");
assert_ne!(key(&raw), key(&bearer));
}
#[test]
fn evict_and_clear_are_safe_when_empty() {
clear();
evict(key(&stdio("never-pooled")));
clear();
assert_eq!(len(), 0);
}
}