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 {
Filter::PlayerSet(IdSet::from_ids(ids.iter().copied()))
}
#[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),
PlayerSet(IdSet),
WhiteSet(IdSet),
BlackSet(IdSet),
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 player_set(ids: IdSet) -> Filter {
Filter::PlayerSet(ids)
}
pub fn white_set(ids: IdSet) -> Filter {
Filter::WhiteSet(ids)
}
pub fn black_set(ids: IdSet) -> Filter {
Filter::BlackSet(ids)
}
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::PlayerSet(ids) => ids.contains(header.white()) || ids.contains(header.black()),
Filter::WhiteSet(ids) => ids.contains(header.white()),
Filter::BlackSet(ids) => ids.contains(header.black()),
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));
}
fn mock_header(white: u32, black: u32) -> GameHeader {
let mut b = [0u8; cbvault_format::cbh::RECORD_SIZE];
b[0] = 1; b[0x09] = (white >> 16) as u8;
b[0x0a] = (white >> 8) as u8;
b[0x0b] = white as u8;
b[0x0c] = (black >> 16) as u8;
b[0x0d] = (black >> 8) as u8;
b[0x0e] = black as u8;
GameHeader::from_bytes(1, &b)
}
#[test]
fn any_player_empty_matches_nothing() {
let f = any_player(&[]);
assert!(!f.matches(1, &mock_header(1, 2)));
assert!(!f.matches(1, &mock_header(0, 0)));
}
#[test]
fn any_player_matches_identically_to_fold_and_handles_large_sets() {
let legacy_fold = |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))))
};
let ids_1 = [42];
let f1 = any_player(&ids_1);
let fold1 = legacy_fold(&ids_1);
for (w, b) in [(42, 1), (1, 42), (42, 42), (99, 100)] {
let h = mock_header(w, b);
assert_eq!(f1.matches(1, &h), fold1.matches(1, &h));
}
let ids_2 = [10, 20];
let f2 = any_player(&ids_2);
let fold2 = legacy_fold(&ids_2);
for (w, b) in [(10, 1), (1, 20), (20, 10), (30, 40)] {
let h = mock_header(w, b);
assert_eq!(f2.matches(1, &h), fold2.matches(1, &h));
}
let ids_100: Vec<u32> = (1..=100).collect();
let f100 = any_player(&ids_100);
let fold100 = legacy_fold(&ids_100);
for id in [1, 50, 100, 101, 500] {
let h_white = mock_header(id, 9999);
let h_black = mock_header(9999, id);
assert_eq!(f100.matches(1, &h_white), fold100.matches(1, &h_white));
assert_eq!(f100.matches(1, &h_black), fold100.matches(1, &h_black));
}
let white_set = Filter::white_set(IdSet::from_ids([10, 20]));
let black_set = Filter::black_set(IdSet::from_ids([10, 20]));
assert!(white_set.matches(1, &mock_header(10, 99)));
assert!(!white_set.matches(1, &mock_header(99, 10)));
assert!(black_set.matches(1, &mock_header(99, 20)));
assert!(!black_set.matches(1, &mock_header(20, 99)));
}
}