Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 31 additions & 0 deletions integration/rust/tests/integration/limit.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,37 @@ async fn limit_across_shards() -> Result<(), Box<dyn std::error::Error>> {
)
.await?;

let rows = sharded
.fetch_all("SELECT id FROM limit_test ORDER BY value")
.await?;
assert_eq!(
rows.iter()
.map(|row| {
assert_eq!(row.len(), 1);
row.get::<i32, _>("id")
})
.collect::<Vec<_>>(),
vec![1, 2, 3, 4, 5, 1, 2, 3, 4, 5]
);

let rows = sharded
.fetch_all("/* pgdog_shard: 0 */ SELECT id FROM limit_test ORDER BY value")
.await?;
assert_eq!(
rows.iter()
.map(|row| {
assert_eq!(row.len(), 1);
row.get::<i32, _>("id")
})
.collect::<Vec<_>>(),
vec![1, 2, 3, 4, 5]
);

let row = sharded
.fetch_one("/* pgdog_shard: 0 */ SELECT stddev(value) FROM limit_test")
.await?;
assert_eq!(row.len(), 1);

// LIMIT 5
let rows = sharded
.fetch_all("SELECT value FROM limit_test ORDER BY value LIMIT 5")
Expand Down
32 changes: 16 additions & 16 deletions pgdog/src/backend/pool/connection/aggregate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ use std::mem;
use crate::{
frontend::router::parser::{
Aggregate, AggregateFunction, AggregateTarget,
rewrite::statement::aggregate::{AggregateRewritePlan, HelperKind},
rewrite::statement::{aggregate::HelperKind, projection::ProjectionRewritePlan},
},
net::{
Decoder,
Expand Down Expand Up @@ -237,37 +237,37 @@ pub(super) struct Aggregates<'a> {
mappings: HashMap<Grouping, GroupState<'a>>,
decoder: &'a Decoder,
aggregate: &'a Aggregate,
helper_columns: HashMap<usize, HelperColumns>,
projected_columns: HashMap<usize, HelperColumns>,
}

impl<'a> Aggregates<'a> {
pub(super) fn new(
rows: &'a VecDeque<DataRow>,
decoder: &'a Decoder,
aggregate: &'a Aggregate,
plan: &AggregateRewritePlan,
plan: &ProjectionRewritePlan,
) -> Option<Self> {
let mut helper_columns: HashMap<usize, HelperColumns> = HashMap::new();
let mut projected_columns: HashMap<usize, HelperColumns> = HashMap::new();

for target in aggregate.targets() {
let key = target.column();
match target.function() {
AggregateFunction::Count => {
helper_columns.entry(key).or_default().count = Some(target.column());
projected_columns.entry(key).or_default().count = Some(target.column());
}
AggregateFunction::Sum => {
helper_columns.entry(key).or_default().sum = Some(target.column());
projected_columns.entry(key).or_default().sum = Some(target.column());
}
_ => {}
}
}

for helper in plan.helpers() {
for helper in plan.aggregate_helpers() {
let Some(index) = decoder.row_description().field_index(&helper.alias) else {
continue;
};

let entry = helper_columns.entry(helper.target_column).or_default();
let entry = projected_columns.entry(helper.target_column).or_default();
match helper.kind {
HelperKind::Count => entry.count = Some(index),
HelperKind::Sum => entry.sum = Some(index),
Expand All @@ -278,14 +278,14 @@ impl<'a> Aggregates<'a> {
let helpers_present = aggregate.targets().iter().all(|target| {
let key = target.column();
match target.function() {
AggregateFunction::Avg => helper_columns
AggregateFunction::Avg => projected_columns
.get(&key)
.and_then(|columns| columns.count)
.is_some(),
AggregateFunction::StddevPop
| AggregateFunction::StddevSamp
| AggregateFunction::VarPop
| AggregateFunction::VarSamp => helper_columns
| AggregateFunction::VarSamp => projected_columns
.get(&key)
.map(|columns| {
columns.count.is_some() && columns.sum.is_some() && columns.sumsq.is_some()
Expand All @@ -305,7 +305,7 @@ impl<'a> Aggregates<'a> {
decoder,
mappings: HashMap::new(),
aggregate,
helper_columns,
projected_columns,
})
} else {
None
Expand Down Expand Up @@ -333,7 +333,7 @@ impl<'a> Aggregates<'a> {
Entry::Occupied(o) => o.into_mut(),
Entry::Vacant(v) => {
let accumulators =
Accumulator::from_aggregate(self.aggregate, &self.helper_columns)?;
Accumulator::from_aggregate(self.aggregate, &self.projected_columns)?;

// Gather all col vals corresponding to list of passthrough indices.
let mut passthrough = Vec::new();
Expand Down Expand Up @@ -523,7 +523,7 @@ mod test {
shard1.add("3");
rows.push_back(shard1);

let plan = AggregateRewritePlan::default();
let plan = ProjectionRewritePlan::default();
let mut result = Aggregates::new(&rows, &decoder, &aggregate, &plan)
.unwrap()
.aggregate()
Expand Down Expand Up @@ -557,7 +557,7 @@ mod test {
&rows,
&decoder,
&aggregate,
&AggregateRewritePlan::default(),
&ProjectionRewritePlan::default(),
)
.unwrap()
.aggregate()
Expand Down Expand Up @@ -602,7 +602,7 @@ mod test {
&rows,
&decoder,
&aggregate,
&AggregateRewritePlan::default(),
&ProjectionRewritePlan::default(),
)
.unwrap()
.aggregate()
Expand Down Expand Up @@ -649,7 +649,7 @@ mod test {
&rows,
&decoder,
&aggregate,
&AggregateRewritePlan::default(),
&ProjectionRewritePlan::default(),
)
.unwrap()
.aggregate()
Expand Down
47 changes: 39 additions & 8 deletions pgdog/src/backend/pool/connection/buffer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ use std::{
use crate::{
frontend::router::parser::{
Aggregate, DistinctBy, DistinctColumn, Limit, OrderBy,
rewrite::statement::aggregate::AggregateRewritePlan,
rewrite::statement::projection::ProjectionRewritePlan,
},
net::{
Decoder,
Expand Down Expand Up @@ -140,31 +140,30 @@ impl Buffer {
&mut self,
aggregate: &Aggregate,
decoder: &Decoder,
plan: &AggregateRewritePlan,
plan: &ProjectionRewritePlan,
) -> Result<(), super::Error> {
let buffer: VecDeque<DataRow> = std::mem::take(&mut self.buffer);
let mut rows = if aggregate.is_empty() {
let rows = if aggregate.is_empty() {
buffer
} else if let Some(aggregates) = Aggregates::new(&buffer, decoder, aggregate, plan) {
aggregates.aggregate()?
} else {
buffer
};

Self::drop_helper_columns(&mut rows, plan);
self.buffer = rows;

Ok(())
}

fn drop_helper_columns(rows: &mut VecDeque<DataRow>, plan: &AggregateRewritePlan) {
pub(super) fn drop_columns(&mut self, plan: &ProjectionRewritePlan) {
if plan.is_noop() {
return;
}

let drop = plan.drop_columns().collect();

for row in rows.iter_mut() {
for row in self.buffer.iter_mut() {
row.drop_columns(&drop);
}
}
Expand Down Expand Up @@ -240,6 +239,7 @@ impl Buffer {
#[cfg(test)]
mod test {
use super::*;
use crate::frontend::router::parser::rewrite::statement::projection::OrderByHelper;
use crate::net::{Datum, Field, Format, RowDescription};
use bytes::Bytes;

Expand Down Expand Up @@ -273,6 +273,37 @@ mod test {
assert_eq!(i, 26);
}

#[test]
fn test_sort_by_hidden_column_before_dropping_it() {
let mut buf = Buffer::default();
let rd = RowDescription::new(&[Field::bigint("id"), Field::text("__pgdog_order_by_0")]);
let decoder = Decoder::from(rd);
let mut plan = ProjectionRewritePlan::default();
plan.add_order_by_helper(OrderByHelper {
sort_position: 0,
projected_column: 1,
});

for (id, name) in [(1_i64, "z"), (2, "a"), (3, "m")] {
let mut row = DataRow::new();
row.add(id).add(name);
buf.add(row.message()).unwrap();
}

buf.sort(&[OrderBy::Asc(2)], &decoder);
buf.drop_columns(&plan);
buf.mark_full();

let ids = std::iter::from_fn(|| buf.take())
.map(|message| {
let row = DataRow::from_bytes(message.to_bytes()).unwrap();
assert_eq!(row.len(), 1);
row.get::<i64>(0, Format::Text).unwrap()
})
.collect::<Vec<_>>();
assert_eq!(ids, [2, 3, 1]);
}

#[test]
fn test_aggregate_buffer() {
let mut buf = Buffer::default();
Expand All @@ -285,7 +316,7 @@ mod test {
buf.add(dr.message()).unwrap();
}

buf.aggregate(&agg, &Decoder::from(rd), &AggregateRewritePlan::default())
buf.aggregate(&agg, &Decoder::from(rd), &ProjectionRewritePlan::default())
.unwrap();
buf.mark_full();

Expand All @@ -312,7 +343,7 @@ mod test {
}
}

buf.aggregate(&agg, &Decoder::from(rd), &AggregateRewritePlan::default())
buf.aggregate(&agg, &Decoder::from(rd), &ProjectionRewritePlan::default())
.unwrap();
buf.mark_full();

Expand Down
22 changes: 19 additions & 3 deletions pgdog/src/backend/pool/connection/multi_shard/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -264,13 +264,15 @@ impl MultiShard {
.aggregate(
self.route.aggregate(),
&self.decoder,
self.route.aggregate_rewrite_plan(),
self.route.projection_rewrite_plan(),
)
.map_err(Error::from)?;

self.buffer.sort(self.route.order_by(), &self.decoder);
self.buffer.distinct(self.route.distinct(), &self.decoder);
self.buffer.limit(self.route.limit());
self.buffer
.drop_columns(self.route.projection_rewrite_plan());
}

if has_rows {
Expand Down Expand Up @@ -311,7 +313,7 @@ impl MultiShard {
{
// Only send it to the client once all shards sent it,
// so we don't get early requests from clients.
let plan = self.route.aggregate_rewrite_plan();
let plan = self.route.projection_rewrite_plan();
if plan.is_noop() {
forward = Some(message);
} else {
Expand Down Expand Up @@ -342,7 +344,7 @@ impl MultiShard {
)
}

fn handle_data_row(&mut self, message: Message) -> Result<Option<Message>, Error> {
fn handle_data_row(&mut self, mut message: Message) -> Result<Option<Message>, Error> {
let mut forward = None;

if self.shards > 1 {
Expand All @@ -366,9 +368,11 @@ impl MultiShard {
{
if self.route.is_omnisharded() {
if self.request_state.first_backend_data == message.source().backend_id() {
self.drop_columns(&mut message)?;
forward = Some(message);
}
} else {
self.drop_columns(&mut message)?;
forward = Some(message);
}
} else {
Expand All @@ -378,6 +382,18 @@ impl MultiShard {
Ok(forward)
}

fn drop_columns(&self, message: &mut Message) -> Result<(), Error> {
let plan = self.route.projection_rewrite_plan();
if plan.is_noop() {
return Ok(());
}

let mut row = DataRow::from_bytes(message.to_bytes())?;
row.drop_columns(&plan.drop_columns().collect());
message.replace_payload(row.to_bytes());
Ok(())
}

fn handle_passthrough(message: Message, counter: &mut usize, shards: usize) -> Option<Message> {
*counter += 1;

Expand Down
Loading
Loading