Skip to main content

uqa_client/
sql_stream.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7use crate::{HttpEngineError, SQLStreamFrame};
8
9const MAX_STREAM_FRAME_BYTES: usize = 64 * 1024 * 1024;
10
11/// Incremental reader for one authenticated UQA NDJSON SQL response.
12pub 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}