use std::sync::Arc;
use bytes::Bytes;
use super::core::{read_u32_le, read_u64_le};
use crate::superfile::{
ReadError,
error::FtsError,
format::{
self,
fts::{U32_BYTES, U64_BYTES, skip_entry, term_meta},
},
fts::{
bm25,
builder::{SKIP_ENTRY_SIZE, TERM_META_POSITIONAL_SIZE, TERM_META_SIZE},
posting::{BLOCK_LEN, decode_block},
},
};
#[derive(Debug, Copy, Clone)]
pub(super) struct TermMeta {
pub(super) df: u64,
pub(super) postings_length: usize,
pub(super) num_blocks: usize,
pub(super) skip_start: usize,
pub(super) positions_offset: u64,
pub(super) positions_length: u32,
}
impl TermMeta {
pub(super) fn parse(
postings: &[u8],
metadata_offset: usize,
positional: bool,
) -> Result<Self, FtsError> {
let term_meta_size = match positional {
true => TERM_META_POSITIONAL_SIZE,
false => TERM_META_SIZE,
};
if metadata_offset + term_meta_size > postings.len() {
return Err(FtsError::Read(ReadError::MalformedVersion(
"term metadata offset out of postings region".into(),
)));
}
let df = read_u32_le(
&postings[metadata_offset + term_meta::DF_OFF
..metadata_offset + term_meta::DF_OFF + U32_BYTES],
) as u64;
let postings_length = read_u32_le(
&postings[metadata_offset + term_meta::POSTINGS_LENGTH_OFF
..metadata_offset + term_meta::POSTINGS_LENGTH_OFF + U32_BYTES],
) as usize;
let num_blocks = read_u32_le(
&postings[metadata_offset + term_meta::NUM_BLOCKS_OFF
..metadata_offset + term_meta::NUM_BLOCKS_OFF + U32_BYTES],
) as usize;
let (positions_offset, positions_length) = match positional {
true => (
read_u64_le(
&postings[metadata_offset + term_meta::POSITIONS_OFFSET_OFF
..metadata_offset + term_meta::POSITIONS_OFFSET_OFF + U64_BYTES],
),
read_u32_le(
&postings[metadata_offset + term_meta::POSITIONS_LENGTH_OFF
..metadata_offset + term_meta::POSITIONS_LENGTH_OFF + U32_BYTES],
),
),
false => (0, 0),
};
if metadata_offset + postings_length > postings.len() {
return Err(FtsError::Read(ReadError::MalformedVersion(
"term postings length exceeds the fetched term range".into(),
)));
}
let skip_start = metadata_offset + term_meta_size;
let skip_end = skip_start + num_blocks * SKIP_ENTRY_SIZE;
if skip_end > postings.len() {
return Err(FtsError::Read(ReadError::MalformedVersion(
"skip table runs past postings region".into(),
)));
}
Ok(Self {
df,
postings_length,
num_blocks,
skip_start,
positions_offset,
positions_length,
})
}
#[inline]
pub(super) fn skip_entry(&self, postings: &[u8], i: usize) -> (u32, usize, f32) {
debug_assert!(i < self.num_blocks, "skip entry {i} >= {}", self.num_blocks);
let entry_off = self.skip_start + i * SKIP_ENTRY_SIZE;
let last_doc_id = read_u32_le(
&postings[entry_off + skip_entry::LAST_DOC_ID_OFF
..entry_off + skip_entry::LAST_DOC_ID_OFF + U32_BYTES],
);
let block_offset = read_u32_le(
&postings[entry_off + skip_entry::BLOCK_OFFSET_OFF
..entry_off + skip_entry::BLOCK_OFFSET_OFF + U32_BYTES],
) as usize;
let max_bm25_x1000 = read_u32_le(
&postings[entry_off + skip_entry::MAX_BM25_OFF
..entry_off + skip_entry::MAX_BM25_OFF + U32_BYTES],
);
(
last_doc_id,
block_offset,
max_bm25_x1000.saturating_add(1) as f32 / format::fts::BLOCK_MAX_BM25_FIXED_POINT_SCALE,
)
}
#[inline]
pub(super) fn positions_block_offset(&self, postings: &[u8], i: usize) -> u32 {
debug_assert!(i < self.num_blocks, "skip entry {i} >= {}", self.num_blocks);
let entry_off = self.skip_start + i * SKIP_ENTRY_SIZE;
read_u32_le(
&postings[entry_off + skip_entry::POSITIONS_BLOCK_OFFSET_OFF
..entry_off + skip_entry::POSITIONS_BLOCK_OFFSET_OFF + U32_BYTES],
)
}
#[inline]
pub(super) fn block_end_in_term(&self, postings: &[u8], i: usize) -> usize {
if i + 1 < self.num_blocks {
let next_off = self.skip_start + (i + 1) * SKIP_ENTRY_SIZE;
read_u32_le(&postings[next_off + 4..next_off + 8]) as usize
} else {
self.postings_length
}
}
}
#[derive(Debug, Clone, Copy)]
pub(super) struct BlockMeta {
pub(super) last_doc_id: u32,
pub(super) block_byte_offset: usize,
pub(super) block_byte_end: usize,
pub(super) block_max_bm25: f32,
}
#[derive(Clone)]
pub(crate) struct TermCursor {
pub(super) idf_x_k1p1: f32,
pub(super) term_max_bm25: f32,
pub(super) df: u64,
pub(super) blocks: Arc<[BlockMeta]>,
pub(super) block_doc_ids: Vec<u32>,
pub(super) block_tfs: Vec<u32>,
pub(super) block_n: usize,
pub(super) current_block: usize,
pub(super) pos: usize,
pub(super) inspect_block: usize,
pub(super) bytes: Bytes,
pub(super) header_probed: bool,
}
impl TermCursor {
pub(super) fn new(
term_bytes: Bytes,
n_docs: u64,
positional: bool,
global_idf: Option<f32>,
header_probed: bool,
) -> Result<Self, FtsError> {
let postings: &[u8] = term_bytes.as_ref();
let metadata_offset = 0usize;
let term_meta = TermMeta::parse(postings, metadata_offset, positional)?;
let local_idf = bm25::idf(n_docs, term_meta.df);
let idf = global_idf.unwrap_or(local_idf);
let idf_rescale = match global_idf {
Some(_) if local_idf > 0.0 && idf != local_idf => Some(idf / local_idf),
_ => None,
};
let mut term_max_bm25: f32 = 0.0;
let blocks: Arc<[BlockMeta]> = (0..term_meta.num_blocks)
.map(|i| {
let (last_doc_id, block_offset_in_term, raw_block_max) =
term_meta.skip_entry(postings, i);
let block_max_bm25 = match idf_rescale {
Some(ratio) => raw_block_max * ratio,
None => raw_block_max,
};
term_max_bm25 = term_max_bm25.max(block_max_bm25);
BlockMeta {
last_doc_id,
block_byte_offset: metadata_offset + block_offset_in_term,
block_byte_end: metadata_offset + term_meta.block_end_in_term(postings, i),
block_max_bm25,
}
})
.collect();
let mut cursor = Self {
idf_x_k1p1: idf * (bm25::K1 + 1.0),
term_max_bm25,
df: term_meta.df,
blocks,
block_doc_ids: vec![0u32; BLOCK_LEN],
block_tfs: vec![0u32; BLOCK_LEN],
block_n: 0,
current_block: 0,
pos: 0,
inspect_block: 0,
bytes: term_bytes,
header_probed,
};
if !cursor.blocks.is_empty() {
cursor.decode_current_block();
}
Ok(cursor)
}
pub(super) fn new_inline(
doc_id: u32,
tf: u32,
n_docs: u64,
dl_norm_k1: f32,
global_idf: Option<f32>,
) -> Self {
let idf = global_idf.unwrap_or_else(|| bm25::idf(n_docs, 1));
let idf_x_k1p1 = idf * (bm25::K1 + 1.0);
let block_max_bm25 = bm25::score_with_dl_norm_k1(idf_x_k1p1, tf, dl_norm_k1);
let blocks: Arc<[BlockMeta]> = Arc::from([BlockMeta {
last_doc_id: doc_id,
block_byte_offset: 0,
block_byte_end: 0,
block_max_bm25,
}]);
let mut block_doc_ids = vec![0u32; BLOCK_LEN];
let mut block_tfs = vec![0u32; BLOCK_LEN];
block_doc_ids[0] = doc_id;
block_tfs[0] = tf;
Self {
idf_x_k1p1,
term_max_bm25: block_max_bm25,
df: 1,
blocks,
block_doc_ids,
block_tfs,
block_n: 1,
current_block: 0,
pos: 0,
inspect_block: 0,
bytes: Bytes::new(),
header_probed: false,
}
}
pub(super) fn decode_current_block(&mut self) {
let block = self.blocks[self.current_block];
let bytes = self
.bytes
.slice(block.block_byte_offset..block.block_byte_end);
self.block_n = decode_block(&bytes, &mut self.block_doc_ids, &mut self.block_tfs);
self.pos = 0;
}
pub(super) fn is_exhausted(&self) -> bool {
self.current_block >= self.blocks.len()
}
#[inline(always)]
pub(super) fn block_count(&self) -> usize {
self.blocks.len()
}
#[inline(always)]
pub(super) fn current_doc_id(&self) -> u32 {
if self.is_exhausted() || self.pos >= self.block_n {
u32::MAX
} else {
self.block_doc_ids[self.pos]
}
}
#[inline(always)]
pub(super) fn current_tf(&self) -> u32 {
debug_assert!(!self.is_exhausted() && self.pos < self.block_n);
self.block_tfs[self.pos]
}
#[inline(always)]
pub(super) fn current_block_max_bm25(&self) -> f32 {
if self.is_exhausted() {
0.0
} else {
self.blocks[self.current_block].block_max_bm25
}
}
#[inline(always)]
pub(super) fn current_block_last_doc_id(&self) -> u32 {
if self.is_exhausted() {
u32::MAX
} else {
self.blocks[self.current_block].last_doc_id
}
}
pub(super) fn shallow_advance_block_to(&mut self, target: u32) {
if self.inspect_block < self.current_block {
self.inspect_block = self.current_block;
}
while self.inspect_block < self.blocks.len()
&& self.blocks[self.inspect_block].last_doc_id < target
{
self.inspect_block += 1;
}
}
pub(super) fn block_max_in_range(&mut self, range_start: u32, range_end: u32) -> f32 {
self.shallow_advance_block_to(range_start);
let mut max: f32 = 0.0;
let mut i = self.inspect_block;
while i < self.blocks.len() {
let block_start = if i == 0 {
0u32
} else {
self.blocks[i - 1].last_doc_id.saturating_add(1)
};
if block_start > range_end {
break;
}
let m = self.blocks[i].block_max_bm25;
if m > max {
max = m;
}
i += 1;
}
max
}
pub(super) fn inspect_block_max_bm25(&self) -> f32 {
if self.inspect_block >= self.blocks.len() {
0.0
} else {
self.blocks[self.inspect_block].block_max_bm25
}
}
pub(super) fn inspect_block_last_doc_id(&self) -> u32 {
if self.inspect_block >= self.blocks.len() {
u32::MAX
} else {
self.blocks[self.inspect_block].last_doc_id
}
}
#[inline(always)]
pub(super) fn next(&mut self) {
if self.is_exhausted() {
return;
}
self.pos += 1;
if self.pos >= self.block_n {
self.advance_block();
}
}
#[inline(always)]
pub(super) fn advance_by(&mut self, count: usize) {
debug_assert!(!self.is_exhausted());
debug_assert!(count > 0 && self.pos + count <= self.block_n);
self.pos += count;
if self.pos == self.block_n {
self.advance_block();
}
}
#[inline(always)]
pub(super) fn advance_block(&mut self) {
self.current_block += 1;
if self.current_block > self.inspect_block {
self.inspect_block = self.current_block;
}
if self.current_block < self.blocks.len() {
self.decode_current_block();
}
}
#[inline(always)]
pub(super) fn skip_to(&mut self, target: u32) {
if self.is_exhausted() {
return;
}
let cur_block = self.current_block;
let cur_block_last = self.blocks[cur_block].last_doc_id;
if cur_block_last >= target {
let n = self.block_n;
while self.pos < n && self.block_doc_ids[self.pos] < target {
self.pos += 1;
}
if self.pos < n {
return;
}
}
self.skip_to_cross_block(target);
}
#[cold]
pub(super) fn skip_to_cross_block(&mut self, target: u32) {
while self.current_block < self.blocks.len()
&& self.blocks[self.current_block].last_doc_id < target
{
self.current_block += 1;
}
if self.current_block > self.inspect_block {
self.inspect_block = self.current_block;
}
if self.is_exhausted() {
return;
}
self.decode_current_block();
while self.pos < self.block_n && self.block_doc_ids[self.pos] < target {
self.pos += 1;
}
if self.pos >= self.block_n {
self.current_block += 1;
if self.current_block > self.inspect_block {
self.inspect_block = self.current_block;
}
if self.current_block < self.blocks.len() {
self.decode_current_block();
}
}
}
}