use std::ffi::{c_int, c_uint, c_ulonglong, c_void};
use std::mem::MaybeUninit;
use std::sync::Arc;
use foreign_types::ForeignType;
use hypervectorscan_sys as hs;
use super::{AsResult, Error, HyperscanErrorCode, Pattern, ScanMode, wrapper};
use crate::SerializedDatabase;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum Scan {
Continue,
Terminate,
}
#[derive(Clone, Debug)]
pub struct BlockDatabase {
inner: wrapper::Database,
}
impl BlockDatabase {
pub fn new(patterns: Vec<Pattern>) -> Result<Self, Error> {
let inner = wrapper::Database::new(patterns, ScanMode::BLOCK)?;
Ok(Self { inner })
}
pub fn create_scanner(self: &Arc<Self>) -> Result<BlockScanner, Error> {
BlockScanner::new(self.clone())
}
pub fn size(&self) -> Result<usize, Error> {
self.inner.size()
}
pub fn serialize(&self) -> Result<SerializedDatabase, Error> {
self.inner.serialize()
}
pub fn deserialize(sdb: SerializedDatabase) -> Result<Self, Error> {
let db = wrapper::Database::deserialize(sdb)?;
Ok(Self { inner: db })
}
}
#[derive(Clone, Debug)]
pub struct BlockScanner {
scratch: wrapper::Scratch,
db: Arc<BlockDatabase>,
}
impl BlockScanner {
pub fn new(db: Arc<BlockDatabase>) -> Result<Self, Error> {
Ok(Self {
scratch: wrapper::Scratch::new(&db.inner)?,
db,
})
}
pub fn scan<F>(&mut self, data: &[u8], on_match: F) -> Result<Scan, Error>
where
F: FnMut(u32, u64, u64, u32) -> Scan,
{
let mut context = Context { on_match };
let res = unsafe {
hs::hs_scan(
self.db.inner.as_ptr(),
data.as_ptr() as *const _,
data.len() as u32,
0,
self.scratch.as_ptr(),
Some(on_match_trampoline::<F>),
&mut context as *mut _ as *mut c_void,
)
.ok()
};
match res {
Ok(_) => Ok(Scan::Continue),
Err(err) => match err {
Error::Hyperscan(HyperscanErrorCode::ScanTerminated) => Ok(Scan::Terminate),
err => Err(err),
},
}
}
pub fn size(&self) -> Result<usize, Error> {
self.scratch.size()
}
}
#[derive(Clone, Debug)]
pub struct StreamingDatabase {
inner: wrapper::Database,
}
impl StreamingDatabase {
pub fn new(patterns: Vec<Pattern>) -> Result<Self, Error> {
let inner = wrapper::Database::new(patterns, ScanMode::STREAM)?;
Ok(Self { inner })
}
pub fn create_scanner(self: &Arc<Self>) -> Result<StreamingScanner, Error> {
StreamingScanner::new(self.clone())
}
pub fn size(&self) -> Result<usize, Error> {
self.inner.size()
}
pub fn stream_size(&self) -> Result<usize, Error> {
self.inner.stream_size()
}
pub fn serialize(&self) -> Result<SerializedDatabase, Error> {
self.inner.serialize()
}
pub fn deserialize(sdb: SerializedDatabase) -> Result<Self, Error> {
let db = wrapper::Database::deserialize(sdb)?;
Ok(Self { inner: db })
}
}
#[derive(Debug)]
pub struct Stream {
inner: *mut hs::hs_stream_t,
}
impl Stream {
fn new(database: &StreamingDatabase) -> Result<Self, Error> {
let mut inner = MaybeUninit::zeroed();
let flags = 0;
unsafe {
hs::hs_open_stream(database.inner.as_ptr(), flags, inner.as_mut_ptr())
.ok()
.map(|()| Self {
inner: inner.assume_init(),
})
}
}
}
#[derive(Clone, Debug)]
pub struct StreamingScanner {
scratch: wrapper::Scratch,
db: Arc<StreamingDatabase>,
}
#[derive(Debug)]
pub struct StreamScanner {
scanner: Arc<StreamingScanner>,
stream: Stream,
}
impl StreamingScanner {
pub fn new(db: Arc<StreamingDatabase>) -> Result<Self, Error> {
Ok(Self {
scratch: wrapper::Scratch::new(&db.inner)?,
db,
})
}
pub fn open_stream(self: &Arc<Self>) -> Result<StreamScanner, Error> {
let stream = Stream::new(&self.db)?;
Ok(StreamScanner {
stream,
scanner: self.clone(),
})
}
}
impl StreamScanner {
pub fn close<F>(self, on_match: F) -> Result<Scan, Error>
where
F: FnMut(u32, u64, u64, u32) -> Scan,
{
let mut context = Context { on_match };
let res = unsafe {
hs::hs_close_stream(
self.stream.inner,
self.scanner.scratch.as_ptr(),
Some(on_match_trampoline::<F>),
&mut context as *mut _ as *mut c_void,
)
.ok()
};
match res {
Ok(_) => Ok(Scan::Continue),
Err(err) => match err {
Error::Hyperscan(HyperscanErrorCode::ScanTerminated) => Ok(Scan::Terminate),
err => Err(err),
},
}
}
pub fn scan<F>(&mut self, data: &[u8], on_match: F) -> Result<Scan, Error>
where
F: FnMut(u32, u64, u64, u32) -> Scan,
{
let mut context = Context { on_match };
let res = unsafe {
hs::hs_scan_stream(
self.stream.inner,
data.as_ptr() as *const _,
data.len() as u32,
0,
self.scanner.scratch.as_ptr(),
Some(on_match_trampoline::<F>),
&mut context as *mut _ as *mut c_void,
)
.ok()
};
match res {
Ok(_) => Ok(Scan::Continue),
Err(err) => match err {
Error::Hyperscan(HyperscanErrorCode::ScanTerminated) => Ok(Scan::Terminate),
err => Err(err),
},
}
}
}
struct Context<F>
where
F: FnMut(u32, u64, u64, u32) -> Scan,
{
on_match: F,
}
unsafe extern "C" fn on_match_trampoline<F>(
id: c_uint,
from: c_ulonglong,
to: c_ulonglong,
flags: c_uint,
ctx: *mut c_void,
) -> c_int
where
F: FnMut(u32, u64, u64, u32) -> Scan,
{
let context = unsafe {
(ctx as *mut Context<F>)
.as_mut()
.expect("context object should be set")
};
match (context.on_match)(id, from, to, flags) {
Scan::Continue => 0,
Scan::Terminate => 1,
}
}