use std::io::Write;
use serde::Serialize;
use super::types::{
AppliedCaps, ErrorPayload, RunOutcome, StopReason, Timings, ToolCallRecord, TranscriptEntry,
Usage,
};
use super::HeadlessError;
const SCHEMA_VERSION: u32 = 1;
const TRUNCATION_MARKER_PREFIX: &str = "…[truncated ";
const TRUNCATION_MARKER_SUFFIX: &str = " bytes]";
const REDACTED_PLACEHOLDER: &str = "[REDACTED]";
const SK_KEY_PREFIX: &str = "sk-";
const SK_KEY_MIN_SUFFIX_LEN: usize = 16;
const AKIA_KEY_PREFIX: &str = "AKIA";
const AKIA_KEY_BODY_LEN: usize = 16;
const BEARER_KEYWORD: &str = "Bearer";
const GENERIC_SECRET_RUN_MIN_LEN: usize = 32;
const HTTP_STATUS_TOKEN_LEN: usize = 3;
const HTTP_STATUS_MIN: u16 = 100;
const HTTP_STATUS_MAX: u16 = 599;
const TIMEOUT_MARKER: &str = "timed out";
const CONNECTION_REFUSED_MARKER: &str = "connection refused";
#[derive(Debug, Serialize)]
struct WireOutcome<'a> {
schema_version: u32,
response: &'a Option<String>,
model: &'a str,
provider: &'a str,
usage: &'a Usage,
timings: &'a Timings,
stop_reason: StopReason,
tool_calls: Vec<ToolCallRecord>,
transcript: Vec<TranscriptEntry>,
consult: &'a Option<serde_json::Value>,
applied_caps: &'a AppliedCaps,
error: &'a Option<ErrorPayload>,
}
pub fn truncate_result(s: &str, cap: usize) -> String {
if s.len() <= cap {
return s.to_string();
}
let mut boundary = cap;
while boundary > 0 && !s.is_char_boundary(boundary) {
boundary -= 1;
}
let kept = s.get(..boundary).unwrap_or_default();
let dropped = s.len() - kept.len();
format!("{kept}{TRUNCATION_MARKER_PREFIX}{dropped}{TRUNCATION_MARKER_SUFFIX}")
}
fn truncate_tool_call(tc: &ToolCallRecord, cap: usize) -> ToolCallRecord {
let mut truncated = tc.clone();
truncated.result = truncate_result(&truncated.result, cap);
truncated
}
fn truncate_transcript_entry(entry: &TranscriptEntry, cap: usize) -> TranscriptEntry {
let mut truncated = entry.clone();
truncated.content = truncate_result(&truncated.content, cap);
truncated.tool_calls = truncated
.tool_calls
.map(|calls| calls.iter().map(|tc| truncate_tool_call(tc, cap)).collect());
truncated
}
pub fn write_json(
out: &mut impl Write,
o: &RunOutcome,
tool_result_cap: usize,
) -> Result<(), HeadlessError> {
let wire = WireOutcome {
schema_version: SCHEMA_VERSION,
response: &o.response,
model: &o.model,
provider: &o.provider,
usage: &o.usage,
timings: &o.timings,
stop_reason: o.stop_reason,
tool_calls: o
.tool_calls
.iter()
.map(|tc| truncate_tool_call(tc, tool_result_cap))
.collect(),
transcript: o
.transcript
.iter()
.map(|e| truncate_transcript_entry(e, tool_result_cap))
.collect(),
consult: &o.consult,
applied_caps: &o.applied_caps,
error: &o.error,
};
serde_json::to_writer(out, &wire).map_err(|e| HeadlessError::Io(e.to_string()))
}
pub fn write_text(out: &mut impl Write, err_out: &mut impl Write, o: &RunOutcome) {
if let Some(response) = &o.response {
let _ = out.write_all(response.as_bytes());
}
if o.applied_caps.max_tool_calls_clamped {
let notice = format!(
"applied_caps: max_tool_calls clamped to {}\n",
o.applied_caps.max_tool_calls
);
let _ = err_out.write_all(notice.as_bytes());
}
}
fn classify_http_status(raw: &str) -> Option<u16> {
const HTTP_KEYWORD: &str = "http";
if !raw.to_ascii_lowercase().contains(HTTP_KEYWORD) {
return None;
}
raw.split(|c: char| !c.is_ascii_alphanumeric())
.filter(|tok| tok.len() == HTTP_STATUS_TOKEN_LEN && tok.bytes().all(|b| b.is_ascii_digit()))
.find_map(|tok| tok.parse::<u16>().ok())
.filter(|code| (HTTP_STATUS_MIN..=HTTP_STATUS_MAX).contains(code))
}
pub fn sanitize_error_message(raw: &str) -> String {
if let Some(status) = classify_http_status(raw) {
return format!("provider error: HTTP {status}");
}
let lower = raw.to_ascii_lowercase();
if lower.contains(TIMEOUT_MARKER) {
return "provider error: request timed out".to_string();
}
if lower.contains(CONNECTION_REFUSED_MARKER) {
return "network error: connection refused".to_string();
}
redact_secret_patterns(raw)
}
fn is_key_body_char(c: char) -> bool {
c.is_ascii_alphanumeric() || c == '-'
}
fn is_generic_secret_char(c: char) -> bool {
c.is_ascii_alphanumeric() || matches!(c, '+' | '/' | '=' | '_' | '-')
}
fn match_bearer_token(chars: &[char], i: usize) -> Option<usize> {
let mut consumed = 0usize;
for (offset, kw_char) in BEARER_KEYWORD.chars().enumerate() {
if !chars.get(i + offset)?.eq_ignore_ascii_case(&kw_char) {
return None;
}
consumed += 1;
}
let ws_start = consumed;
let mut j = ws_start;
while matches!(chars.get(i + j), Some(c) if c.is_whitespace()) {
j += 1;
}
if j == ws_start {
return None;
}
let token_start = j;
while matches!(chars.get(i + j), Some(c) if !c.is_whitespace()) {
j += 1;
}
if j == token_start {
return None;
}
Some(j)
}
fn match_sk_key(chars: &[char], i: usize) -> Option<usize> {
let mut consumed = 0usize;
for (offset, p_char) in SK_KEY_PREFIX.chars().enumerate() {
if *chars.get(i + offset)? != p_char {
return None;
}
consumed += 1;
}
let mut run_len = 0usize;
while matches!(chars.get(i + consumed + run_len), Some(c) if is_key_body_char(*c)) {
run_len += 1;
}
(run_len >= SK_KEY_MIN_SUFFIX_LEN).then_some(consumed + run_len)
}
fn match_akia_key(chars: &[char], i: usize) -> Option<usize> {
let mut consumed = 0usize;
for (offset, p_char) in AKIA_KEY_PREFIX.chars().enumerate() {
if *chars.get(i + offset)? != p_char {
return None;
}
consumed += 1;
}
let mut run_len = 0usize;
while run_len < AKIA_KEY_BODY_LEN {
match chars.get(i + consumed + run_len) {
Some(c) if c.is_ascii_uppercase() || c.is_ascii_digit() => run_len += 1,
_ => break,
}
}
(run_len == AKIA_KEY_BODY_LEN).then_some(consumed + run_len)
}
fn match_generic_secret_run(chars: &[char], i: usize) -> Option<usize> {
let mut run_len = 0usize;
while matches!(chars.get(i + run_len), Some(c) if is_generic_secret_char(*c)) {
run_len += 1;
}
(run_len >= GENERIC_SECRET_RUN_MIN_LEN).then_some(run_len)
}
pub fn redact_secret_patterns(raw: &str) -> String {
let chars: Vec<char> = raw.chars().collect();
let mut out = String::with_capacity(raw.len());
let mut i = 0usize;
while i < chars.len() {
if let Some(consumed) = match_bearer_token(&chars, i)
.or_else(|| match_sk_key(&chars, i))
.or_else(|| match_akia_key(&chars, i))
.or_else(|| match_generic_secret_run(&chars, i))
{
out.push_str(REDACTED_PLACEHOLDER);
i += consumed;
continue;
}
if let Some(c) = chars.get(i) {
out.push(*c);
}
i += 1;
}
out
}
#[doc(hidden)]
pub fn fuzz_sanitize_error_entrypoint(data: &[u8]) {
let s = String::from_utf8_lossy(data);
let _ = sanitize_error_message(&s);
let redacted = redact_secret_patterns(&s);
debug_assert_eq!(
redact_secret_patterns(&redacted),
redacted,
"redaction must be idempotent (no key-pattern left un-redacted)"
);
}
#[cfg(test)]
impl RunOutcome {
pub(crate) fn sample() -> Self {
let sample_tool_call = ToolCallRecord {
name: "ls".to_string(),
input: serde_json::json!({"path": "."}),
result: "file1\nfile2".to_string(),
ms: 12,
ok: true,
};
RunOutcome {
response: Some("Hello from magi.".to_string()),
model: "claude-sonnet-4-6".to_string(),
provider: "anthropic".to_string(),
usage: Usage {
input_tokens: 100,
output_tokens: 50,
},
timings: Timings {
total_ms: 1234,
ttfb_ms: Some(200),
per_turn_ms: vec![600, 634],
},
stop_reason: StopReason::Done,
tool_calls: vec![sample_tool_call.clone()],
transcript: vec![
TranscriptEntry {
role: "user".to_string(),
content: "list files".to_string(),
tool_calls: None,
},
TranscriptEntry {
role: "assistant".to_string(),
content: "Hello from magi.".to_string(),
tool_calls: Some(vec![sample_tool_call]),
},
],
consult: None,
applied_caps: AppliedCaps {
max_tool_calls: 15,
max_tool_calls_clamped: false,
timeout_secs: None,
system_override_applied: false,
},
error: None,
}
}
}
#[cfg(test)]
mod tests {
use super::super::limits::TOOL_RESULT_CAP;
use super::*;
#[test]
fn test_write_json_has_schema_version_and_truncates_large_results() {
let mut o = RunOutcome::sample();
o.tool_calls[0].result = "x".repeat(70_000);
let mut buf = Vec::new();
write_json(&mut buf, &o, TOOL_RESULT_CAP).unwrap();
let v: serde_json::Value = serde_json::from_slice(&buf).unwrap();
assert_eq!(v["schema_version"], 1);
let r = v["tool_calls"][0]["result"].as_str().unwrap();
assert!(r.len() <= 64 * 1024 + 32 && r.ends_with("bytes]"));
}
#[test]
fn test_write_json_field_order_matches_contract() {
let o = RunOutcome::sample();
let mut buf = Vec::new();
write_json(&mut buf, &o, TOOL_RESULT_CAP).unwrap();
let text = String::from_utf8(buf).unwrap();
assert!(text.starts_with("{\"schema_version\":1"));
let order = [
"schema_version",
"response",
"model",
"provider",
"usage",
"timings",
"stop_reason",
"tool_calls",
"transcript",
"consult",
"applied_caps",
"error",
];
let mut search_from = 0usize;
for key in order {
let needle = format!("\"{key}\"");
let rest = text.get(search_from..).unwrap();
let pos = rest
.find(&needle)
.unwrap_or_else(|| panic!("missing key `{key}` after byte {search_from}"));
search_from += pos + needle.len();
}
}
#[test]
fn test_write_json_matches_golden_shape() {
let o = RunOutcome::sample();
let mut buf = Vec::new();
write_json(&mut buf, &o, TOOL_RESULT_CAP).unwrap();
let produced: serde_json::Value = serde_json::from_slice(&buf).unwrap();
let golden: serde_json::Value =
serde_json::from_str(include_str!("../../tests/golden/headless_output_v1.json"))
.unwrap();
assert_eq!(produced, golden);
}
#[test]
fn test_write_text_streams_response_without_clamp_notice() {
let o = RunOutcome::sample();
let mut out = Vec::new();
let mut err = Vec::new();
write_text(&mut out, &mut err, &o);
assert_eq!(String::from_utf8(out).unwrap(), "Hello from magi.");
assert!(err.is_empty());
}
#[test]
fn test_write_text_emits_clamp_notice_to_stderr_when_clamped() {
let mut o = RunOutcome::sample();
o.applied_caps.max_tool_calls_clamped = true;
let mut out = Vec::new();
let mut err = Vec::new();
write_text(&mut out, &mut err, &o);
assert_eq!(String::from_utf8(out).unwrap(), "Hello from magi.");
let err_text = String::from_utf8(err).unwrap();
assert!(err_text.starts_with("applied_caps: max_tool_calls clamped"));
}
#[test]
fn test_write_text_with_no_response_writes_nothing_to_out() {
let mut o = RunOutcome::sample();
o.response = None;
let mut out = Vec::new();
let mut err = Vec::new();
write_text(&mut out, &mut err, &o);
assert!(out.is_empty());
assert!(err.is_empty());
}
#[test]
fn test_truncate_result_leaves_short_strings_untouched() {
let s = "short result";
assert_eq!(truncate_result(s, TOOL_RESULT_CAP), s);
}
#[test]
fn test_truncate_result_respects_custom_effective_cap() {
let small_cap = 16usize;
let s = "x".repeat(small_cap + 4);
let truncated = truncate_result(&s, small_cap);
let expected_prefix = format!("{}{TRUNCATION_MARKER_PREFIX}", "x".repeat(small_cap));
assert!(
truncated.starts_with(&expected_prefix),
"a custom (smaller) effective cap must truncate at {small_cap} bytes, \
not the module constant: {truncated}"
);
assert!(truncated.contains("[truncated 4 bytes]"));
}
#[test]
fn test_truncate_result_caps_and_appends_marker() {
let s = "y".repeat(TOOL_RESULT_CAP + 10);
let truncated = truncate_result(&s, TOOL_RESULT_CAP);
assert!(truncated.len() <= TOOL_RESULT_CAP + 32);
assert!(truncated.ends_with("bytes]"));
assert!(truncated.contains("[truncated 10 bytes]"));
}
#[test]
fn test_truncate_result_backs_off_to_char_boundary_on_multibyte_split() {
let prefix = "a".repeat(TOOL_RESULT_CAP - 1);
let s = format!("{prefix}€{}", "b".repeat(50));
let truncated = truncate_result(&s, TOOL_RESULT_CAP);
let marker_pos = truncated.find(TRUNCATION_MARKER_PREFIX).unwrap();
let kept = truncated.get(..marker_pos).unwrap();
assert!(s.starts_with(kept));
assert!(kept.len() < TOOL_RESULT_CAP);
}
#[test]
fn test_error_message_redacts_multiple_key_formats() {
let anthropic_like = format!("sk-ant-{}", "SECRET".repeat(3));
let secrets = [
anthropic_like.as_str(),
"sk-proj-OPENAISECRET",
"Bearer eyJhbGciOiJ...",
"AKIAIOSFODNN7EXAMPLE",
];
for secret in secrets {
let msg = sanitize_error_message(&format!("http 401: {secret} rejected"));
assert!(!msg.contains(secret), "leaked: {secret}");
}
assert_eq!(
sanitize_error_message("plain network timeout"),
"plain network timeout"
);
assert_eq!(
sanitize_error_message("http 401 Unauthorized: sk-ant-xxxxxxxxxxxxxxxx"),
"provider error: HTTP 401"
);
}
#[test]
fn test_sanitize_fence_redacts_each_pattern_when_no_known_class_matches() {
let sk_like = format!("sk-{}", "a".repeat(20));
let out = sanitize_error_message(&format!("upstream rejected key {sk_like} for tenant"));
assert!(!out.contains(&sk_like));
assert!(out.contains(REDACTED_PLACEHOLDER));
let bearer_msg = "auth failed: Bearer abc123DEF456token";
let out2 = sanitize_error_message(bearer_msg);
assert!(!out2.contains("abc123DEF456token"));
let akia_like = format!("AKIA{}", "B".repeat(16));
let out3 = sanitize_error_message(&format!("leaked credential {akia_like} found"));
assert!(!out3.contains(&akia_like));
let generic_run = "Z".repeat(40);
let out4 = sanitize_error_message(&format!("token dump: {generic_run} end"));
assert!(!out4.contains(&generic_run));
}
#[test]
fn test_generic_secret_run_threshold_is_exact_boundary() {
let below = "a".repeat(GENERIC_SECRET_RUN_MIN_LEN - 1);
assert_eq!(sanitize_error_message(&below), below);
let at_threshold = "a".repeat(GENERIC_SECRET_RUN_MIN_LEN);
let out = sanitize_error_message(&at_threshold);
assert!(out.contains(REDACTED_PLACEHOLDER));
assert!(!out.contains(&at_threshold));
}
#[test]
fn test_fuzz_sanitize_error_entrypoint_never_panics_on_arbitrary_input() {
let long_run = "Z".repeat(5_000);
let cases: &[&[u8]] = &[
b"",
b"\xff\xfe\x00\x80",
b"http 401: leaked key rejected",
b"Bearer tokenvaluewithsomelength",
b"AKIAABCDEFGHIJKLMNOP",
b"plain network timeout",
long_run.as_bytes(),
];
for case in cases {
fuzz_sanitize_error_entrypoint(case);
}
}
#[test]
fn test_sanitize_error_message_classifies_timeout_and_connection_refused() {
assert_eq!(
sanitize_error_message("upstream request timed out after 30s"),
"provider error: request timed out"
);
assert_eq!(
sanitize_error_message("io error: connection refused (os error 111)"),
"network error: connection refused"
);
}
#[test]
fn test_sanitize_redacts_lowercase_bearer() {
let token = format!("{}{}", "tok", "EN1234567890abcdef");
let lower = sanitize_error_message(&format!("auth failed: bearer {token}"));
assert!(!lower.contains(&token), "lowercase bearer leaked: {lower}");
assert!(lower.contains(REDACTED_PLACEHOLDER));
let upper = sanitize_error_message(&format!("auth failed: BEARER {token}"));
assert!(!upper.contains(&token), "uppercase BEARER leaked: {upper}");
assert!(upper.contains(REDACTED_PLACEHOLDER));
}
}