use std::collections::VecDeque;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use crate::protocol::messages::{
SqlFieldsFirstPage, SqlFieldsPage, encode_cursor_get_page, encode_resource_close,
};
use crate::protocol::op_code;
use crate::transport::{IgniteConnection, next_request_id};
use futures::Stream;
use crate::error::{IgniteError, Result};
use crate::query::{Column, Row};
pub struct QueryStream {
pub columns: Vec<Column>,
inner: Pin<Box<dyn Stream<Item = Result<Row>> + Send>>,
}
impl QueryStream {
pub(crate) fn new(
columns: Vec<Column>,
inner: Pin<Box<dyn Stream<Item = Result<Row>> + Send>>,
) -> Self {
Self { columns, inner }
}
}
impl QueryStream {
pub async fn collect_all(mut self) -> crate::error::Result<Vec<Row>> {
use futures::StreamExt;
let mut rows = Vec::new();
while let Some(result) = self.next().await {
rows.push(result?);
}
Ok(rows)
}
}
impl Stream for QueryStream {
type Item = Result<Row>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.inner.as_mut().poll_next(cx)
}
}
struct CursorState {
buffer: VecDeque<Row>,
columns: Vec<Column>,
cursor_id: i64,
has_more: bool,
field_count: usize,
conn: Arc<IgniteConnection>,
}
impl Drop for CursorState {
fn drop(&mut self) {
let req_id = next_request_id();
let payload = encode_resource_close(req_id, self.cursor_id);
let conn = Arc::clone(&self.conn);
tokio::spawn(async move {
let _ = conn.request(req_id, payload).await;
});
}
}
pub(crate) fn build_stream(conn: Arc<IgniteConnection>, first: SqlFieldsFirstPage) -> QueryStream {
let columns: Vec<Column> = first
.field_names
.iter()
.map(|n| Column { name: n.clone() })
.collect();
let field_count = columns.len();
let cursor_id = first.cursor_id;
let has_more = first.has_more;
let buffer: VecDeque<Row> = first
.rows
.into_iter()
.map(|vals| Row::new(columns.clone(), vals))
.collect();
let state = CursorState {
buffer,
columns: columns.clone(),
cursor_id,
has_more,
field_count,
conn,
};
let inner = futures::stream::unfold(state, |mut s| async move {
loop {
if let Some(row) = s.buffer.pop_front() {
return Some((Ok(row), s));
}
if !s.has_more {
return None;
}
let req_id = next_request_id();
let payload = encode_cursor_get_page(
op_code::QUERY_SQL_FIELDS_CURSOR_GET_PAGE,
req_id,
s.cursor_id,
);
let mut response = match s.conn.request(req_id, payload).await {
Err(e) => return Some((Err(IgniteError::Transport(e)), s)),
Ok(r) => r,
};
let page = match SqlFieldsPage::decode(&mut response, s.field_count) {
Err(e) => return Some((Err(IgniteError::Protocol(e)), s)),
Ok(p) => p,
};
s.has_more = page.has_more;
s.buffer = page
.rows
.into_iter()
.map(|vals| Row::new(s.columns.clone(), vals))
.collect();
}
});
QueryStream::new(columns, Box::pin(inner))
}