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, sanitize, 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());
let body = sanitize(&body);
format!("{header}\n{}", paint(Style::new().dimmed(), &body))
}
fn parse(raw: &str) -> Value {
serde_json::from_str(raw).unwrap_or_else(|_| Value::String(scrub_malformed(raw)))
}
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_auth_schemes(s)),
other => other.clone(),
}
}
fn scrub_malformed(raw: &str) -> String {
let keyed = scrub_json_like_secrets(raw);
let headers = scrub_header_lines(&keyed);
mask_auth_schemes(&headers)
}
fn scrub_json_like_secrets(raw: &str) -> String {
let bytes = raw.as_bytes();
let mut output = String::with_capacity(raw.len());
let mut copied = 0usize;
let mut cursor = 0usize;
while cursor < bytes.len() {
if bytes[cursor] != b'"' {
cursor += 1;
continue;
}
let Some(key_end) = quoted_end(bytes, cursor) else {
break;
};
let mut colon = key_end + 1;
while colon < bytes.len() && bytes[colon].is_ascii_whitespace() {
colon += 1;
}
if bytes.get(colon) != Some(&b':') {
cursor = key_end + 1;
continue;
}
let key = serde_json::from_str::<String>(&raw[cursor..=key_end]).ok();
if !key.as_deref().is_some_and(is_secret_key) {
cursor = key_end + 1;
continue;
}
let mut value_start = colon + 1;
while value_start < bytes.len() && bytes[value_start].is_ascii_whitespace() {
value_start += 1;
}
let value_end = malformed_value_end(bytes, value_start);
if value_end > value_start {
output.push_str(&raw[copied..value_start]);
output.push('"');
output.push_str(REDACTED);
output.push('"');
copied = value_end;
cursor = value_end;
} else {
cursor = value_start.max(key_end + 1);
}
}
if copied == 0 {
raw.to_string()
} else {
output.push_str(&raw[copied..]);
output
}
}
fn quoted_end(bytes: &[u8], start: usize) -> Option<usize> {
let mut escaped = false;
for (offset, byte) in bytes[start + 1..].iter().enumerate() {
if escaped {
escaped = false;
} else if *byte == b'\\' {
escaped = true;
} else if *byte == b'"' {
return Some(start + 1 + offset);
}
}
None
}
fn malformed_value_end(bytes: &[u8], start: usize) -> usize {
let Some(first) = bytes.get(start) else {
return start;
};
match first {
b'"' => quoted_end(bytes, start).map_or(bytes.len(), |end| end + 1),
b'{' | b'[' => balanced_end(bytes, start).unwrap_or(bytes.len()),
_ => {
let mut end = start;
while end < bytes.len() && !matches!(bytes[end], b',' | b'}' | b']' | b'\r' | b'\n') {
end += 1;
}
while end > start && bytes[end - 1].is_ascii_whitespace() {
end -= 1;
}
end
}
}
}
fn balanced_end(bytes: &[u8], start: usize) -> Option<usize> {
let mut stack = vec![bytes[start]];
let mut cursor = start + 1;
while cursor < bytes.len() {
match bytes[cursor] {
b'"' => cursor = quoted_end(bytes, cursor)? + 1,
b'{' | b'[' => {
stack.push(bytes[cursor]);
cursor += 1;
}
b'}' if stack.last() == Some(&b'{') => {
stack.pop();
cursor += 1;
if stack.is_empty() {
return Some(cursor);
}
}
b']' if stack.last() == Some(&b'[') => {
stack.pop();
cursor += 1;
if stack.is_empty() {
return Some(cursor);
}
}
_ => cursor += 1,
}
}
None
}
fn scrub_header_lines(raw: &str) -> String {
let mut output = String::with_capacity(raw.len());
for segment in raw.split_inclusive('\n') {
let line_end = segment.trim_end_matches(['\r', '\n']);
let leading = line_end.len() - line_end.trim_start_matches([' ', '\t']).len();
let line = &line_end[leading..];
let Some(colon) = line.find(':') else {
output.push_str(segment);
continue;
};
let name = line[..colon].trim_end();
let header_shaped = !name.is_empty()
&& name
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_'));
if !header_shaped || !is_secret_key(name) {
output.push_str(segment);
continue;
}
let value_offset = leading + colon + 1;
let whitespace = segment[value_offset..]
.bytes()
.take_while(u8::is_ascii_whitespace)
.take_while(|byte| !matches!(byte, b'\r' | b'\n'))
.count();
let value_start = value_offset + whitespace;
output.push_str(&segment[..value_start]);
output.push_str(REDACTED);
if segment.ends_with("\r\n") {
output.push_str("\r\n");
} else if segment.ends_with('\n') {
output.push('\n');
}
}
output
}
const AUTH_SCHEMES: &[&str] = &["bearer ", "basic ", "digest ", "token "];
fn mask_auth_schemes(s: &str) -> String {
let lowered = s.to_ascii_lowercase();
let mut output = String::with_capacity(s.len());
let mut copied = 0usize;
let mut cursor = 0usize;
while cursor < s.len() {
let Some((scheme_start, scheme)) = next_auth_scheme(lowered.as_bytes(), cursor) else {
break;
};
let credential_start = scheme_start + scheme.len();
let credential_end = if scheme == "digest " {
s[credential_start..]
.find(['\r', '\n'])
.map_or(s.len(), |offset| credential_start + offset)
} else {
s[credential_start..]
.find(|character: char| {
character.is_ascii_whitespace()
|| matches!(character, '"' | '\'' | ',' | ';' | '}' | ']' | ')' | '&')
})
.map_or(s.len(), |offset| credential_start + offset)
};
if credential_end == credential_start {
cursor = credential_start;
continue;
}
output.push_str(&s[copied..credential_start]);
output.push_str(REDACTED);
copied = credential_end;
cursor = credential_end;
}
if copied == 0 {
s.to_string()
} else {
output.push_str(&s[copied..]);
output
}
}
fn next_auth_scheme(lowered: &[u8], start: usize) -> Option<(usize, &'static str)> {
(start..lowered.len()).find_map(|at| {
AUTH_SCHEMES
.iter()
.find(|scheme| lowered[at..].starts_with(scheme.as_bytes()))
.map(|scheme| (at, *scheme))
})
}
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::*;
use crate::property::{GENERATED_CASES, Generator, WIRE_REGRESSIONS};
fn assert_terminal_safe(text: &str) {
assert_eq!(
sanitize(text).as_ref(),
text,
"terminal control survived: {text:?}"
);
}
#[test]
fn property_wire_frames_are_total_terminal_safe_and_never_retain_seeded_secrets() {
let wire = Wire::new(true);
for (raw, secrets) in WIRE_REGRESSIONS {
let frame = wire.record(Direction::Received, raw);
let stored = serde_json::to_string(&frame.json).unwrap();
let rendered = render(Direction::Received, &frame);
for secret in *secrets {
assert!(
!stored.contains(secret),
"stored secret {secret:?}: {stored}"
);
assert!(
!rendered.contains(secret),
"rendered secret {secret:?}: {rendered}"
);
}
assert_terminal_safe(&rendered);
}
let mut generator = Generator::new(0x04);
for case in 0..GENERATED_CASES {
let arbitrary = generator.text(192);
let frame = wire.record(Direction::Received, &arbitrary);
assert_terminal_safe(&render(Direction::Received, &frame));
let secret = format!("fuzz-secret-{case:04x}-{:016x}", generator.next());
let note = serde_json::to_string(&generator.text(48)).unwrap();
let inputs = [
format!(r#"{{"password":"{secret}","note":{note}}}"#),
format!(r#"{{"clientSecret":"{secret}","note":{note}"#),
format!("X-Api-Key: {secret}\nNote: keep"),
format!(r#"{{"message":"Bearer {secret}","note":{note}}}"#),
];
for raw in inputs {
let frame = wire.record(Direction::Received, &raw);
let stored = serde_json::to_string(&frame.json).unwrap();
let rendered = render(Direction::Received, &frame);
assert!(
!stored.contains(&secret),
"stored secret from {raw:?}: {stored}"
);
assert!(
!rendered.contains(&secret),
"rendered secret from {raw:?}: {rendered}"
);
assert_terminal_safe(&rendered);
}
}
}
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"
);
let malformed = format!(
"{{\"password\":\"do-not-retain{}",
"x".repeat(MAX_FRAME_BYTES)
);
let frame = wire.record(Direction::Received, &malformed);
let rendered = render(Direction::Received, &frame);
assert!(frame.json.get("mcp-repl/truncated").is_some());
assert!(!rendered.contains("do-not-retain"), "{rendered}");
assert!(
rendered.len() < 1024,
"malformed oversize frame was retained"
);
}
#[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 malformed_objects_and_arrays_scrub_secret_keys_but_keep_context() {
for raw in [
r#"{"params":{"password":"hunter2","taskToken":"visible"}"#,
r#"[{"api_token":"ghp_one"},{"clientSecret":"stripe_two"},{"note":"keep me"}"#,
] {
let scrubbed = parse(raw);
let scrubbed = scrubbed.as_str().expect("malformed frames stay strings");
for secret in ["hunter2", "ghp_one", "stripe_two"] {
assert!(!scrubbed.contains(secret), "{secret} leaked: {scrubbed}");
}
assert!(scrubbed.contains(REDACTED), "{scrubbed}");
if raw.contains("taskToken") {
assert!(scrubbed.contains("\"taskToken\":\"visible\""), "{scrubbed}");
}
}
let object =
scrub_malformed("{\"password\":\"line one\nline two\",\"note\":\"still visible\"");
assert!(!object.contains("line one"), "{object}");
assert!(!object.contains("line two"), "{object}");
assert!(object.contains("still visible"), "{object}");
}
#[test]
fn malformed_header_lines_scrub_multiple_secrets_and_preserve_other_lines() {
let raw = "Authorization: Bearer first\nX-Api-Key: second\r\nCookie: third\nContent-Type: application/json\n<broken>";
let scrubbed = scrub_malformed(raw);
for secret in ["first", "second", "third"] {
assert!(!scrubbed.contains(secret), "{secret} leaked: {scrubbed}");
}
for context in [
"Authorization:",
"X-Api-Key:",
"Cookie:",
"Content-Type: application/json",
"<broken>",
] {
assert!(
scrubbed.contains(context),
"missing {context:?}: {scrubbed}"
);
}
assert_eq!(scrubbed.matches(REDACTED).count(), 3, "{scrubbed}");
}
#[test]
fn ordinary_malformed_text_is_unchanged() {
let raw = "<html>502 Bad Gateway</html>\nupstream reset {";
assert_eq!(scrub_malformed(raw), raw);
}
#[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_auth_schemes(value);
assert_eq!(masked, format!("{kept}{REDACTED}"), "{value:?}");
}
}
#[test]
fn multiple_auth_schemes_are_masked_without_hiding_context() {
let masked = mask_auth_schemes("Bearer aaa and Basic bbb");
assert!(!masked.contains("aaa"), "{masked}");
assert!(!masked.contains("bbb"), "{masked}");
assert_eq!(masked, "Bearer <redacted> and Basic <redacted>");
}
#[test]
fn a_value_without_a_scheme_is_untouched() {
assert_eq!(mask_auth_schemes("just a sentence"), "just a sentence");
assert_eq!(mask_auth_schemes("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());
}
}