1use crate::{HttpEngineError, SQLStreamFrame};
8
9const MAX_STREAM_FRAME_BYTES: usize = 64 * 1024 * 1024;
10
11pub struct SQLStream {
13 response: reqwest::Response,
14 request_id: String,
15 buffer: Vec<u8>,
16 scan_start: usize,
17 phase: StreamPhase,
18 body_finished: bool,
19}
20
21#[derive(Clone, Copy, Eq, PartialEq)]
22enum StreamPhase {
23 AwaitingMetadata,
24 Rows,
25 Terminal,
26 Finished,
27}
28
29impl SQLStream {
30 pub(crate) fn new(response: reqwest::Response, request_id: String) -> Self {
31 Self {
32 response,
33 request_id,
34 buffer: Vec::new(),
35 scan_start: 0,
36 phase: StreamPhase::AwaitingMetadata,
37 body_finished: false,
38 }
39 }
40
41 pub fn request_id(&self) -> &str {
42 &self.request_id
43 }
44
45 pub async fn next_frame(&mut self) -> Result<Option<SQLStreamFrame>, HttpEngineError> {
46 loop {
47 if self.phase == StreamPhase::Finished {
48 return Ok(None);
49 }
50 if self.phase == StreamPhase::Terminal {
51 return self.finish().await;
52 }
53 if let Some(newline) = self.next_newline() {
54 if newline > MAX_STREAM_FRAME_BYTES {
55 return Err(HttpEngineError::StreamFrameTooLarge);
56 }
57 let line = take_line(&mut self.buffer, newline);
58 self.scan_start = 0;
59 if line_is_empty(&line) {
60 continue;
61 }
62 return self.decode_frame(&line).map(Some);
63 }
64 if self.buffer.len() > MAX_STREAM_FRAME_BYTES {
65 return Err(HttpEngineError::StreamFrameTooLarge);
66 }
67 if self.body_finished {
68 if self.buffer.is_empty() {
69 return Err(HttpEngineError::TruncatedStream);
70 }
71 if self.buffer.len() > MAX_STREAM_FRAME_BYTES {
72 return Err(HttpEngineError::StreamFrameTooLarge);
73 }
74 let line = std::mem::take(&mut self.buffer);
75 return self.decode_frame(&line).map(Some);
76 }
77 match self
78 .response
79 .chunk()
80 .await
81 .map_err(HttpEngineError::transport)?
82 {
83 Some(chunk) => {
84 self.buffer.extend_from_slice(&chunk);
85 }
86 None => self.body_finished = true,
87 }
88 }
89 }
90
91 fn decode_frame(&mut self, line: &[u8]) -> Result<SQLStreamFrame, HttpEngineError> {
92 let line = line.strip_suffix(b"\r").unwrap_or(line);
93 let frame = serde_json::from_slice::<SQLStreamFrame>(line)
94 .map_err(HttpEngineError::InvalidResponse)?;
95 if frame
96 .request_id()
97 .is_some_and(|request_id| request_id != self.request_id)
98 {
99 return Err(HttpEngineError::StreamRequestIdMismatch);
100 }
101 self.phase = match (&self.phase, &frame) {
102 (StreamPhase::AwaitingMetadata, SQLStreamFrame::Metadata { .. })
103 | (StreamPhase::Rows, SQLStreamFrame::Row { .. }) => StreamPhase::Rows,
104 (StreamPhase::AwaitingMetadata | StreamPhase::Rows, SQLStreamFrame::Error { .. })
105 | (StreamPhase::Rows, SQLStreamFrame::Complete { .. }) => StreamPhase::Terminal,
106 _ => return Err(HttpEngineError::InvalidStreamSequence),
107 };
108 Ok(frame)
109 }
110
111 async fn finish(&mut self) -> Result<Option<SQLStreamFrame>, HttpEngineError> {
112 loop {
113 if let Some(newline) = self.next_newline() {
114 if newline > MAX_STREAM_FRAME_BYTES {
115 return Err(HttpEngineError::StreamFrameTooLarge);
116 }
117 let line = take_line(&mut self.buffer, newline);
118 self.scan_start = 0;
119 if !line_is_empty(&line) {
120 return Err(HttpEngineError::InvalidStreamSequence);
121 }
122 continue;
123 }
124 if self.buffer.len() > MAX_STREAM_FRAME_BYTES {
125 return Err(HttpEngineError::StreamFrameTooLarge);
126 }
127 if self.body_finished {
128 if !line_is_empty(&self.buffer) {
129 return Err(HttpEngineError::InvalidStreamSequence);
130 }
131 self.buffer.clear();
132 self.phase = StreamPhase::Finished;
133 return Ok(None);
134 }
135 match self
136 .response
137 .chunk()
138 .await
139 .map_err(HttpEngineError::transport)?
140 {
141 Some(chunk) => {
142 self.buffer.extend_from_slice(&chunk);
143 }
144 None => self.body_finished = true,
145 }
146 }
147 }
148
149 fn next_newline(&mut self) -> Option<usize> {
150 let start = self.scan_start.min(self.buffer.len());
151 let newline = self.buffer[start..]
152 .iter()
153 .position(|byte| *byte == b'\n')
154 .map(|offset| start + offset);
155 if newline.is_none() {
156 self.scan_start = self.buffer.len();
157 }
158 newline
159 }
160}
161
162fn line_is_empty(line: &[u8]) -> bool {
163 line.is_empty() || line == b"\r"
164}
165
166fn take_line(buffer: &mut Vec<u8>, newline: usize) -> Vec<u8> {
167 let remainder = buffer.split_off(newline + 1);
168 let mut line = std::mem::replace(buffer, remainder);
169 line.truncate(newline);
170 line
171}