use std::sync::Arc;
use pgwire::api::results::{DataRowEncoder, FieldFormat, FieldInfo, QueryResponse, Response};
use pgwire::error::{ErrorInfo, PgWireError, PgWireResult};
use crate::control::server::response_shape::compose::shape_decoded_rows;
use crate::control::server::response_shape::schema::OutputSchema;
use crate::control::server::response_shape::types::DdlColType;
use crate::control::server::result_stream::ResultStream;
use crate::data::executor::response_codec::decode_payload_to_json;
use super::super::ddl_encode::col_type_to_field_with_format;
use super::super::types::{error_to_sqlstate, text_field};
use super::shape_encode::{encode_shaped_row, shaped_query_response};
pub(crate) fn streaming_multirow_response(stream: ResultStream, limit: usize) -> Response {
use futures::StreamExt;
let schema = Arc::new(vec![text_field("result")]);
let row_schema = schema.clone();
let row_stream = async_stream::try_stream! {
let mut emitted: usize = 0;
let mut batches = stream;
while let Some(batch) = batches.next().await {
let batch = batch.map_err(|e| {
let (severity, code, message) = error_to_sqlstate(&e);
PgWireError::UserError(Box::new(ErrorInfo::new(
severity.to_owned(),
code.to_owned(),
message,
)))
})?;
if emitted >= limit {
break;
}
let text = decode_payload_to_json(&batch.payload);
if let Ok(serde_json::Value::Array(items)) =
sonic_rs::from_str::<serde_json::Value>(&text)
{
for item in items {
if emitted >= limit {
break;
}
let mut encoder = DataRowEncoder::new(row_schema.clone());
encoder.encode_field(&item.to_string()).map_err(|e| {
PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"XX000".to_owned(),
format!("failed to encode streamed row: {e}"),
)))
})?;
emitted += 1;
yield encoder.take_row();
}
}
}
};
Response::Query(QueryResponse::new(schema, row_stream))
}
pub(crate) fn streaming_shaped_response(
stream: ResultStream,
limit: usize,
schema_out: OutputSchema,
formats: &[FieldFormat],
) -> Response {
use futures::StreamExt;
let display_columns: Vec<String> = schema_out
.columns
.iter()
.map(|c| c.display_name.clone())
.collect();
let column_types: Vec<DdlColType> = schema_out.columns.iter().map(|c| c.ty).collect();
let row_formats: Vec<FieldFormat> = schema_out
.columns
.iter()
.enumerate()
.map(|(i, _)| formats.get(i).copied().unwrap_or(FieldFormat::Text))
.collect();
let fields: Vec<FieldInfo> = schema_out
.columns
.iter()
.enumerate()
.map(|(i, c)| col_type_to_field_with_format(&c.display_name, c.ty, row_formats[i]))
.collect();
let schema = Arc::new(fields);
let row_schema = schema.clone();
let row_stream = async_stream::try_stream! {
let mut emitted: usize = 0;
let mut batches = stream;
while let Some(batch) = batches.next().await {
let batch = batch.map_err(|e| {
let (severity, code, message) = error_to_sqlstate(&e);
PgWireError::UserError(Box::new(ErrorInfo::new(
severity.to_owned(),
code.to_owned(),
message,
)))
})?;
if emitted >= limit {
break;
}
let text = decode_payload_to_json(&batch.payload);
let value = sonic_rs::from_str::<serde_json::Value>(&text).map_err(|e| {
PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"XX000".to_owned(),
format!("failed to decode streamed batch: {e}"),
)))
})?;
let shaped = shape_decoded_rows(&value, Some(&schema_out));
for row in &shaped.rows {
if emitted >= limit {
break;
}
let encoded =
encode_shaped_row(&row_schema, &display_columns, &column_types, &row_formats, row)?;
emitted += 1;
yield encoded;
}
}
};
Response::Query(QueryResponse::new(schema, row_stream))
}
fn single_pgwire_error(err: PgWireError) -> Response {
let schema = Arc::new(vec![text_field("result")]);
let errored: Vec<PgWireResult<_>> = vec![Err(err)];
Response::Query(QueryResponse::new(schema, futures::stream::iter(errored)))
}
pub(crate) async fn streaming_star_response(stream: ResultStream, limit: usize) -> Response {
use futures::StreamExt;
let mut values: Vec<serde_json::Value> = Vec::new();
let mut batches = stream;
while let Some(batch) = batches.next().await {
let batch = match batch {
Ok(b) => b,
Err(e) => {
let (severity, code, message) = error_to_sqlstate(&e);
return single_pgwire_error(PgWireError::UserError(Box::new(ErrorInfo::new(
severity.to_owned(),
code.to_owned(),
message,
))));
}
};
if values.len() >= limit {
break;
}
let text = decode_payload_to_json(&batch.payload);
match sonic_rs::from_str::<serde_json::Value>(&text) {
Ok(serde_json::Value::Array(items)) => {
for item in items {
if values.len() >= limit {
break;
}
values.push(item);
}
}
Ok(_) => {
return single_pgwire_error(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"XX000".to_owned(),
"streamed batch payload was not a JSON array".to_owned(),
))));
}
Err(e) => {
return single_pgwire_error(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"XX000".to_owned(),
format!("failed to decode streamed batch: {e}"),
))));
}
}
}
if values.is_empty() {
let schema = Arc::new(vec![text_field("result")]);
return Response::Query(QueryResponse::new(
schema,
futures::stream::iter(Vec::<PgWireResult<_>>::new()),
));
}
let shaped = shape_decoded_rows(&serde_json::Value::Array(values), None);
let (response, _notice) = shaped_query_response(shaped, &[]);
response
}