use std::io::{ErrorKind, Read};
use serde::Deserialize;
use ytsaurus_yson::{Scan, YsonFormat, YsonNode, YsonValue, from_slice, scan::scan_value};
use crate::error::{JobError, Result};
const DEFAULT_BUFFER_BYTES: usize = 1024 * 1024;
const DEFAULT_MAX_RECORD_BYTES: usize = 256 * 1024 * 1024;
#[derive(Debug)]
pub enum Event<'a> {
Row(Row<'a>),
KeySwitch,
}
#[derive(Debug, Clone, Copy)]
pub struct Row<'a> {
pub table_index: i64,
pub row_index: Option<i64>,
pub range_index: Option<i64>,
bytes: &'a [u8],
format: YsonFormat,
offset: u64,
}
impl<'a> Row<'a> {
#[must_use]
pub fn raw(&self) -> &'a [u8] {
self.bytes
}
#[must_use]
pub fn offset(&self) -> u64 {
self.offset
}
pub fn parse<T: Deserialize<'a>>(&self) -> Result<T> {
from_slice(self.bytes, self.format).map_err(|source| JobError::Yson {
offset: self.offset,
source,
})
}
pub fn value(&self) -> Result<YsonValue> {
self.parse()
}
}
#[derive(Debug, Clone, Copy)]
enum Pending {
Row { len: usize },
KeySwitch { len: usize },
}
#[derive(Debug)]
pub struct JobReader<R> {
input: R,
format: YsonFormat,
buf: Vec<u8>,
pos: usize,
filled: usize,
base_offset: u64,
input_done: bool,
pending: Option<Pending>,
max_record_bytes: usize,
table_index: i64,
row_index: Option<i64>,
range_index: Option<i64>,
}
impl JobReader<std::io::Stdin> {
#[must_use]
pub fn from_stdin() -> Self {
Self::binary(std::io::stdin())
}
}
impl<R: Read> JobReader<R> {
#[must_use]
pub fn binary(input: R) -> Self {
Self::with_format(input, YsonFormat::Binary)
}
#[must_use]
pub fn text(input: R) -> Self {
Self::with_format(input, YsonFormat::Text)
}
#[must_use]
pub fn with_format(input: R, format: YsonFormat) -> Self {
Self {
input,
format,
buf: vec![0; DEFAULT_BUFFER_BYTES],
pos: 0,
filled: 0,
base_offset: 0,
input_done: false,
pending: None,
max_record_bytes: DEFAULT_MAX_RECORD_BYTES,
table_index: 0,
row_index: None,
range_index: None,
}
}
#[must_use]
pub fn with_buffer_size(mut self, bytes: usize) -> Self {
self.buf = vec![0; bytes.max(64)];
self
}
#[must_use]
pub fn with_max_record_bytes(mut self, bytes: usize) -> Self {
self.max_record_bytes = bytes;
self
}
pub fn next_event(&mut self) -> Result<Option<Event<'_>>> {
let Some(pending) = self.ensure_pending()? else {
return Ok(None);
};
match pending {
Pending::KeySwitch { len } => {
self.consume(len);
Ok(Some(Event::KeySwitch))
}
Pending::Row { len } => {
let start = self.pos;
let offset = self.base_offset + start as u64;
self.consume(len);
Ok(Some(Event::Row(Row {
table_index: self.table_index,
row_index: self.row_index,
range_index: self.range_index,
bytes: &self.buf[start..start + len],
format: self.format,
offset,
})))
}
}
}
pub fn groups(&mut self) -> Groups<'_, R> {
Groups {
reader: self,
in_group: false,
key_columns: Vec::new(),
}
}
pub fn groups_by<I>(&mut self, columns: I) -> Groups<'_, R>
where
I: IntoIterator,
I::Item: AsRef<str>,
{
Groups {
reader: self,
in_group: false,
key_columns: columns
.into_iter()
.map(|c| c.as_ref().as_bytes().to_vec())
.collect(),
}
}
fn decode_key(&mut self, len: usize, columns: &[Vec<u8>]) -> Result<GroupKey> {
let start = self.pos;
let offset = self.base_offset + start as u64;
let record = &self.buf[start..start + len];
let value: YsonValue =
from_slice(record, self.format).map_err(|source| JobError::Yson { offset, source })?;
let YsonNode::Map(fields) = &value.node else {
return Err(JobError::Yson {
offset,
source: ytsaurus_yson::YsonError::Custom(
"a reduce key can only be read from a row that is a map".to_owned(),
),
});
};
let mut decoded = Vec::with_capacity(columns.len());
for name in columns {
if let Some(v) = fields.get(name) {
decoded.push((name.clone(), v.clone()));
}
}
Ok(GroupKey { columns: decoded })
}
fn consume(&mut self, len: usize) {
self.pos += len;
self.pending = None;
}
fn discard_pending(&mut self) {
if let Some(Pending::Row { len } | Pending::KeySwitch { len }) = self.pending {
self.consume(len);
}
}
fn ensure_pending(&mut self) -> Result<Option<Pending>> {
if let Some(pending) = self.pending {
return Ok(Some(pending));
}
loop {
let Some(len) = self.next_record_len()? else {
return Ok(None);
};
let offset = self.base_offset + self.pos as u64;
let record = &self.buf[self.pos..self.pos + len];
let classified = if first_significant_byte(record, self.format) == Some(b'<') {
classify_attributed(record, self.format, offset)?
} else {
Classified::Row
};
match classified {
Classified::Row => {
let pending = Pending::Row { len };
self.pending = Some(pending);
return Ok(Some(pending));
}
Classified::KeySwitch => {
let pending = Pending::KeySwitch { len };
self.pending = Some(pending);
return Ok(Some(pending));
}
Classified::TableIndex(i) => {
self.table_index = i;
self.row_index = None;
self.pos += len;
}
Classified::RowIndex(i) => {
self.row_index = Some(i);
self.pos += len;
}
Classified::RangeIndex(i) => {
self.range_index = Some(i);
self.pos += len;
}
Classified::Skip => self.pos += len,
}
}
}
fn next_record_len(&mut self) -> Result<Option<usize>> {
loop {
self.skip_separators();
if self.pos < self.filled {
match scan_value(&self.buf[self.pos..self.filled], self.format) {
Ok(Scan::Complete { len }) => return Ok(Some(len)),
Ok(Scan::Incomplete) => {}
Err(source) => {
return Err(JobError::Yson {
offset: self.base_offset + self.pos as u64,
source,
});
}
}
}
if self.input_done {
return if self.pos == self.filled {
Ok(None)
} else {
Err(JobError::TruncatedRecord {
offset: self.base_offset + self.pos as u64,
buffered: self.filled - self.pos,
})
};
}
self.fill()?;
}
}
fn skip_separators(&mut self) {
while self.pos < self.filled {
match self.buf[self.pos] {
b';' => self.pos += 1,
b if b.is_ascii_whitespace() => self.pos += 1,
_ => break,
}
}
}
fn fill(&mut self) -> Result<()> {
if self.pos > 0 {
self.buf.copy_within(self.pos..self.filled, 0);
self.filled -= self.pos;
self.base_offset += self.pos as u64;
self.pos = 0;
}
if self.filled == self.buf.len() {
let new_len = self.buf.len().saturating_mul(2);
if self.buf.len() >= self.max_record_bytes {
return Err(JobError::RecordTooLarge {
offset: self.base_offset,
limit: self.max_record_bytes,
});
}
self.buf.resize(new_len.min(self.max_record_bytes), 0);
}
loop {
match self.input.read(&mut self.buf[self.filled..]) {
Ok(0) => {
self.input_done = true;
return Ok(());
}
Ok(n) => {
self.filled += n;
return Ok(());
}
Err(e) if e.kind() == ErrorKind::Interrupted => {}
Err(e) => return Err(JobError::Read(e)),
}
}
}
}
enum Classified {
Row,
Skip,
TableIndex(i64),
RowIndex(i64),
RangeIndex(i64),
KeySwitch,
}
fn first_significant_byte(record: &[u8], format: YsonFormat) -> Option<u8> {
match format {
YsonFormat::Binary => record.first().copied(),
YsonFormat::Text => record.iter().find(|b| !b.is_ascii_whitespace()).copied(),
}
}
fn classify_attributed(record: &[u8], format: YsonFormat, offset: u64) -> Result<Classified> {
let value: YsonValue =
from_slice(record, format).map_err(|source| JobError::Yson { offset, source })?;
if !matches!(value.node, YsonNode::Entity) {
return Ok(Classified::Row);
}
let Some(attributes) = value.attributes.as_ref() else {
return Ok(Classified::Skip);
};
let as_i64 = |name: &str, v: &YsonValue| -> Result<i64> {
v.as_i64().ok_or_else(|| JobError::BadControlRecord {
offset,
reason: format!("{name} must be an int64, got {:?}", v.node),
})
};
for (key, v) in attributes {
match key.as_slice() {
b"key_switch" => {
return match v.node {
YsonNode::Boolean(true) => Ok(Classified::KeySwitch),
YsonNode::Boolean(false) => Ok(Classified::Skip),
ref other => Err(JobError::BadControlRecord {
offset,
reason: format!("key_switch must be a boolean, got {other:?}"),
}),
};
}
b"table_index" => return Ok(Classified::TableIndex(as_i64("table_index", v)?)),
b"row_index" => return Ok(Classified::RowIndex(as_i64("row_index", v)?)),
b"range_index" => return Ok(Classified::RangeIndex(as_i64("range_index", v)?)),
_ => {}
}
}
Ok(Classified::Skip)
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct GroupKey {
columns: Vec<(Vec<u8>, YsonValue)>,
}
impl GroupKey {
#[must_use]
pub fn get(&self, name: &str) -> Option<&YsonValue> {
self.columns
.iter()
.find(|(k, _)| k == name.as_bytes())
.map(|(_, v)| v)
}
#[must_use]
pub fn bytes(&self, name: &str) -> Option<&[u8]> {
match &self.get(name)?.node {
YsonNode::String(bytes) => Some(bytes),
_ => None,
}
}
#[must_use]
pub fn str(&self, name: &str) -> Option<&str> {
std::str::from_utf8(self.bytes(name)?).ok()
}
#[must_use]
pub fn i64(&self, name: &str) -> Option<i64> {
self.get(name)?.as_i64()
}
#[must_use]
pub fn columns(&self) -> &[(Vec<u8>, YsonValue)] {
&self.columns
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.columns.is_empty()
}
}
#[derive(Debug)]
pub struct Groups<'r, R> {
reader: &'r mut JobReader<R>,
in_group: bool,
key_columns: Vec<Vec<u8>>,
}
impl<R: Read> Groups<'_, R> {
pub fn next_group(&mut self) -> Result<Option<Group<'_, R>>> {
if self.in_group {
loop {
match self.reader.ensure_pending()? {
None => {
self.in_group = false;
return Ok(None);
}
Some(Pending::KeySwitch { .. }) => {
self.reader.discard_pending();
break;
}
Some(Pending::Row { .. }) => self.reader.discard_pending(),
}
}
}
match self.reader.ensure_pending()? {
None => {
self.in_group = false;
Ok(None)
}
Some(Pending::KeySwitch { .. }) => {
self.reader.discard_pending();
self.in_group = true;
Ok(Some(Group {
reader: self.reader,
done: false,
key: GroupKey::default(),
}))
}
Some(Pending::Row { len }) => {
let key = if self.key_columns.is_empty() {
GroupKey::default()
} else {
self.reader.decode_key(len, &self.key_columns)?
};
self.in_group = true;
Ok(Some(Group {
reader: self.reader,
done: false,
key,
}))
}
}
}
}
#[derive(Debug)]
pub struct Group<'g, R> {
reader: &'g mut JobReader<R>,
done: bool,
key: GroupKey,
}
impl<R: Read> Group<'_, R> {
#[must_use]
pub fn key(&self) -> &GroupKey {
&self.key
}
pub fn next_row(&mut self) -> Result<Option<Row<'_>>> {
if self.done {
return Ok(None);
}
match self.reader.ensure_pending()? {
None | Some(Pending::KeySwitch { .. }) => {
self.done = true;
Ok(None)
}
Some(Pending::Row { len }) => {
let start = self.reader.pos;
let offset = self.reader.base_offset + start as u64;
self.reader.consume(len);
Ok(Some(Row {
table_index: self.reader.table_index,
row_index: self.reader.row_index,
range_index: self.reader.range_index,
bytes: &self.reader.buf[start..start + len],
format: self.reader.format,
offset,
}))
}
}
}
}