use crate::error::Error;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use std::fmt;
const MAX_PATTERN_BYTES: usize = 1024;
const MAX_SEGMENTS: usize = 64;
const MAX_ALTERNATIVES: usize = 256;
#[derive(Debug, Clone, PartialEq, Eq)]
enum Token {
Star,
Any,
Literal(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum Segment {
DoubleStar,
Literal(String),
Wild(Vec<Token>),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct KeyPattern {
source: String,
alternatives: Vec<Vec<Segment>>,
}
impl KeyPattern {
pub fn whole_kb() -> Self {
Self::parse("**").expect("`**` is a valid pattern")
}
pub fn is_whole_kb(&self) -> bool {
self.source == "**"
}
pub fn parse(input: &str) -> Result<Self, Error> {
if input.is_empty() {
return Err(invalid("pattern must not be empty"));
}
if input.len() > MAX_PATTERN_BYTES {
return Err(invalid(format!(
"pattern exceeds {MAX_PATTERN_BYTES} bytes"
)));
}
if input.starts_with('/') {
return Err(invalid(
"pattern must not start with '/'; keys have no leading slash",
));
}
if input.contains('\\') {
return Err(invalid("pattern must not contain backslash"));
}
if input.contains('\0') {
return Err(invalid("pattern must not contain NUL byte"));
}
let raw_segments: Vec<&str> = input.split('/').collect();
if raw_segments.len() > MAX_SEGMENTS {
return Err(invalid(format!("pattern exceeds {MAX_SEGMENTS} segments")));
}
let mut alternatives: Vec<Vec<Segment>> = vec![Vec::new()];
for raw in raw_segments {
if raw.is_empty() {
return Err(invalid(
"pattern must not contain empty segments (double slash or trailing slash)",
));
}
if raw == "." || raw == ".." {
return Err(invalid(
"pattern must not contain '.' or '..' segments; no key can contain them",
));
}
let expansions = expand_braces(raw)?;
if alternatives.len().saturating_mul(expansions.len()) > MAX_ALTERNATIVES {
return Err(invalid(format!(
"pattern expands to more than {MAX_ALTERNATIVES} brace alternatives"
)));
}
let mut next = Vec::with_capacity(alternatives.len() * expansions.len());
for prefix in &alternatives {
for expansion in &expansions {
let mut extended = prefix.clone();
extended.push(parse_segment(expansion)?);
next.push(extended);
}
}
alternatives = next;
}
Ok(Self {
source: input.to_string(),
alternatives,
})
}
pub fn matches(&self, key: &str) -> bool {
let segments: Vec<&str> = key.split('/').collect();
self.alternatives
.iter()
.any(|alternative| match_segments(alternative, &segments))
}
pub fn as_str(&self) -> &str {
&self.source
}
pub fn literal_prefix(&self) -> Option<&str> {
let mut shortest: Option<&str> = None;
for alternative in &self.alternatives {
let Some(Segment::Literal(first)) = alternative.first() else {
return None;
};
shortest = Some(match shortest {
None => first.as_str(),
Some(existing) if existing == first.as_str() => existing,
Some(_) => return None,
});
}
shortest
}
}
fn match_segments(pattern: &[Segment], key: &[&str]) -> bool {
let (mut p, mut k) = (0_usize, 0_usize);
let mut star_p: Option<usize> = None;
let mut star_k = 0_usize;
while k < key.len() {
match pattern.get(p) {
Some(Segment::DoubleStar) => {
star_p = Some(p);
star_k = k;
p += 1;
}
Some(segment) if segment_matches(segment, key[k]) => {
p += 1;
k += 1;
}
_ => {
let Some(resume) = star_p else { return false };
p = resume + 1;
star_k += 1;
k = star_k;
}
}
}
pattern[p..]
.iter()
.all(|segment| matches!(segment, Segment::DoubleStar))
}
fn segment_matches(segment: &Segment, candidate: &str) -> bool {
match segment {
Segment::DoubleStar => true,
Segment::Literal(literal) => literal == candidate,
Segment::Wild(tokens) => match_tokens(tokens, candidate),
}
}
fn match_tokens(tokens: &[Token], candidate: &str) -> bool {
let chars: Vec<char> = candidate.chars().collect();
let (mut t, mut c) = (0_usize, 0_usize);
let mut star_t: Option<usize> = None;
let mut star_c = 0_usize;
while c < chars.len() {
match tokens.get(t) {
Some(Token::Star) => {
star_t = Some(t);
star_c = c;
t += 1;
}
Some(Token::Any) => {
t += 1;
c += 1;
}
Some(Token::Literal(literal)) if literal_at(&chars, c, literal) => {
t += 1;
c += literal.chars().count();
}
_ => {
let Some(resume) = star_t else { return false };
t = resume + 1;
star_c += 1;
c = star_c;
}
}
}
tokens[t..].iter().all(|token| matches!(token, Token::Star))
}
fn literal_at(chars: &[char], at: usize, literal: &str) -> bool {
for (index, expected) in (at..).zip(literal.chars()) {
if chars.get(index) != Some(&expected) {
return false;
}
}
true
}
fn parse_segment(segment: &str) -> Result<Segment, Error> {
if segment == "**" {
return Ok(Segment::DoubleStar);
}
if segment.contains("**") {
return Err(invalid(format!(
"'**' must be a whole segment; found it inside '{segment}'"
)));
}
if !segment.contains(['*', '?']) {
return Ok(Segment::Literal(segment.to_string()));
}
let mut tokens = Vec::new();
let mut literal = String::new();
for ch in segment.chars() {
match ch {
'*' | '?' => {
if !literal.is_empty() {
tokens.push(Token::Literal(std::mem::take(&mut literal)));
}
tokens.push(if ch == '*' { Token::Star } else { Token::Any });
}
other => literal.push(other),
}
}
if !literal.is_empty() {
tokens.push(Token::Literal(literal));
}
Ok(Segment::Wild(tokens))
}
fn expand_braces(segment: &str) -> Result<Vec<String>, Error> {
if !segment.contains('{') {
if segment.contains('}') {
return Err(invalid(format!(
"'}}' without a matching '{{' in '{segment}'"
)));
}
return Ok(vec![segment.to_string()]);
}
let mut results = vec![String::new()];
let mut rest = segment;
while let Some(open) = rest.find('{') {
let (literal, tail) = rest.split_at(open);
let inner_and_rest = &tail[1..];
let Some(close) = inner_and_rest.find('}') else {
return Err(invalid(format!("unclosed '{{' in '{segment}'")));
};
let inner = &inner_and_rest[..close];
if inner.contains('{') {
return Err(invalid(format!(
"nested braces are not supported: '{segment}'"
)));
}
if inner.contains("**") {
return Err(invalid(format!(
"'**' must be a whole segment and cannot appear inside braces: '{segment}'"
)));
}
let branches: Vec<&str> = inner.split(',').collect();
if branches.iter().any(|branch| branch.is_empty()) {
return Err(invalid(format!(
"brace alternation must not contain an empty branch: '{segment}'"
)));
}
if results.len().saturating_mul(branches.len()) > MAX_ALTERNATIVES {
return Err(invalid(format!(
"pattern expands to more than {MAX_ALTERNATIVES} brace alternatives"
)));
}
let mut next = Vec::with_capacity(results.len() * branches.len());
for prefix in &results {
for branch in &branches {
next.push(format!("{prefix}{literal}{branch}"));
}
}
results = next;
rest = &inner_and_rest[close + 1..];
}
if rest.contains('}') {
return Err(invalid(format!(
"'}}' without a matching '{{' in '{segment}'"
)));
}
for result in &mut results {
result.push_str(rest);
}
Ok(results)
}
fn invalid(message: impl Into<String>) -> Error {
Error::InvalidInput {
message: message.into(),
}
}
impl fmt::Display for KeyPattern {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.source.fmt(f)
}
}
impl TryFrom<&str> for KeyPattern {
type Error = Error;
fn try_from(value: &str) -> Result<Self, Self::Error> {
Self::parse(value)
}
}
impl Serialize for KeyPattern {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(&self.source)
}
}
impl<'de> Deserialize<'de> for KeyPattern {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let source = String::deserialize(deserializer)?;
Self::parse(&source).map_err(serde::de::Error::custom)
}
}