remem-ai 0.6.7

Local-first coding agent memory for Claude Code and OpenAI Codex
Documentation
use brush_parser::word::{
    BraceExpressionMember, BraceExpressionOrText, WordPiece, WordPieceWithSource,
};

use super::{DYNAMIC_SHELL_WORD, MAX_STATIC_WORD_VARIANTS};
use crate::rules::evaluator::git_push_arg_changes_force_state;

const CRITICAL_STATIC_TOKENS: &[&str] = &[
    "git",
    "push",
    "commit",
    "--force",
    "-f",
    "--mirror",
    "--trailer",
    "command",
    "env",
    "exec",
    "eval",
    "bash",
    "dash",
    "ksh",
    "sh",
    "zsh",
];

pub(super) fn append_word_variants(segments: &mut [Vec<String>], variants: Vec<String>) {
    for segment in segments {
        if variants.is_empty() {
            segment.push(DYNAMIC_SHELL_WORD.to_string());
        } else {
            segment.extend(variants.iter().cloned());
        }
    }
}

pub(super) fn critical_brace_variants(pieces: &[BraceExpressionOrText]) -> Vec<String> {
    let mut variants = security_brace_variants(pieces)
        .into_iter()
        .filter(|value| is_critical_static_token(value))
        .collect::<Vec<_>>();
    let mut seen = std::collections::HashSet::new();
    variants.reverse();
    variants.retain(|variant| seen.insert(variant.clone()));
    variants.reverse();
    variants
}

fn is_critical_static_token(value: &str) -> bool {
    CRITICAL_STATIC_TOKENS.contains(&value)
        || git_push_arg_changes_force_state(value)
        || value.starts_with("--trailer=")
        || value.starts_with('+') && value.len() > 1
}

fn security_brace_variants(pieces: &[BraceExpressionOrText]) -> Vec<String> {
    let mut variants = vec![String::new()];
    for piece in pieces {
        let suffixes = match piece {
            BraceExpressionOrText::Text(text) => vec![text.clone()],
            BraceExpressionOrText::Expr(expression) => {
                let mut suffixes = Vec::new();
                for member in expression {
                    match member {
                        BraceExpressionMember::Child(child) => {
                            suffixes.extend(security_brace_variants(child));
                        }
                        BraceExpressionMember::NumberSequence { start, end, .. } => {
                            suffixes.push(start.to_string());
                            suffixes.push(end.to_string());
                        }
                        BraceExpressionMember::CharSequence {
                            start,
                            end,
                            increment,
                        } => {
                            suffixes.push(start.to_string());
                            suffixes.push(end.to_string());
                            for candidate in ['-', '+', ':', 'f', 'm', 'i', 'o'] {
                                if sequence_contains(
                                    *start as i64,
                                    *end as i64,
                                    *increment,
                                    candidate as i64,
                                ) {
                                    suffixes.push(candidate.to_string());
                                }
                            }
                        }
                    }
                    summarize_security_variants(&mut suffixes);
                }
                suffixes
            }
        };
        let prefixes = std::mem::take(&mut variants);
        for prefix in prefixes {
            for suffix in &suffixes {
                variants.push(format!("{prefix}{suffix}"));
            }
        }
        summarize_security_variants(&mut variants);
    }
    variants
}

fn summarize_security_variants(variants: &mut Vec<String>) {
    let mut seen = std::collections::HashSet::new();
    variants.reverse();
    variants.retain(|variant| seen.insert(variant.clone()));
    variants.reverse();
    if variants.len() <= MAX_STATIC_WORD_VARIANTS {
        return;
    }
    let mut retained = (0..variants.len()).collect::<Vec<_>>();
    retained.sort_unstable_by(|left, right| {
        security_variant_score(&variants[*right])
            .cmp(&security_variant_score(&variants[*left]))
            .then_with(|| right.cmp(left))
    });
    retained.truncate(MAX_STATIC_WORD_VARIANTS);
    retained.sort_unstable();
    *variants = retained
        .into_iter()
        .map(|index| variants[index].clone())
        .collect();
}

fn security_variant_score(value: &str) -> usize {
    if is_critical_static_token(value) {
        return 100;
    }
    let before_force = value.split_once('f').map_or(value, |(prefix, _)| prefix);
    let bare = value.trim_start_matches('-');
    if value.contains('f') && !before_force.contains('o') {
        return 50;
    }
    if !bare.is_empty() && ("force".ends_with(bare) || "mirror".ends_with(bare)) {
        return 50;
    }
    if value.starts_with('+') || "mirror".starts_with(bare) {
        return 40;
    }
    usize::from(value.chars().any(|ch| matches!(ch, '-' | '+' | 'f' | 'm')))
}

fn sequence_contains(start: i64, end: i64, increment: i64, value: i64) -> bool {
    if increment == 0
        || value < start.min(end)
        || value > start.max(end)
        || start < end && increment < 0
        || start > end && increment > 0
    {
        return false;
    }
    value
        .checked_sub(start)
        .is_some_and(|distance| distance % increment == 0)
}

pub(super) enum StaticExpansionError {
    Limit,
    Invalid(String),
}

pub(super) fn expand_brace_pieces(
    pieces: &[BraceExpressionOrText],
) -> Result<Vec<String>, StaticExpansionError> {
    let mut variants = vec![String::new()];
    for piece in pieces {
        let suffixes = match piece {
            BraceExpressionOrText::Text(text) => vec![text.clone()],
            BraceExpressionOrText::Expr(expression) => expand_brace_expression(expression)?,
        };
        append_text_variants(&mut variants, &suffixes)?;
    }
    Ok(variants)
}

fn expand_brace_expression(
    expression: &[BraceExpressionMember],
) -> Result<Vec<String>, StaticExpansionError> {
    let mut variants = Vec::new();
    for member in expression {
        match member {
            BraceExpressionMember::Child(pieces) => variants.extend(expand_brace_pieces(pieces)?),
            BraceExpressionMember::NumberSequence {
                start,
                end,
                increment,
            } => {
                let values = inclusive_i64_sequence(*start, *end, *increment)?;
                variants.extend(values.into_iter().map(|value| value.to_string()));
            }
            BraceExpressionMember::CharSequence {
                start,
                end,
                increment,
            } => {
                let values = inclusive_i64_sequence(*start as i64, *end as i64, *increment)?;
                for value in values {
                    let value = u32::try_from(value)
                        .ok()
                        .and_then(char::from_u32)
                        .ok_or_else(|| {
                            StaticExpansionError::Invalid(
                                "Bash brace expansion produced an invalid character".to_string(),
                            )
                        })?;
                    variants.push(value.to_string());
                }
            }
        }
        if variants.len() > MAX_STATIC_WORD_VARIANTS {
            return Err(StaticExpansionError::Limit);
        }
    }
    Ok(variants)
}

fn inclusive_i64_sequence(
    start: i64,
    end: i64,
    increment: i64,
) -> Result<Vec<i64>, StaticExpansionError> {
    if increment == 0 || (start < end && increment < 0) || (start > end && increment > 0) {
        return Err(StaticExpansionError::Invalid(
            "Bash brace expansion has an invalid sequence increment".to_string(),
        ));
    }
    let mut values = Vec::new();
    let mut value = start;
    while if increment > 0 {
        value <= end
    } else {
        value >= end
    } {
        if values.len() == MAX_STATIC_WORD_VARIANTS {
            return Err(StaticExpansionError::Limit);
        }
        values.push(value);
        let Some(next) = value.checked_add(increment) else {
            break;
        };
        value = next;
    }
    Ok(values)
}

fn append_text_variants(
    variants: &mut Vec<String>,
    suffixes: &[String],
) -> Result<(), StaticExpansionError> {
    if suffixes.is_empty()
        || suffixes.len() > MAX_STATIC_WORD_VARIANTS
        || variants.len().saturating_mul(suffixes.len()) > MAX_STATIC_WORD_VARIANTS
    {
        return Err(StaticExpansionError::Limit);
    }
    let prefixes = std::mem::take(variants);
    for prefix in prefixes {
        for suffix in suffixes {
            let mut expanded = prefix.clone();
            expanded.push_str(suffix);
            variants.push(expanded);
        }
    }
    Ok(())
}

pub(super) fn static_word_pieces(pieces: &[WordPieceWithSource]) -> Option<String> {
    let mut value = String::new();
    for piece in pieces {
        match &piece.piece {
            WordPiece::Text(text) | WordPiece::SingleQuotedText(text) => value.push_str(text),
            WordPiece::EscapeSequence(text) => {
                let escaped = text.strip_prefix('\\')?;
                if escaped != "\n" {
                    value.push_str(escaped);
                }
            }
            WordPiece::DoubleQuotedSequence(pieces)
            | WordPiece::GettextDoubleQuotedSequence(pieces) => {
                value.push_str(&static_word_pieces(pieces)?);
            }
            WordPiece::AnsiCQuotedText(text) => {
                value.push_str(&decode_ansi_c_quoted_text(text)?);
            }
            WordPiece::TildeExpansion(_)
            | WordPiece::ParameterExpansion(_)
            | WordPiece::CommandSubstitution(_)
            | WordPiece::BackquotedCommandSubstitution(_)
            | WordPiece::ArithmeticExpression(_) => return None,
        }
    }
    Some(value)
}

fn decode_ansi_c_quoted_text(text: &str) -> Option<String> {
    let mut bytes = Vec::with_capacity(text.len());
    let mut chars = text.chars().peekable();
    while let Some(ch) = chars.next() {
        if ch != '\\' {
            push_char_bytes(&mut bytes, ch);
            continue;
        }
        let escaped = chars.next()?;
        match escaped {
            'a' => bytes.push(0x07),
            'b' => bytes.push(0x08),
            'e' | 'E' => bytes.push(0x1b),
            'f' => bytes.push(0x0c),
            'n' => bytes.push(b'\n'),
            'r' => bytes.push(b'\r'),
            't' => bytes.push(b'\t'),
            'v' => bytes.push(0x0b),
            '\\' => bytes.push(b'\\'),
            '\'' => bytes.push(b'\''),
            'c' => {
                let control = chars.next()?;
                if !control.is_ascii() {
                    return None;
                }
                let control = control.to_ascii_uppercase() as u8;
                bytes.push(if control == b'?' {
                    0x7f
                } else {
                    control & 0x1f
                });
            }
            'x' => bytes.push(take_digits(&mut chars, 16, 2)? as u8),
            'u' => {
                let decoded = char::from_u32(take_digits(&mut chars, 16, 4)?)?;
                push_char_bytes(&mut bytes, decoded);
            }
            'U' => {
                let decoded = char::from_u32(take_digits(&mut chars, 16, 8)?)?;
                push_char_bytes(&mut bytes, decoded);
            }
            '0' => bytes.push(take_digits(&mut chars, 8, 3).unwrap_or(0) as u8),
            '1'..='7' => {
                let mut value = escaped.to_digit(8)?;
                for _ in 0..2 {
                    let Some(digit) = chars.peek().and_then(|ch| ch.to_digit(8)) else {
                        break;
                    };
                    chars.next();
                    value = value * 8 + digit;
                }
                bytes.push(value as u8);
            }
            _ => {
                bytes.push(b'\\');
                push_char_bytes(&mut bytes, escaped);
            }
        }
    }
    if bytes.contains(&0) {
        return None;
    }
    String::from_utf8(bytes).ok()
}

fn take_digits<I>(chars: &mut std::iter::Peekable<I>, radix: u32, max: usize) -> Option<u32>
where
    I: Iterator<Item = char>,
{
    let mut value = 0;
    let mut count = 0;
    while count < max {
        let Some(digit) = chars.peek().and_then(|ch| ch.to_digit(radix)) else {
            break;
        };
        chars.next();
        value = value * radix + digit;
        count += 1;
    }
    (count > 0).then_some(value)
}

fn push_char_bytes(bytes: &mut Vec<u8>, ch: char) {
    let mut encoded = [0; 4];
    bytes.extend_from_slice(ch.encode_utf8(&mut encoded).as_bytes());
}