use cbvault_format::cbh::{Entity, GameHeader};
use cbvault_format::error::{Error, Result};
use cbvault_format::game::{Eco, GameResult, RecordKind};
use rayon::prelude::*;
use super::namebase::Name;
use super::{Database, GameBuf};
pub fn any_player(ids: &[u32]) -> Filter {
if ids.is_empty() {
return Filter::Not(Box::new(Filter::All));
}
let mut iter = ids.iter();
let first = Filter::Either(*iter.next().expect("non-empty"));
iter.fold(first, |acc, id| Filter::AnyOf(Box::new(acc), Box::new(Filter::Either(*id))))
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct AllOf(Filter);
impl AllOf {
pub fn new() -> AllOf {
AllOf(Filter::All)
}
#[must_use]
pub fn and(mut self, filter: Filter) -> AllOf {
self.0 = match std::mem::take(&mut self.0) {
Filter::All => filter,
existing => Filter::AllOf(Box::new(existing), Box::new(filter)),
};
self
}
#[must_use]
pub fn and_opt(self, filter: Option<Filter>) -> AllOf {
match filter {
Some(f) => self.and(f),
None => self,
}
}
pub fn filter(self) -> Filter {
self.0
}
}
impl From<AllOf> for Filter {
fn from(all: AllOf) -> Filter {
all.0
}
}
#[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, Copy, Debug, Default, PartialEq, Eq)]
pub struct Range {
pub from: Option<i32>,
pub to: Option<i32>,
}
impl Range {
pub fn new(from: i32, to: i32) -> Range {
Range { from: Some(from), to: Some(to) }
}
pub fn at_least(min: i32) -> Range {
Range { from: Some(min), to: None }
}
pub fn at_most(max: i32) -> Range {
Range { from: None, to: Some(max) }
}
#[inline]
pub fn contains(&self, value: i32) -> bool {
if let Some(from) = self.from
&& value < from
{
return false;
}
if let Some(to) = self.to
&& value > to
{
return false;
}
self.from.is_some() || self.to.is_some()
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub enum Filter {
#[default]
All,
WhiteEloAtLeast(i32),
BlackEloAtLeast(i32),
EloAtLeast(i32),
WhiteEloBetween(Range),
BlackEloBetween(Range),
EloBetween(Range),
YearBetween(Range),
Result(GameResult),
Round {
round: u8,
sub: Option<u8>,
},
Eco {
code: u16,
sub: Option<u8>,
},
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 eco_text(text: &str) -> Option<Filter> {
let (code, sub) = parse_eco(text)?;
Some(Filter::Eco { code, sub })
}
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::WhiteEloBetween(range) => range.contains(i32::from(header.white_elo())),
Filter::BlackEloBetween(range) => range.contains(i32::from(header.black_elo())),
Filter::EloBetween(range) => {
range.contains(i32::from(header.white_elo())) || range.contains(i32::from(header.black_elo()))
}
Filter::YearBetween(range) => range.contains(i32::from(header.played_date().year())),
Filter::Result(want) => header.result() == *want,
Filter::Round { round, sub } => header.round() == *round && sub.is_none_or(|s| header.subround() == s),
Filter::Eco { code, sub } => match header.eco() {
Eco::Code { code: c, sub: s } => c == *code && sub.is_none_or(|w| s == w),
_ => false,
},
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),
}
}
}
fn parse_eco(text: &str) -> Option<(u16, Option<u8>)> {
let text = text.trim().as_bytes();
let letter = u16::from(match text.first()? {
b'A'..=b'E' => text[0] - b'A',
b'a'..=b'e' => text[0] - b'a',
_ => return None,
});
let digits = &text[1..];
if digits.len() > 3 || !digits.iter().all(u8::is_ascii_digit) {
return None;
}
let number: u16 = std::str::from_utf8(digits).ok()?.parse().unwrap_or(0);
let code = letter * 100 + number;
if code > 499 {
return None;
}
match digits.len() {
0 => Some((code, None)),
2 => Some((code, Some(0))),
3 => Some((code, Some(digits[2] - b'0'))),
_ => None,
}
}
#[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)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn eco_text_round_trips_every_code() {
for code in 0..=499u16 {
let text = Eco::Code { code, sub: 0 }.code_text().unwrap();
let text = std::str::from_utf8(&text).unwrap();
assert_eq!(parse_eco(text), Some((code, Some(0))), "{text} must parse back to code {code}");
for sub in [1u8, 7, 42, 127] {
let field = Eco::Code { code, sub }.field();
let back = Eco::from_field(field);
assert_eq!(back, Eco::Code { code, sub }, "code {code} sub {sub} must survive the field round trip");
}
}
}
#[test]
fn eco_text_rejects_what_is_not_a_code() {
for bad in ["", " ", "F20", "B20X", "20", "B999", "B2O", "b2O", "Z"] {
assert_eq!(parse_eco(bad), None, "{bad:?} is not an ECO code");
}
}
#[test]
fn a_bare_letter_covers_the_whole_hundred() {
assert_eq!(parse_eco("B"), Some((100, None)));
assert_eq!(parse_eco("E"), Some((400, None)));
assert_eq!(parse_eco("b"), Some((100, None)));
assert_eq!(parse_eco("F"), None);
}
#[test]
fn range_treats_none_as_open_not_zero() {
assert!(Range::at_least(2800).contains(2800));
assert!(Range::at_least(2800).contains(4000));
assert!(!Range::at_least(2800).contains(0));
assert!(!Range::at_least(2800).contains(2799));
assert!(Range::new(2000, 2010).contains(2000));
assert!(Range::new(2000, 2010).contains(2010));
assert!(!Range::new(2000, 2010).contains(2011));
assert!(!Range::new(2000, 2010).contains(1999));
assert!(Range::at_most(1500).contains(0));
assert!(!Range::at_most(1500).contains(1501));
assert!(!Range { from: None, to: None }.contains(0));
assert!(!Range::new(2010, 2000).contains(2005));
}
}