use cbvault_format::cbh::{Entity, GameHeader};
use cbvault_format::error::{Error, Result};
use cbvault_format::game::RecordKind;
use rayon::prelude::*;
use super::namebase::Name;
use super::{Database, GameBuf};
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct IdSet {
ids: Vec<u32>,
}
impl IdSet {
pub fn new() -> IdSet {
IdSet::default()
}
pub fn from_ids(ids: impl IntoIterator<Item = u32>) -> IdSet {
let mut ids: Vec<u32> = ids.into_iter().collect();
ids.sort_unstable();
ids.dedup();
IdSet { ids }
}
pub fn len(&self) -> usize {
self.ids.len()
}
pub fn is_empty(&self) -> bool {
self.ids.is_empty()
}
pub fn contains(&self, id: u32) -> bool {
self.ids.binary_search(&id).is_ok()
}
}
impl FromIterator<u32> for IdSet {
fn from_iter<T: IntoIterator<Item = u32>>(ids: T) -> IdSet {
IdSet::from_ids(ids)
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub enum Filter {
#[default]
All,
WhiteEloAtLeast(i32),
BlackEloAtLeast(i32),
EloAtLeast(i32),
Player(u32),
Opponent(u32),
Either(u32),
Tournament(u32),
Annotator(u32),
Source(u32),
Ids(IdSet),
AllOf(Box<Filter>, Box<Filter>),
AnyOf(Box<Filter>, Box<Filter>),
Not(Box<Filter>),
}
impl Filter {
pub fn player(id: u32) -> Filter {
Filter::Either(id)
}
pub fn and(self, other: Filter) -> Filter {
Filter::AllOf(Box::new(self), Box::new(other))
}
pub fn or(self, other: Filter) -> Filter {
Filter::AnyOf(Box::new(self), Box::new(other))
}
#[allow(clippy::should_implement_trait)]
pub fn not(self) -> Filter {
Filter::Not(Box::new(self))
}
pub fn matches(&self, id: u32, header: &GameHeader) -> bool {
match self {
Filter::All => true,
Filter::WhiteEloAtLeast(min) => i32::from(header.white_elo()) >= *min,
Filter::BlackEloAtLeast(min) => i32::from(header.black_elo()) >= *min,
Filter::EloAtLeast(min) => i32::from(header.white_elo()) >= *min || i32::from(header.black_elo()) >= *min,
Filter::Player(id) => header.white() == *id,
Filter::Opponent(id) => header.black() == *id,
Filter::Either(id) => header.white() == *id || header.black() == *id,
Filter::Tournament(id) => header.tournament() == *id,
Filter::Annotator(id) => header.annotator() == *id,
Filter::Source(id) => header.source() == *id,
Filter::Ids(ids) => ids.contains(id),
Filter::AllOf(a, b) => a.matches(id, header) && b.matches(id, header),
Filter::AnyOf(a, b) => a.matches(id, header) || b.matches(id, header),
Filter::Not(inner) => !inner.matches(id, header),
}
}
}
#[derive(Clone, Copy)]
pub struct Match<'a> {
pub id: u32,
pub header: GameHeader,
pub white: Name<'a>,
pub black: Name<'a>,
pub event: Name<'a>,
}
#[derive(Debug, Default)]
pub struct Scan<'a> {
matches: Vec<Match<'a>>,
records: u64,
threads: usize,
}
impl<'a> Scan<'a> {
pub fn matches(&self) -> &[Match<'a>] {
&self.matches
}
pub fn len(&self) -> usize {
self.matches.len()
}
pub fn is_empty(&self) -> bool {
self.matches.is_empty()
}
pub fn into_matches(self) -> Vec<Match<'a>> {
self.matches
}
pub fn ids(&self) -> Vec<u32> {
self.matches.iter().map(|m| m.id).collect()
}
pub fn stats(&self) -> SearchStats {
SearchStats { records: self.records, matches: self.matches.len() as u64, threads: self.threads }
}
}
impl Match<'_> {
pub fn player(name: &Name<'_>, out: &mut String) -> bool {
if name.is_empty() {
return false;
}
name.push_text(out);
true
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct SearchStats {
pub records: u64,
pub matches: u64,
pub threads: usize,
}
const SCAN_BATCH: u32 = 8192;
pub fn scan<'db>(db: &'db Database, filter: &Filter, threads: usize) -> Result<Scan<'db>> {
let ranges = ranges_of(db, 1, 0);
scan_ranges(db, &ranges, filter, threads)
}
pub fn scan_range<'db>(db: &'db Database, first: u32, last: u32, filter: &Filter, threads: usize) -> Result<Scan<'db>> {
scan_ranges(db, &ranges_of(db, first, last), filter, threads)
}
fn ranges_of(db: &Database, first: u32, last: u32) -> Vec<(u32, u32)> {
let last = if last == 0 { db.records() } else { last.min(db.records()) };
let mut out = Vec::new();
let mut at = first.max(1);
while at <= last {
let end = at.saturating_add(SCAN_BATCH - 1).min(last);
out.push((at, end));
at = end + 1;
}
out
}
fn scan_ranges<'db>(db: &'db Database, ranges: &[(u32, u32)], filter: &Filter, threads: usize) -> Result<Scan<'db>> {
let entities = db.entities();
let scan_chunk = |&(first, last): &(u32, u32)| -> Result<Vec<Match<'db>>> {
let mut out = Vec::new();
let mut batch = db.record_batch(first, last)?;
for header in batch.by_ref() {
if !matches!(header.kind(), RecordKind::Game) || header.is_deleted() {
continue;
}
if !filter.matches(header.id(), &header) {
continue;
}
out.push(Match {
id: header.id(),
header,
white: entities.name(Entity::Player, header.white())?.unwrap_or(BLANK),
black: entities.name(Entity::Player, header.black())?.unwrap_or(BLANK),
event: Name::of(entities.name(Entity::Tournament, header.tournament())?.unwrap_or(BLANK).last()),
});
}
Ok(out)
};
let found: Vec<Vec<Match<'db>>> = if threads <= 1 || ranges.len() <= 1 {
ranges.iter().map(&scan_chunk).collect::<Result<_>>()?
} else {
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.map_err(|e| Error::corrupt(db.base(), 0, format!("thread pool creation failed: {e}")))?;
pool.install(|| ranges.par_iter().map(scan_chunk).collect::<Result<Vec<_>>>())?
};
let records: u64 = ranges.iter().map(|(first, last)| u64::from(last - first + 1)).sum();
let matches = found.into_iter().flatten().collect::<Vec<_>>();
Ok(Scan { matches, records, threads: if threads <= 1 { 1 } else { threads } })
}
const BLANK: Name<'static> = Name::blank();
impl std::fmt::Debug for Match<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Match")
.field("id", &self.id)
.field("white", &String::from_utf8_lossy(self.white.last()))
.field("black", &String::from_utf8_lossy(self.black.last()))
.field("event", &String::from_utf8_lossy(self.event.last()))
.finish()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Hit {
pub game: u32,
pub ply: u32,
pub key: u64,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct PositionSearch {
pub games: u64,
pub plies: u64,
pub hits: u64,
pub failures: u64,
pub complete: bool,
}
#[derive(Clone, Copy)]
pub struct PositionQuery<'p> {
pub key: u64,
pub every: u64,
pub cancelled: Option<&'p (dyn Fn() -> bool + Sync)>,
}
impl std::fmt::Debug for PositionQuery<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PositionQuery")
.field("key", &format_args!("{:#018x}", self.key))
.field("every", &self.every)
.field("cancelled", &self.cancelled.is_some())
.finish()
}
}
impl Default for PositionQuery<'_> {
fn default() -> PositionQuery<'static> {
PositionQuery { key: 0, every: 0, cancelled: None }
}
}
impl<'p> PositionQuery<'p> {
pub fn of(key: u64) -> PositionQuery<'static> {
PositionQuery { key, every: 0, cancelled: None }
}
pub fn every(mut self, every: u64) -> PositionQuery<'p> {
self.every = every;
self
}
pub fn cancel_with(mut self, cancelled: &'p (dyn Fn() -> bool + Sync)) -> PositionQuery<'p> {
self.cancelled = Some(cancelled);
self
}
}
#[derive(Debug, Default)]
struct Replay {
stats: PositionSearch,
hits: Vec<Hit>,
}
pub fn for_each_position_key(
db: &Database,
query: &PositionQuery<'_>,
threads: usize,
mut on_hit: impl FnMut(Hit) + Send,
) -> Result<PositionSearch> {
let ranges = ranges_of(db, 1, 0);
for_each_position_key_range(db, &ranges, query, threads, &mut on_hit)
}
pub fn for_each_position_key_range(
db: &Database,
ranges: &[(u32, u32)],
query: &PositionQuery<'_>,
threads: usize,
on_hit: &mut (impl FnMut(Hit) + Send),
) -> Result<PositionSearch> {
let want = query.key;
let moves = db.moves()?;
let path = db.members().moves.clone();
let replay_chunk = |&(first, last): &(u32, u32)| -> Result<Replay> {
let mut out = Replay::default();
let mut buf = GameBuf::with_capacity(super::walk::DEFAULT_PLY_ROOM);
buf.set_wants(true, false);
let mut record = Vec::new();
let mut headers = db.record_batch(first, last)?;
for header in headers.by_ref() {
if !matches!(header.kind(), RecordKind::Game) || header.is_deleted() {
continue;
}
let Some(at) = db.offset_of(header.id(), header.moves_offset(), header.annotations_offset()).ok() else {
out.stats.failures += 1;
continue;
};
let Ok(bytes) = super::convert::move_record(moves, at, &mut record) else {
out.stats.failures += 1;
continue;
};
let Ok(game) = cbvault_format::cbh::moves::GameMoves::parse(&path, bytes) else {
out.stats.failures += 1;
continue;
};
if buf.walk(header.id(), at, &game).is_err() {
out.stats.failures += 1;
continue;
}
out.stats.games += 1;
out.stats.plies += buf.moves().len() as u64;
for (ply, key) in buf.keys().iter().enumerate() {
if *key == want {
out.hits.push(Hit { game: header.id(), ply: ply as u32, key: *key });
}
}
}
Ok(out)
};
let replays: Vec<Replay> = if threads <= 1 || ranges.len() <= 1 {
ranges.iter().map(&replay_chunk).collect::<Result<_>>()?
} else {
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.map_err(|e| Error::corrupt(db.base(), 0, format!("thread pool creation failed: {e}")))?;
pool.install(|| ranges.par_iter().map(replay_chunk).collect::<Result<Vec<_>>>())?
};
let mut total = PositionSearch { complete: true, ..PositionSearch::default() };
for replay in replays {
total.games += replay.stats.games;
total.plies += replay.stats.plies;
total.failures += replay.stats.failures;
for hit in replay.hits {
on_hit(hit);
total.hits += 1;
}
if query.cancelled.is_some_and(|cancel| cancel()) {
total.complete = false;
}
}
Ok(total)
}