1use std::vec::IntoIter;
16
17use futures::stream::FusedStream;
18use futures::{StreamExt, TryStreamExt};
19use postgres_types::FromSql;
20
21use crate::error::{PsqlError, PsqlResult};
22use crate::pg_message::{BeCommandCompleteMessage, BeMessage};
23use crate::pg_protocol::{PgByteStream, PgStream};
24use crate::pg_response::{PgResponse, ValuesStream};
25use crate::types::{Format, Row};
26
27pub struct ResultCache<VS>
28where
29 VS: ValuesStream,
30{
31 result: PgResponse<VS>,
32 row_cache: IntoIter<Row>,
33}
34
35impl<VS> ResultCache<VS>
36where
37 VS: ValuesStream,
38{
39 pub fn new(result: PgResponse<VS>) -> Self {
40 ResultCache {
41 result,
42 row_cache: vec![].into_iter(),
43 }
44 }
45
46 pub async fn consume<S: PgByteStream>(
48 &mut self,
49 row_limit: usize,
50 msg_stream: &mut PgStream<S>,
51 ) -> PsqlResult<bool> {
52 for notice in self.result.notices() {
53 msg_stream.write_no_flush(BeMessage::NoticeResponse(notice))?;
54 }
55
56 let status = self.result.status();
57 if let Some(ref application_name) = status.application_name {
58 msg_stream.write_no_flush(BeMessage::ParameterStatus(
59 crate::pg_message::BeParameterStatusMessage::ApplicationName(application_name),
60 ))?;
61 }
62
63 if self.result.is_empty() {
64 self.result.run_callback().await?;
66
67 msg_stream.write_no_flush(BeMessage::EmptyQueryResponse)?;
68 return Ok(true);
69 }
70
71 let mut query_end = false;
72 if self.result.is_copy_query_to_stdout() {
73 msg_stream.write_no_flush(BeMessage::CopyOutResponse(self.result.row_desc().len()))?;
74
75 let mut count = 0;
76 while let Some(row_set) = self.result.values_stream().next().await {
77 let row_set = row_set.map_err(PsqlError::SimpleQueryError)?;
78 for row in row_set {
79 msg_stream
80 .write_streaming(BeMessage::CopyData(&row))
81 .await?;
82 count += 1;
83 }
84 }
85
86 msg_stream.write_no_flush(BeMessage::CopyDone)?;
87
88 self.result.run_callback().await?;
90
91 msg_stream.write_no_flush(BeMessage::CommandComplete(BeCommandCompleteMessage {
92 stmt_type: self.result.stmt_type(),
93 rows_cnt: count,
94 }))?;
95
96 query_end = true;
97 } else if self.result.is_query() {
98 let mut query_row_count = 0;
99
100 while row_limit == 0 || query_row_count < row_limit {
104 if self.row_cache.len() > 0 {
105 for row in self.row_cache.by_ref() {
106 msg_stream.write_streaming(BeMessage::DataRow(&row)).await?;
107 query_row_count += 1;
108 if row_limit > 0 && query_row_count >= row_limit {
109 break;
110 }
111 }
112 } else {
113 self.row_cache = match self
114 .result
115 .values_stream()
116 .try_next()
117 .await
118 .map_err(PsqlError::ExtendedExecuteError)?
119 {
120 Some(rows) => rows.into_iter(),
121 _ => {
122 query_end = true;
123 break;
124 }
125 };
126 }
127 }
128
129 if self.row_cache.len() == 0 && self.result.values_stream().peekable().is_terminated() {
132 query_end = true;
133 }
134 if query_end {
135 self.result.run_callback().await?;
137
138 msg_stream.write_no_flush(BeMessage::CommandComplete(
139 BeCommandCompleteMessage {
140 stmt_type: self.result.stmt_type(),
141 rows_cnt: query_row_count as i32,
142 },
143 ))?;
144 } else {
145 msg_stream.write_no_flush(BeMessage::PortalSuspended)?;
146 }
147 } else if self.result.stmt_type().is_dml() && !self.result.stmt_type().is_returning() {
148 let first_row_set = self.result.values_stream().next().await;
149 let first_row_set = match first_row_set {
150 None => {
151 return Err(PsqlError::Uncategorized(
152 "no affected rows in output".into(),
153 ));
154 }
155 Some(row) => row.map_err(PsqlError::SimpleQueryError)?,
156 };
157 let affected_rows_str = first_row_set[0].values()[0]
158 .as_ref()
159 .expect("compute node should return affected rows in output");
160
161 let affected_rows_cnt: i32 = match self.result.row_cnt_format() {
162 Some(Format::Binary) => {
163 i64::from_sql(&postgres_types::Type::INT8, affected_rows_str)
164 .unwrap()
165 .try_into()
166 .expect("affected rows count large than i64")
167 }
168 Some(Format::Text) => String::from_utf8(affected_rows_str.to_vec())
169 .unwrap()
170 .parse()
171 .unwrap_or_default(),
172 None => panic!("affected rows count should be set"),
173 };
174
175 self.result.run_callback().await?;
177
178 msg_stream.write_no_flush(BeMessage::CommandComplete(BeCommandCompleteMessage {
179 stmt_type: self.result.stmt_type(),
180 rows_cnt: affected_rows_cnt,
181 }))?;
182
183 query_end = true;
184 } else {
185 self.result.run_callback().await?;
187
188 msg_stream.write_no_flush(BeMessage::CommandComplete(BeCommandCompleteMessage {
189 stmt_type: self.result.stmt_type(),
190 rows_cnt: self
191 .result
192 .affected_rows_cnt()
193 .expect("row count should be set"),
194 }))?;
195
196 query_end = true;
197 }
198
199 Ok(query_end)
200 }
201}