use serde::{Deserialize, Serialize};
use crate::hir::IrRule;
pub const HIR_SCHEMA_VERSION: u32 = 1;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct HirCacheHeader {
pub ir_schema_version: u32,
pub rsigma_version: String,
}
impl HirCacheHeader {
pub fn current() -> Self {
Self {
ir_schema_version: HIR_SCHEMA_VERSION,
rsigma_version: env!("CARGO_PKG_VERSION").to_string(),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct HirCache {
pub header: HirCacheHeader,
pub rules: Vec<IrRule>,
}
#[derive(Debug, thiserror::Error)]
pub enum CacheError {
#[error("HIR cache encode failed: {0}")]
Encode(String),
#[error("HIR cache decode failed: {0}")]
Decode(String),
#[error("HIR cache schema mismatch: blob is v{found}, this build expects v{expected}")]
SchemaMismatch { expected: u32, found: u32 },
}
pub fn encode_rules(rules: &[IrRule]) -> Result<Vec<u8>, CacheError> {
let mut buf = Vec::new();
ciborium::into_writer(&HirCacheHeader::current(), &mut buf)
.map_err(|e| CacheError::Encode(e.to_string()))?;
ciborium::into_writer(&rules, &mut buf).map_err(|e| CacheError::Encode(e.to_string()))?;
Ok(buf)
}
pub fn decode_rules(bytes: &[u8]) -> Result<Vec<IrRule>, CacheError> {
let mut cursor = std::io::Cursor::new(bytes);
let header: HirCacheHeader =
ciborium::from_reader(&mut cursor).map_err(|e| CacheError::Decode(e.to_string()))?;
if header.ir_schema_version != HIR_SCHEMA_VERSION {
return Err(CacheError::SchemaMismatch {
expected: HIR_SCHEMA_VERSION,
found: header.ir_schema_version,
});
}
ciborium::from_reader(&mut cursor).map_err(|e| CacheError::Decode(e.to_string()))
}
pub fn to_json(rules: &[IrRule]) -> Result<String, CacheError> {
let cache = HirCache {
header: HirCacheHeader::current(),
rules: rules.to_vec(),
};
serde_json::to_string_pretty(&cache).map_err(|e| CacheError::Encode(e.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hir::{IrRule, IrRuleMetadata};
use rsigma_parser::LogSource;
fn sample_rule(title: &str) -> IrRule {
IrRule {
metadata: IrRuleMetadata {
title: title.to_string(),
..Default::default()
},
logsource: LogSource::default(),
sigma_version: None,
detections: Default::default(),
conditions: Vec::new(),
}
}
#[test]
fn round_trips_through_postcard() {
let rules = vec![sample_rule("a"), sample_rule("b")];
let blob = encode_rules(&rules).unwrap();
let decoded = decode_rules(&blob).unwrap();
assert_eq!(decoded, rules);
}
#[test]
fn rejects_schema_mismatch() {
let rules = vec![sample_rule("a")];
let mut blob = Vec::new();
let header = HirCacheHeader {
ir_schema_version: HIR_SCHEMA_VERSION + 1,
rsigma_version: "test".to_string(),
};
ciborium::into_writer(&header, &mut blob).unwrap();
ciborium::into_writer(&rules, &mut blob).unwrap();
match decode_rules(&blob) {
Err(CacheError::SchemaMismatch { expected, found }) => {
assert_eq!(expected, HIR_SCHEMA_VERSION);
assert_eq!(found, HIR_SCHEMA_VERSION + 1);
}
other => panic!("expected schema mismatch, got {other:?}"),
}
}
#[test]
fn json_export_is_readable() {
let rules = vec![sample_rule("json")];
let json = to_json(&rules).unwrap();
assert!(json.contains("\"ir_schema_version\""));
assert!(json.contains("\"json\""));
}
}