use super::limits::{MAX_HEX_INPUT_LEN, MIN_HEX_CANDIDATE_LEN};
use super::pipeline::{with_extracted_value_spans, DecodedReplacementBatcher, ExtractedValue};
use super::{DecodeAdmissionSketch, DecodeOutputSink, Decoder, EncodedString};
use keyhog_core::Chunk;
use zeroize::{Zeroize, Zeroizing};
pub(super) struct HexDecoder;
impl Decoder for HexDecoder {
fn name(&self) -> &'static str {
"hex"
}
fn admission_sketch(&self, chunk: &Chunk) -> DecodeAdmissionSketch {
with_extracted_value_spans(&chunk.data, |candidates| {
let mut count = 0usize;
let mut bytes = 0usize;
for candidate in candidates
.iter()
.filter(|candidate| is_hex_candidate(candidate, MIN_HEX_CANDIDATE_LEN))
{
count = count.saturating_add(1);
bytes = bytes.saturating_add(candidate.value.len());
}
if count == 0 {
DecodeAdmissionSketch::NONE
} else {
DecodeAdmissionSketch::possible(DecodeAdmissionSketch::HEX, count, bytes)
}
})
}
fn decode_chunk_into(&self, chunk: &Chunk, sink: &mut dyn DecodeOutputSink) {
let mut batch = DecodedReplacementBatcher::new(sink, chunk, self.name());
let mut open = true;
with_extracted_value_spans(&chunk.data, |candidates| {
for candidate in candidates
.iter()
.filter(|candidate| is_hex_candidate(candidate, MIN_HEX_CANDIDATE_LEN))
{
if !open {
break;
}
let Some(text) = try_decode_hex_candidate_to_utf8(&candidate.value) else {
continue;
};
let (start, end) = candidate.span();
open = batch.push(start, end, text);
}
});
if open {
batch.finish();
}
}
}
const HEX_STACK_INPUT_LIMIT: usize = if 256 < MAX_HEX_INPUT_LEN {
256
} else {
MAX_HEX_INPUT_LEN
};
fn hex_decode_to_stack_buf(input: &str, stack_dst: &mut [u8; 128]) -> Result<usize, ()> {
if !input.as_bytes().contains(&b'_') {
if !input.len().is_multiple_of(2) || input.len() > HEX_STACK_INPUT_LIMIT {
return Err(());
}
let len = input.len() / 2;
hex_simd::decode(
input.as_bytes(),
hex_simd::Out::from_slice(&mut stack_dst[..len]),
)
.map_err(|_| ())?;
Ok(len)
} else {
let mut cleaned = Zeroizing::new([0u8; 256]);
let mut len = 0usize;
for &b in input.as_bytes() {
if b != b'_' {
if len >= HEX_STACK_INPUT_LIMIT {
return Err(());
}
cleaned[len] = b;
len += 1;
}
}
if !len.is_multiple_of(2) {
return Err(());
}
let decoded_len = len / 2;
hex_simd::decode(
&cleaned[..len],
hex_simd::Out::from_slice(&mut stack_dst[..decoded_len]),
)
.map_err(|_| ())?;
Ok(decoded_len)
}
}
fn try_decode_hex_candidate_to_utf8(value: &str) -> Option<String> {
if value.len() <= HEX_STACK_INPUT_LIMIT {
let mut stack_dst = Zeroizing::new([0u8; 128]);
let Ok(decoded_len) = hex_decode_to_stack_buf(value, &mut *stack_dst) else {
return None;
};
let Ok(text) = std::str::from_utf8(&stack_dst[..decoded_len]) else {
return None;
};
return Some(text.to_string());
}
let Ok(decoded) = hex_decode(value) else {
return None;
};
match String::from_utf8(decoded) {
Ok(text) => Some(text),
Err(err) => {
let mut bytes = err.into_bytes();
bytes.zeroize();
None
}
}
}
pub fn find_hex_strings(text: &str, min_length: usize) -> Vec<EncodedString> {
find_hex_string_spans(text, min_length)
.into_iter()
.map(|candidate| EncodedString {
value: candidate.value.to_string(),
})
.collect()
}
fn find_hex_string_spans(text: &str, min_length: usize) -> Vec<ExtractedValue> {
let mut results = Vec::new();
with_extracted_value_spans(text, |candidates| {
for candidate in candidates {
if is_hex_candidate(candidate, min_length) {
results.push(candidate.clone());
}
}
});
results
}
fn is_hex_candidate(candidate: &ExtractedValue, min_length: usize) -> bool {
let hex_len = candidate.value.bytes().filter(|byte| *byte != b'_').count();
hex_len >= min_length
&& hex_len.is_multiple_of(2)
&& candidate
.value
.bytes()
.all(|byte| byte == b'_' || byte.is_ascii_hexdigit())
}
#[allow(clippy::result_unit_err)]
pub fn hex_decode(input: &str) -> Result<Vec<u8>, ()> {
if input.len() <= HEX_STACK_INPUT_LIMIT {
let mut stack_dst = Zeroizing::new([0u8; 128]);
let len = hex_decode_to_stack_buf(input, &mut *stack_dst)?;
return Ok(stack_dst[..len].to_vec());
}
if !input.as_bytes().contains(&b'_') {
if !input.len().is_multiple_of(2) || input.len() > MAX_HEX_INPUT_LEN {
return Err(());
}
let decoded_len = input.len() / 2;
let mut out = vec![0u8; decoded_len];
if hex_simd::decode(input.as_bytes(), hex_simd::Out::from_slice(&mut out)).is_err() {
out.zeroize();
return Err(());
}
return Ok(out);
}
let mut cleaned = Zeroizing::new(Vec::with_capacity(input.len().min(MAX_HEX_INPUT_LEN)));
for &b in input.as_bytes() {
if b != b'_' {
if cleaned.len() >= MAX_HEX_INPUT_LEN {
return Err(());
}
cleaned.push(b);
}
}
if !cleaned.len().is_multiple_of(2) {
return Err(());
}
let decoded_len = cleaned.len() / 2;
let mut out = vec![0u8; decoded_len];
if hex_simd::decode(&cleaned, hex_simd::Out::from_slice(&mut out)).is_err() {
out.zeroize();
return Err(());
}
Ok(out)
}
#[cfg(test)]
#[path = "../../tests/unit/decode_hex.rs"]
mod tests;