use crate::{
Column, CopyRowReader, DumpId, EntryDataReader, FieldRef, Limits, OwnedRow, PgDumpError, Row,
ScanLimits, copy_metadata::TableDataMetadata,
};
use std::io::Read;
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum ColumnEqualityResult {
Match(OwnedRow),
NoMatch,
ColumnNotFound,
}
pub struct TableRowReader<'a, R> {
data_id: DumpId,
metadata: &'a TableDataMetadata,
expected_field_count: Option<usize>,
next_row_number: u64,
failed: bool,
rows: CopyRowReader<EntryDataReader<'a, R>>,
}
impl<'a, R: Read> TableRowReader<'a, R> {
pub(crate) fn new_with_limits(
data_id: DumpId,
metadata: &'a TableDataMetadata,
entry: EntryDataReader<'a, R>,
limits: Limits,
) -> Self {
Self::new_with_scan_limits(data_id, metadata, entry, limits, ScanLimits::unlimited())
}
pub(crate) fn new_with_scan_limits(
data_id: DumpId,
metadata: &'a TableDataMetadata,
entry: EntryDataReader<'a, R>,
limits: Limits,
scan_limits: ScanLimits,
) -> Self {
let expected_field_count = metadata.columns(data_id).ok().map(|columns| columns.len());
Self {
data_id,
metadata,
expected_field_count,
next_row_number: 1,
failed: false,
rows: CopyRowReader::with_limits_and_scan_limits(entry, limits, scan_limits),
}
}
pub fn columns(&self) -> Result<&[Column], PgDumpError> {
self.metadata.columns(self.data_id)
}
pub fn column_index(&self, name: &[u8]) -> Result<Option<usize>, PgDumpError> {
self.metadata.column_index(self.data_id, name)
}
pub fn next_row(&mut self) -> Result<Option<Row<'_>>, PgDumpError> {
if self.failed {
return Ok(None);
}
let data_id = self.data_id.as_i32();
let row_number = self.next_row_number;
let expected_field_count = self.expected_field_count;
let Some(row) = self.rows.next_row()? else {
return Ok(None);
};
if let Some(expected) = expected_field_count {
let actual = row.len();
if actual != expected {
self.failed = true;
return Err(PgDumpError::CopyRowFieldCountMismatch {
dump_id: data_id,
row: row_number,
expected: u64::try_from(expected).unwrap_or(u64::MAX),
actual: u64::try_from(actual).unwrap_or(u64::MAX),
});
}
}
self.next_row_number = row_number
.checked_add(1)
.ok_or(PgDumpError::CopyRowNumberOverflow { row: row_number })?;
Ok(Some(row))
}
pub fn find_first<F>(&mut self, predicate: F) -> Result<Option<OwnedRow>, PgDumpError>
where
F: FnMut(&Row<'_>) -> bool,
{
self.find_first_with_limits(ScanLimits::unlimited(), predicate)
}
pub fn find_first_with_limits<F>(
&mut self,
scan_limits: ScanLimits,
mut predicate: F,
) -> Result<Option<OwnedRow>, PgDumpError>
where
F: FnMut(&Row<'_>) -> bool,
{
if self.failed {
return Ok(None);
}
let expected_field_count = self.expected_field_count;
let mut next_row_number = self.next_row_number;
let mut mismatch = None;
let result = self.rows.find_first_with_limits(scan_limits, |row| {
let row_number = next_row_number;
next_row_number = next_row_number.saturating_add(1);
if let Some(expected) = expected_field_count {
let actual = row.len();
if actual != expected {
mismatch = Some((row_number, expected, actual));
return true;
}
}
predicate(row)
});
self.next_row_number = next_row_number;
if let Some((row, expected, actual)) = mismatch {
self.failed = true;
return Err(PgDumpError::CopyRowFieldCountMismatch {
dump_id: self.data_id.as_i32(),
row,
expected: u64::try_from(expected).unwrap_or(u64::MAX),
actual: u64::try_from(actual).unwrap_or(u64::MAX),
});
}
result
}
pub fn find_first_equal(
&mut self,
column: &[u8],
expected: FieldRef<'_>,
) -> Result<ColumnEqualityResult, PgDumpError> {
self.find_first_equal_with_limits(ScanLimits::unlimited(), column, expected)
}
pub fn find_first_equal_with_limits(
&mut self,
scan_limits: ScanLimits,
column: &[u8],
expected: FieldRef<'_>,
) -> Result<ColumnEqualityResult, PgDumpError> {
let Some(column_index) = self.column_index(column)? else {
return Ok(ColumnEqualityResult::ColumnNotFound);
};
match self
.find_first_with_limits(scan_limits, |row| row.field(column_index) == Some(expected))?
{
Some(row) => Ok(ColumnEqualityResult::Match(row)),
None => Ok(ColumnEqualityResult::NoMatch),
}
}
#[cfg(test)]
pub(crate) const fn consumed_input_bytes(&self) -> u64 {
self.rows.consumed_input_bytes()
}
}