use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Mutex, OnceLock};
use std::time::{Duration, Instant};
use async_trait::async_trait;
use nu_ansi_term::Style;
use serde_json::Value;
use tower_mcp::client::ClientTransport;
use tower_mcp::error::Result;
use crate::style::{paint, tag};
use crate::timing;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Direction {
Sent,
Received,
}
impl Direction {
fn label(self) -> &'static str {
match self {
Direction::Sent => "wire ->",
Direction::Received => "wire <-",
}
}
}
#[derive(Clone, Debug)]
pub struct Frame {
pub json: Value,
pub at: Duration,
pub elapsed: Option<Duration>,
}
pub struct Wire {
trace: AtomicBool,
started: Instant,
state: Mutex<State>,
}
#[derive(Default)]
struct State {
pending: HashMap<String, Instant>,
last_request: Option<Frame>,
last_response: Option<Frame>,
}
const PENDING_CAP: usize = 256;
const MAX_FRAME_BYTES: usize = 1 << 20;
fn summarize(raw_len: usize, json: &Value) -> Value {
let mut summary = serde_json::Map::new();
for key in ["jsonrpc", "id", "method"] {
if let Some(value) = json.get(key) {
summary.insert(key.to_string(), value.clone());
}
}
summary.insert(
"mcp-repl/truncated".to_string(),
Value::String(format!(
"{raw_len} bytes, over the {MAX_FRAME_BYTES} byte cap; body not retained"
)),
);
Value::Object(summary)
}
impl Wire {
pub fn new(trace: bool) -> Self {
Self {
trace: AtomicBool::new(trace),
started: Instant::now(),
state: Mutex::new(State::default()),
}
}
pub fn set_trace(&self, on: bool) {
self.trace.store(on, Ordering::Relaxed);
}
pub fn trace_enabled(&self) -> bool {
self.trace.load(Ordering::Relaxed)
}
pub fn sent(&self, raw: &str) -> Option<String> {
let frame = self.record(Direction::Sent, raw);
self.trace_enabled()
.then(|| render(Direction::Sent, &frame))
}
pub fn received(&self, raw: &str) -> Option<String> {
let frame = self.record(Direction::Received, raw);
self.trace_enabled()
.then(|| render(Direction::Received, &frame))
}
pub fn last_exchange(&self) -> Option<(Frame, Option<Frame>)> {
let state = self.state.lock().unwrap();
let request = state.last_request.clone()?;
Some((request, state.last_response.clone()))
}
fn record(&self, dir: Direction, raw: &str) -> Frame {
let now = Instant::now();
let json = redact(&parse(raw));
let id = frame_id(&json);
let json = if raw.len() > MAX_FRAME_BYTES {
summarize(raw.len(), &json)
} else {
json
};
let has_method = json.get("method").is_some();
let mut state = self.state.lock().unwrap();
let mut elapsed = None;
if dir == Direction::Received
&& !has_method
&& let Some(id) = &id
{
elapsed = state
.pending
.remove(id)
.map(|sent| now.saturating_duration_since(sent));
}
let frame = Frame {
json,
at: now.saturating_duration_since(self.started),
elapsed,
};
match dir {
Direction::Sent => {
if has_method && let Some(id) = id {
if state.pending.len() >= PENDING_CAP {
state.pending.clear();
}
state.pending.insert(id, now);
state.last_request = Some(frame.clone());
state.last_response = None;
}
}
Direction::Received => {
if !has_method
&& id.is_some()
&& state.last_request.as_ref().and_then(|f| frame_id(&f.json)) == id
{
state.last_response = Some(frame.clone());
}
}
}
frame
}
}
static WIRE: OnceLock<Wire> = OnceLock::new();
pub fn init(trace: bool) {
let _ = WIRE.set(Wire::new(trace));
}
pub fn wire() -> &'static Wire {
WIRE.get_or_init(|| Wire::new(false))
}
pub fn render(dir: Direction, frame: &Frame) -> String {
let mut header = format!(
"{} {}",
tag(Style::new().dimmed(), dir.label()),
paint(
Style::new().dimmed(),
&format!("+{:.3}s", frame.at.as_secs_f64())
)
);
if let Some(elapsed) = frame.elapsed {
header.push(' ');
header.push_str(&timing(elapsed));
}
let body = serde_json::to_string_pretty(&frame.json).unwrap_or_else(|_| frame.json.to_string());
format!("{header}\n{}", paint(Style::new().dimmed(), &body))
}
fn parse(raw: &str) -> Value {
serde_json::from_str(raw).unwrap_or_else(|_| Value::String(raw.to_string()))
}
fn frame_id(json: &Value) -> Option<String> {
match json.get("id")? {
Value::Null => None,
Value::String(s) => Some(s.clone()),
other => Some(other.to_string()),
}
}
const REDACTED: &str = "<redacted>";
const SECRET_KEYS: &[&str] = &[
"authorization",
"proxyauthorization",
"wwwauthenticate",
"bearer",
"bearertoken",
"token",
"accesstoken",
"refreshtoken",
"idtoken",
"sessiontoken",
"apitoken",
"authtoken",
"apikey",
"xapikey",
"apisecret",
"accesskey",
"accesskeyid",
"secretaccesskey",
"privatekey",
"secret",
"clientsecret",
"clientassertion",
"assertion",
"password",
"passwd",
"passphrase",
"credential",
"credentials",
"cookie",
"setcookie",
"signature",
];
fn normalize_key(key: &str) -> String {
key.chars()
.filter(|c| c.is_ascii_alphanumeric())
.map(|c| c.to_ascii_lowercase())
.collect()
}
const NOT_SECRETS: &[&str] = &[
"tasktoken",
"progresstoken",
"requesttoken",
"continuationtoken",
"pagetoken",
"nexttoken",
"publickey",
"keys",
"key",
];
const STRONG_ENDINGS: &[&str] = &["token", "secret", "password", "passphrase", "credential"];
const QUALIFIED_ENDINGS: &[&str] = &["key"];
const SECRET_QUALIFIERS: &[&str] = &[
"api",
"auth",
"access",
"private",
"client",
"session",
"signing",
"encryption",
"secret",
];
fn is_secret_key(key: &str) -> bool {
let normalized = normalize_key(key);
if NOT_SECRETS.contains(&normalized.as_str()) {
return false;
}
if SECRET_KEYS.contains(&normalized.as_str()) {
return true;
}
if STRONG_ENDINGS
.iter()
.any(|ending| normalized.ends_with(ending))
{
return true;
}
QUALIFIED_ENDINGS.iter().any(|ending| {
normalized.ends_with(ending)
&& SECRET_QUALIFIERS
.iter()
.any(|qualifier| normalized.contains(qualifier))
})
}
const CREDENTIAL_SHAPES: &[&str] = &[
"token",
"secret",
"password",
"passwd",
"passphrase",
"credential",
"apikey",
"privatekey",
"authorization",
];
pub(crate) fn looks_like_credential(name: &str) -> bool {
let normalized = normalize_key(name);
CREDENTIAL_SHAPES
.iter()
.any(|shape| normalized.contains(shape))
}
fn redact(value: &Value) -> Value {
match value {
Value::Object(map) => Value::Object(
map.iter()
.map(|(key, val)| {
if is_secret_key(key) {
(key.clone(), Value::String(REDACTED.to_string()))
} else {
(key.clone(), redact(val))
}
})
.collect(),
),
Value::Array(items) => Value::Array(items.iter().map(redact).collect()),
Value::String(s) => Value::String(mask_bearer(s)),
other => other.clone(),
}
}
const AUTH_SCHEMES: &[&str] = &["bearer ", "basic ", "digest ", "token "];
fn mask_bearer(s: &str) -> String {
let lowered = s.to_ascii_lowercase();
let earliest = AUTH_SCHEMES
.iter()
.filter_map(|scheme| lowered.find(scheme).map(|at| at + scheme.len()))
.min();
match earliest {
Some(end) => format!("{}{REDACTED}", &s[..end]),
None => s.to_string(),
}
}
pub struct TracingTransport<T> {
inner: T,
wire: &'static Wire,
}
impl<T: ClientTransport> TracingTransport<T> {
pub fn new(inner: T) -> Self {
Self::with_wire(inner, wire())
}
pub fn with_wire(inner: T, wire: &'static Wire) -> Self {
Self { inner, wire }
}
}
#[async_trait]
impl<T: ClientTransport> ClientTransport for TracingTransport<T> {
async fn send(&mut self, message: &str) -> Result<()> {
if let Some(block) = self.wire.sent(message) {
eprintln!("{block}");
}
self.inner.send(message).await
}
async fn recv(&mut self) -> Result<Option<String>> {
let message = self.inner.recv().await?;
if let Some(raw) = &message
&& let Some(block) = self.wire.received(raw)
{
eprintln!("{block}");
}
Ok(message)
}
fn is_connected(&self) -> bool {
self.inner.is_connected()
}
async fn close(&mut self) -> Result<()> {
self.inner.close().await
}
async fn reset_session(&mut self) {
self.inner.reset_session().await;
}
fn supports_session_recovery(&self) -> bool {
self.inner.supports_session_recovery()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn request(id: u32, method: &str) -> String {
serde_json::json!({"jsonrpc": "2.0", "id": id, "method": method, "params": {}}).to_string()
}
fn response(id: u32) -> String {
serde_json::json!({"jsonrpc": "2.0", "id": id, "result": {"ok": true}}).to_string()
}
#[test]
fn an_oversized_frame_is_summarized_rather_than_kept() {
let wire = Wire::new(true);
let body = "x".repeat(4 * MAX_FRAME_BYTES);
let huge =
serde_json::json!({"jsonrpc": "2.0", "id": 1, "result": {"text": body}}).to_string();
wire.sent(&request(1, "resources/read"));
wire.received(&huge);
let (_, response) = wire.last_exchange().unwrap();
let response = response.expect("the response is still paired with its request");
assert_eq!(response.json["id"], 1);
assert!(response.json.get("result").is_none(), "{:?}", response.json);
let note = response.json["mcp-repl/truncated"]
.as_str()
.expect("the truncation is explained rather than silent");
assert!(note.contains(&huge.len().to_string()), "{note}");
assert!(
serde_json::to_string(&response.json).unwrap().len() < 1024,
"the summary is small"
);
}
#[test]
fn a_normal_frame_is_kept_whole() {
let wire = Wire::new(true);
wire.sent(&request(1, "tools/call"));
wire.received(&response(1));
let (_, response) = wire.last_exchange().unwrap();
assert_eq!(response.unwrap().json["result"]["ok"], true);
}
#[test]
fn a_response_is_paired_with_the_request_it_answers() {
let wire = Wire::new(true);
wire.sent(&request(1, "tools/call"));
wire.received(&response(1));
let (req, resp) = wire.last_exchange().expect("an exchange was recorded");
assert_eq!(req.json["method"], "tools/call");
let resp = resp.expect("the response was paired");
assert_eq!(resp.json["result"]["ok"], true);
assert!(
resp.elapsed.is_some(),
"a paired response carries its round-trip time"
);
}
#[test]
fn a_new_request_clears_the_previous_response() {
let wire = Wire::new(false);
wire.sent(&request(1, "tools/list"));
wire.received(&response(1));
wire.sent(&request(2, "tools/call"));
let (req, resp) = wire.last_exchange().unwrap();
assert_eq!(req.json["method"], "tools/call");
assert!(resp.is_none(), "the new request has not been answered yet");
}
#[test]
fn notifications_are_not_exchanges() {
let wire = Wire::new(false);
wire.sent(
&serde_json::json!({"jsonrpc": "2.0", "method": "notifications/initialized"})
.to_string(),
);
assert!(wire.last_exchange().is_none());
}
#[test]
fn a_server_initiated_request_does_not_answer_ours() {
let wire = Wire::new(false);
wire.sent(&request(1, "tools/call"));
wire.received(&request(7, "sampling/createMessage"));
let (_, resp) = wire.last_exchange().unwrap();
assert!(resp.is_none());
}
#[test]
fn a_mismatched_response_is_not_the_last_response() {
let wire = Wire::new(false);
wire.sent(&request(1, "tools/list"));
wire.sent(&request(2, "tools/call"));
wire.received(&response(1));
let (req, resp) = wire.last_exchange().unwrap();
assert_eq!(req.json["id"], 2);
assert!(resp.is_none());
}
#[test]
fn recording_happens_with_tracing_off_but_nothing_renders() {
let wire = Wire::new(false);
assert!(wire.sent(&request(1, "tools/list")).is_none());
assert!(wire.received(&response(1)).is_none());
assert!(
wire.last_exchange().is_some(),
"`last` works without --trace"
);
wire.set_trace(true);
assert!(wire.sent(&request(2, "tools/list")).is_some());
}
#[test]
fn a_rendered_frame_shows_direction_timestamp_and_elapsed() {
let wire = Wire::new(true);
let sent = wire.sent(&request(1, "tools/call")).unwrap();
assert!(sent.contains("wire ->"), "{sent}");
assert!(sent.contains("+0."), "a session-relative timestamp: {sent}");
assert!(sent.contains("tools/call"), "{sent}");
assert!(!sent.contains("elapsed"));
let received = wire.received(&response(1)).unwrap();
assert!(received.contains("wire <-"), "{received}");
assert!(
received.contains("ms]") || received.contains("s]"),
"a response carries its round-trip time: {received}"
);
}
#[test]
fn an_unparseable_frame_still_traces() {
let wire = Wire::new(true);
let rendered = wire.received("<html>502 Bad Gateway</html>").unwrap();
assert!(rendered.contains("502 Bad Gateway"), "{rendered}");
}
#[test]
fn secrets_are_masked_by_key_name() {
let frame = redact(&serde_json::json!({
"params": {
"headers": {"Authorization": "Bearer sk-live-123", "X-Api-Key": "k1"},
"arguments": {"apiKey": "k2", "password": "hunter2", "nested": [{"token": "t"}]},
}
}));
let rendered = frame.to_string();
for secret in ["sk-live-123", "k1", "k2", "hunter2", "\"t\""] {
assert!(!rendered.contains(secret), "{secret} leaked: {rendered}");
}
assert_eq!(frame["params"]["headers"]["Authorization"], REDACTED);
assert_eq!(frame["params"]["arguments"]["nested"][0]["token"], REDACTED);
}
#[test]
fn a_bearer_token_inside_a_string_is_masked() {
let frame = redact(&serde_json::json!({
"error": {"message": "rejected Authorization: Bearer sk-live-123"}
}));
let message = frame["error"]["message"].as_str().unwrap();
assert!(!message.contains("sk-live-123"), "{message}");
assert!(message.starts_with("rejected Authorization: Bearer "));
}
#[test]
fn ordinary_values_are_left_alone() {
let original = serde_json::json!({
"params": {"name": "add", "arguments": {"a": 2, "b": 3, "taskToken": "visible"}},
"flags": [true, null, 1.5],
});
assert_eq!(redact(&original), original);
}
#[test]
fn credential_headers_and_fields_are_masked() {
for key in [
"Cookie",
"Set-Cookie",
"api_token",
"auth_token",
"x-api-token",
"accessKey",
"secretAccessKey",
"AWS_SECRET_ACCESS_KEY",
"private_key",
"client_assertion",
"signature",
"githubToken",
"session_key",
"signingSecret",
"WWW-Authenticate",
] {
let frame = serde_json::json!({ key.to_string(): "s3cret" });
let redacted = redact(&frame);
assert_eq!(redacted[key], REDACTED, "{key} leaked: {redacted}");
}
}
#[test]
fn identifiers_that_merely_end_like_secrets_stay_readable() {
for key in [
"taskToken",
"progressToken",
"nextToken",
"continuationToken",
"pageToken",
"publicKey",
"sortKey",
"partitionKey",
"idempotencyKey",
"name",
"uri",
] {
let frame = serde_json::json!({ key.to_string(): "visible" });
assert_eq!(redact(&frame)[key], "visible", "{key} was over-redacted");
}
}
#[test]
fn every_auth_scheme_is_masked_inside_a_string() {
for (value, kept) in [
("Bearer abc.def.ghi", "Bearer "),
("bearer abc", "bearer "),
("Basic dXNlcjpwYXNz", "Basic "),
("Digest username=\"u\", response=\"r\"", "Digest "),
("token ghp_xxx", "token "),
("Authorization: Bearer abc", "Authorization: Bearer "),
] {
let masked = mask_bearer(value);
assert_eq!(masked, format!("{kept}{REDACTED}"), "{value:?}");
}
}
#[test]
fn the_earliest_scheme_wins_so_nothing_trails_it() {
let masked = mask_bearer("Bearer aaa and Basic bbb");
assert!(!masked.contains("aaa"), "{masked}");
assert!(!masked.contains("bbb"), "{masked}");
}
#[test]
fn a_value_without_a_scheme_is_untouched() {
assert_eq!(mask_bearer("just a sentence"), "just a sentence");
assert_eq!(mask_bearer("bearer"), "bearer");
}
struct FakeTransport {
sent: Vec<String>,
incoming: Vec<String>,
}
#[async_trait]
impl ClientTransport for FakeTransport {
async fn send(&mut self, message: &str) -> Result<()> {
self.sent.push(message.to_string());
Ok(())
}
async fn recv(&mut self) -> Result<Option<String>> {
Ok(if self.incoming.is_empty() {
None
} else {
Some(self.incoming.remove(0))
})
}
fn is_connected(&self) -> bool {
true
}
async fn close(&mut self) -> Result<()> {
Ok(())
}
fn supports_session_recovery(&self) -> bool {
true
}
}
#[tokio::test]
async fn the_wrapper_records_both_directions_and_delegates() {
let wire: &'static Wire = Box::leak(Box::new(Wire::new(false)));
let mut transport = TracingTransport::with_wire(
FakeTransport {
sent: Vec::new(),
incoming: vec![response(1)],
},
wire,
);
transport.send(&request(1, "tools/call")).await.unwrap();
let received = transport.recv().await.unwrap();
assert_eq!(received.as_deref(), Some(response(1).as_str()));
assert_eq!(
transport.inner.sent.len(),
1,
"the frame reached the inner transport"
);
assert!(
transport.supports_session_recovery(),
"the wrapper must not change how the client handles sessions"
);
let (req, resp) = wire.last_exchange().unwrap();
assert_eq!(req.json["method"], "tools/call");
assert!(resp.is_some());
}
}