use std::collections::hash_map::{Entry, HashMap};
use std::fmt;
use std::fs;
use std::io;
use std::str;
use byteorder::{BigEndian, WriteBytesExt};
use csv;
use crate::config::{Config, Delimiter};
use crate::index::Indexed;
use crate::select::{SelectColumns, Selection};
use crate::CliResult;
use clap::Parser;
type ByteString = Vec<u8>;
#[derive(Parser, Debug)]
pub struct Args {
#[arg()]
pub arg_columns1: SelectColumns,
#[arg()]
pub arg_input1: String,
#[arg()]
pub arg_columns2: SelectColumns,
#[arg()]
pub arg_input2: String,
#[arg(long = "left")]
pub flag_left: bool,
#[arg(long = "right")]
pub flag_right: bool,
#[arg(long = "full")]
pub flag_full: bool,
#[arg(long = "cross")]
pub flag_cross: bool,
#[arg(short = 'o', long = "output", value_name = "file")]
pub flag_output: Option<String>,
#[arg(short = 'n', long = "no-headers")]
pub flag_no_headers: bool,
#[arg(long = "no-case")]
pub flag_no_case: bool,
#[arg(long = "nulls")]
pub flag_nulls: bool,
#[arg(short = 'd', long = "delimiter", value_name = "arg")]
pub flag_delimiter: Option<Delimiter>,
}
pub fn run(args: &Args) -> CliResult<()> {
let mut state = args.new_io_state()?;
match (
args.flag_left,
args.flag_right,
args.flag_full,
args.flag_cross,
) {
(true, false, false, false) => {
state.write_headers()?;
state.outer_join(false)
}
(false, true, false, false) => {
state.write_headers()?;
state.outer_join(true)
}
(false, false, true, false) => {
state.write_headers()?;
state.full_outer_join()
}
(false, false, false, true) => {
state.write_headers()?;
state.cross_join()
}
(false, false, false, false) => {
state.write_headers()?;
state.inner_join()
}
_ => fail!("Please pick exactly one join operation."),
}
}
struct IoState<R, W: io::Write> {
wtr: csv::Writer<W>,
rdr1: csv::Reader<R>,
sel1: Selection,
rdr2: csv::Reader<R>,
sel2: Selection,
no_headers: bool,
casei: bool,
nulls: bool,
}
impl<R: io::Read + io::Seek, W: io::Write> IoState<R, W> {
fn write_headers(&mut self) -> CliResult<()> {
if !self.no_headers {
let mut headers = self.rdr1.byte_headers()?.clone();
headers.extend(self.rdr2.byte_headers()?.iter());
self.wtr.write_record(&headers)?;
}
Ok(())
}
fn inner_join(mut self) -> CliResult<()> {
let mut scratch = csv::ByteRecord::new();
let mut validx = ValueIndex::new(self.rdr2, &self.sel2, self.casei, self.nulls)?;
for row in self.rdr1.byte_records() {
let row = row?;
let key = get_row_key(&self.sel1, &row, self.casei);
match validx.values.get(&key) {
None => continue,
Some(rows) => {
for &rowi in rows.iter() {
validx.idx.seek(rowi as u64)?;
validx.idx.read_byte_record(&mut scratch)?;
let combined = row.iter().chain(scratch.iter());
self.wtr.write_record(combined)?;
}
}
}
}
Ok(())
}
fn outer_join(mut self, right: bool) -> CliResult<()> {
if right {
::std::mem::swap(&mut self.rdr1, &mut self.rdr2);
::std::mem::swap(&mut self.sel1, &mut self.sel2);
}
let mut scratch = csv::ByteRecord::new();
let (_, pad2) = self.get_padding()?;
let mut validx = ValueIndex::new(self.rdr2, &self.sel2, self.casei, self.nulls)?;
for row in self.rdr1.byte_records() {
let row = row?;
let key = get_row_key(&self.sel1, &row, self.casei);
match validx.values.get(&key) {
None => {
if right {
self.wtr.write_record(pad2.iter().chain(&row))?;
} else {
self.wtr.write_record(row.iter().chain(&pad2))?;
}
}
Some(rows) => {
for &rowi in rows.iter() {
validx.idx.seek(rowi as u64)?;
let row1 = row.iter();
validx.idx.read_byte_record(&mut scratch)?;
if right {
self.wtr.write_record(scratch.iter().chain(row1))?;
} else {
self.wtr.write_record(row1.chain(&scratch))?;
}
}
}
}
}
Ok(())
}
fn full_outer_join(mut self) -> CliResult<()> {
let mut scratch = csv::ByteRecord::new();
let (pad1, pad2) = self.get_padding()?;
let mut validx = ValueIndex::new(self.rdr2, &self.sel2, self.casei, self.nulls)?;
let mut rdr2_written: Vec<_> = std::iter::repeat_n(false, validx.num_rows).collect();
for row1 in self.rdr1.byte_records() {
let row1 = row1?;
let key = get_row_key(&self.sel1, &row1, self.casei);
match validx.values.get(&key) {
None => {
self.wtr.write_record(row1.iter().chain(&pad2))?;
}
Some(rows) => {
for &rowi in rows.iter() {
rdr2_written[rowi] = true;
validx.idx.seek(rowi as u64)?;
validx.idx.read_byte_record(&mut scratch)?;
self.wtr.write_record(row1.iter().chain(&scratch))?;
}
}
}
}
for (i, &written) in rdr2_written.iter().enumerate() {
if !written {
validx.idx.seek(i as u64)?;
validx.idx.read_byte_record(&mut scratch)?;
self.wtr.write_record(pad1.iter().chain(&scratch))?;
}
}
Ok(())
}
fn cross_join(mut self) -> CliResult<()> {
let mut pos = csv::Position::new();
pos.set_byte(0);
let mut row2 = csv::ByteRecord::new();
for row1 in self.rdr1.byte_records() {
let row1 = row1?;
self.rdr2.seek(pos.clone())?;
if self.rdr2.has_headers() {
self.rdr2.read_byte_record(&mut row2)?;
}
while self.rdr2.read_byte_record(&mut row2)? {
self.wtr.write_record(row1.iter().chain(&row2))?;
}
}
Ok(())
}
fn get_padding(&mut self) -> CliResult<(csv::ByteRecord, csv::ByteRecord)> {
let len1 = self.rdr1.byte_headers()?.len();
let len2 = self.rdr2.byte_headers()?.len();
Ok((
std::iter::repeat_n(b"", len1).collect(),
std::iter::repeat_n(b"", len2).collect(),
))
}
}
impl Args {
fn new_io_state(&self) -> CliResult<IoState<fs::File, Box<dyn io::Write + 'static>>> {
let rconf1 = Config::new(&Some(self.arg_input1.clone()))
.delimiter(self.flag_delimiter)
.no_headers(self.flag_no_headers)
.select(self.arg_columns1.clone());
let rconf2 = Config::new(&Some(self.arg_input2.clone()))
.delimiter(self.flag_delimiter)
.no_headers(self.flag_no_headers)
.select(self.arg_columns2.clone());
let mut rdr1 = rconf1.reader_file()?;
let mut rdr2 = rconf2.reader_file()?;
let (sel1, sel2) = self.get_selections(&rconf1, &mut rdr1, &rconf2, &mut rdr2)?;
Ok(IoState {
wtr: Config::new(&self.flag_output).writer()?,
rdr1,
sel1,
rdr2,
sel2,
no_headers: rconf1.no_headers,
casei: self.flag_no_case,
nulls: self.flag_nulls,
})
}
fn get_selections<R: io::Read>(
&self,
rconf1: &Config,
rdr1: &mut csv::Reader<R>,
rconf2: &Config,
rdr2: &mut csv::Reader<R>,
) -> CliResult<(Selection, Selection)> {
let headers1 = rdr1.byte_headers()?;
let headers2 = rdr2.byte_headers()?;
let select1 = rconf1.selection(headers1)?;
let select2 = rconf2.selection(headers2)?;
if select1.len() != select2.len() {
return fail!(format!(
"Column selections must have the same number of columns, \
but found column selections with {} and {} columns.",
select1.len(),
select2.len()
));
}
Ok((select1, select2))
}
}
struct ValueIndex<R> {
values: HashMap<Vec<ByteString>, Vec<usize>>,
idx: Indexed<R, io::Cursor<Vec<u8>>>,
num_rows: usize,
}
impl<R: io::Read + io::Seek> ValueIndex<R> {
fn new(
mut rdr: csv::Reader<R>,
sel: &Selection,
casei: bool,
nulls: bool,
) -> CliResult<ValueIndex<R>> {
let mut val_idx = HashMap::with_capacity(10000);
let mut row_idx = io::Cursor::new(Vec::with_capacity(8 * 10000));
let (mut rowi, mut count) = (0usize, 0usize);
if !rdr.has_headers() {
let mut pos = csv::Position::new();
pos.set_byte(0);
rdr.seek(pos)?;
} else {
rdr.byte_headers()?;
row_idx.write_u64::<BigEndian>(0)?;
count += 1;
}
let mut row = csv::ByteRecord::new();
while rdr.read_byte_record(&mut row)? {
row_idx.write_u64::<BigEndian>(row.position().unwrap().byte())?;
let fields: Vec<_> = sel.select(&row).map(|v| transform(v, casei)).collect();
if nulls || !fields.iter().any(|f| f.is_empty()) {
match val_idx.entry(fields) {
Entry::Vacant(v) => {
let mut rows = Vec::with_capacity(4);
rows.push(rowi);
v.insert(rows);
}
Entry::Occupied(mut v) => {
v.get_mut().push(rowi);
}
}
}
rowi += 1;
count += 1;
}
row_idx.write_u64::<BigEndian>(count as u64)?;
let idx = Indexed::open(rdr, io::Cursor::new(row_idx.into_inner()))?;
Ok(ValueIndex {
values: val_idx,
idx,
num_rows: rowi,
})
}
}
impl<R> fmt::Debug for ValueIndex<R> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let mut kvs = self.values.iter().collect::<Vec<_>>();
kvs.sort_by(|&(_, v1), &(_, v2)| v1[0].cmp(&v2[0]));
for (keys, rows) in kvs.into_iter() {
let keys = keys
.iter()
.map(|k| String::from_utf8(k.to_vec()).unwrap())
.collect::<Vec<_>>();
writeln!(f, "({}) => {:?}", keys.join(", "), rows)?
}
Ok(())
}
}
fn get_row_key(sel: &Selection, row: &csv::ByteRecord, casei: bool) -> Vec<ByteString> {
sel.select(row).map(|v| transform(v, casei)).collect()
}
fn transform(bs: &[u8], casei: bool) -> ByteString {
match str::from_utf8(bs) {
Err(_) => bs.to_vec(),
Ok(s) => {
if !casei {
s.trim().as_bytes().to_vec()
} else {
let norm: String = s
.trim()
.chars()
.map(|c| c.to_lowercase().next().unwrap())
.collect();
norm.into_bytes()
}
}
}
}