1use std::ffi::OsStr;
8use std::fmt;
9use std::net::IpAddr;
10use std::time::Duration;
11
12use reqwest::header::CONTENT_TYPE;
13use reqwest::{Client, Response, Url};
14use secrecy::{ExposeSecret, SecretString};
15use serde::de::DeserializeOwned;
16use serde::Serialize;
17use uqa_sql::{AsyncSQLEngine, SQLParam, SQLResult};
18
19use crate::cli_connection;
20use crate::server_error_envelope::ServerErrorEnvelope;
21use crate::sql_batch_execution::SQLBatchWireResponse;
22use crate::sql_execution::SQLWireResponse;
23use crate::{HttpEngineError, SQLBatchExecution, SQLExecution, SQLStatement, SQLStream};
24
25const JSON_CONTENT_TYPE: &str = "application/json";
26const NDJSON_CONTENT_TYPE: &str = "application/x-ndjson";
27const REQUEST_ID_HEADER: &str = "x-request-id";
28const MAX_JSON_RESPONSE_BYTES: usize = 65 * 1024 * 1024;
29const MAX_ERROR_RESPONSE_BYTES: usize = 64 * 1024;
30
31#[derive(Clone)]
33pub struct HttpEngine {
34 http: Client,
35 base_url: Url,
36 credential: SecretString,
37}
38
39#[derive(Serialize)]
40struct SQLBatchRequest<'a> {
41 statements: &'a [SQLStatement],
42}
43
44impl HttpEngine {
45 pub fn new(base_url: &str, credential: SecretString) -> Result<Self, HttpEngineError> {
49 if credential.expose_secret().is_empty() {
50 return Err(HttpEngineError::InvalidCredential);
51 }
52 let base_url = parse_base_url(base_url)?;
53 let http = Client::builder()
54 .no_proxy()
55 .connect_timeout(Duration::from_secs(10))
56 .redirect(reqwest::redirect::Policy::none())
57 .user_agent(concat!("uqa-client/", env!("CARGO_PKG_VERSION")))
58 .build()
59 .map_err(HttpEngineError::build_client)?;
60 Ok(Self {
61 http,
62 base_url,
63 credential,
64 })
65 }
66
67 pub fn from_env() -> Result<Self, HttpEngineError> {
69 let base_url = std::env::var("UQA_URL")
70 .map_err(|_| HttpEngineError::MissingEnvironmentVariable("UQA_URL"))?;
71 let credential = std::env::var("UQA_TOKEN")
72 .map_err(|_| HttpEngineError::MissingEnvironmentVariable("UQA_TOKEN"))?;
73 Self::new(&base_url, SecretString::from(credential))
74 }
75
76 pub async fn local(project: &str) -> Result<Self, HttpEngineError> {
80 Self::local_with_cli(project, "uqa").await
81 }
82
83 pub async fn local_with_cli(
85 project: &str,
86 cli_path: impl AsRef<OsStr>,
87 ) -> Result<Self, HttpEngineError> {
88 let connection = cli_connection::resolve_local(cli_path.as_ref(), project).await?;
89 Self::new(&connection.url, connection.token)
90 }
91
92 pub async fn cloud(project: &str, organization: Option<&str>) -> Result<Self, HttpEngineError> {
97 Self::cloud_with_cli(project, organization, "uqa").await
98 }
99
100 pub async fn cloud_with_cli(
102 project: &str,
103 organization: Option<&str>,
104 cli_path: impl AsRef<OsStr>,
105 ) -> Result<Self, HttpEngineError> {
106 let connection =
107 cli_connection::resolve_cloud(cli_path.as_ref(), project, organization).await?;
108 Self::new(&connection.url, connection.token)
109 }
110
111 pub async fn sql(
113 &self,
114 query: &str,
115 params: &[SQLParam],
116 ) -> Result<SQLResult, HttpEngineError> {
117 Ok(self.sql_with_metadata(query, params).await?.into_result())
118 }
119
120 pub async fn sql_with_metadata(
122 &self,
123 query: &str,
124 params: &[SQLParam],
125 ) -> Result<SQLExecution, HttpEngineError> {
126 let statement = SQLStatement::new(query, params)?;
127 let response = self
128 .authorized(self.http.post(self.endpoint("v1/sql")?))
129 .json(&statement)
130 .send()
131 .await
132 .map_err(HttpEngineError::transport)?;
133 let request_id = response_request_id(&response)?;
134 let result = decode_json_response::<SQLWireResponse>(response).await?;
135 validate_request_id(&request_id, &result.request_id)?;
136 Ok(SQLExecution::from_wire(result))
137 }
138
139 pub async fn sql_batch(
141 &self,
142 statements: &[(&str, &[SQLParam])],
143 ) -> Result<Vec<SQLResult>, HttpEngineError> {
144 Ok(self
145 .sql_batch_with_metadata(statements)
146 .await?
147 .into_results())
148 }
149
150 pub async fn sql_batch_with_metadata(
152 &self,
153 statements: &[(&str, &[SQLParam])],
154 ) -> Result<SQLBatchExecution, HttpEngineError> {
155 let statements = statements
156 .iter()
157 .map(|(query, params)| SQLStatement::new(*query, params))
158 .collect::<Result<Vec<_>, _>>()?;
159 let response = self
160 .authorized(self.http.post(self.endpoint("v1/sql/batch")?))
161 .json(&SQLBatchRequest {
162 statements: &statements,
163 })
164 .send()
165 .await
166 .map_err(HttpEngineError::transport)?;
167 let request_id = response_request_id(&response)?;
168 let result = decode_json_response::<SQLBatchWireResponse>(response).await?;
169 validate_request_id(&request_id, &result.request_id)?;
170 Ok(SQLBatchExecution::from_wire(result))
171 }
172
173 pub async fn sql_stream(
175 &self,
176 query: &str,
177 params: &[SQLParam],
178 ) -> Result<SQLStream, HttpEngineError> {
179 let statement = SQLStatement::new(query, params)?;
180 let response = self
181 .authorized(self.http.post(self.endpoint("v1/sql/stream")?))
182 .header(reqwest::header::ACCEPT, NDJSON_CONTENT_TYPE)
183 .json(&statement)
184 .send()
185 .await
186 .map_err(HttpEngineError::transport)?;
187 if !response.status().is_success() {
188 return Err(error_from_response(response).await);
189 }
190 validate_content_type(&response, NDJSON_CONTENT_TYPE)?;
191 let request_id = response_request_id(&response)?;
192 Ok(SQLStream::new(response, request_id))
193 }
194
195 pub async fn subscribe_notifications(
197 &self,
198 channels: &[&str],
199 options: crate::notifications::HttpNotificationOptions,
200 ) -> Result<
201 crate::notifications::HttpNotificationSubscription,
202 crate::notifications::HttpNotificationError,
203 > {
204 self.subscribe_notifications_with_cancellation(
205 channels,
206 options,
207 &crate::notifications::NotificationCancellation::new(),
208 )
209 .await
210 }
211
212 pub async fn subscribe_notifications_with_cancellation(
214 &self,
215 channels: &[&str],
216 options: crate::notifications::HttpNotificationOptions,
217 cancellation: &crate::notifications::NotificationCancellation,
218 ) -> Result<
219 crate::notifications::HttpNotificationSubscription,
220 crate::notifications::HttpNotificationError,
221 > {
222 let endpoint = self
223 .endpoint("v1/notifications/subscribe")
224 .map_err(|_| crate::notifications::HttpNotificationError::invalid_options())?;
225 crate::notifications::http::subscribe(
226 endpoint,
227 self.credential.clone(),
228 channels,
229 options,
230 cancellation.clone(),
231 )
232 .await
233 }
234
235 fn authorized(&self, request: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
236 request.bearer_auth(self.credential.expose_secret())
237 }
238
239 fn endpoint(&self, path: &str) -> Result<Url, HttpEngineError> {
240 self.base_url
241 .join(path)
242 .map_err(|_| HttpEngineError::InvalidBaseURL)
243 }
244}
245
246impl AsyncSQLEngine for HttpEngine {
247 type Error = HttpEngineError;
248
249 async fn sql<'a>(
250 &'a self,
251 query: &'a str,
252 params: &'a [SQLParam],
253 ) -> Result<SQLResult, Self::Error> {
254 HttpEngine::sql(self, query, params).await
255 }
256
257 async fn sql_batch<'a>(
258 &'a self,
259 statements: &'a [(&'a str, &'a [SQLParam])],
260 ) -> Result<Vec<SQLResult>, Self::Error> {
261 HttpEngine::sql_batch(self, statements).await
262 }
263}
264
265impl fmt::Debug for HttpEngine {
266 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
267 formatter
268 .debug_struct("HttpEngine")
269 .field("base_url", &"[REDACTED]")
270 .field("credential", &"[REDACTED]")
271 .finish()
272 }
273}
274
275fn parse_base_url(source: &str) -> Result<Url, HttpEngineError> {
276 let url = Url::parse(source).map_err(|_| HttpEngineError::InvalidBaseURL)?;
277 let valid_origin = url.username().is_empty()
278 && url.password().is_none()
279 && url.query().is_none()
280 && url.fragment().is_none()
281 && url.path() == "/"
282 && url.host_str().is_some();
283 if !valid_origin || !matches!(url.scheme(), "http" | "https") {
284 return Err(HttpEngineError::InvalidBaseURL);
285 }
286 if url.scheme() == "http" && !url.host_str().is_some_and(is_loopback_host) {
287 return Err(HttpEngineError::InsecureRemoteURL);
288 }
289 Ok(url)
290}
291
292fn is_loopback_host(host: &str) -> bool {
293 let host = host
294 .strip_prefix('[')
295 .and_then(|host| host.strip_suffix(']'))
296 .unwrap_or(host);
297 host.eq_ignore_ascii_case("localhost")
298 || host
299 .parse::<IpAddr>()
300 .is_ok_and(|address| address.is_loopback())
301}
302
303async fn decode_json_response<T: DeserializeOwned>(
304 response: Response,
305) -> Result<T, HttpEngineError> {
306 if !response.status().is_success() {
307 return Err(error_from_response(response).await);
308 }
309 validate_content_type(&response, JSON_CONTENT_TYPE)?;
310 let body = read_bounded(response, MAX_JSON_RESPONSE_BYTES).await?;
311 serde_json::from_slice(&body).map_err(HttpEngineError::InvalidResponse)
312}
313
314async fn error_from_response(response: Response) -> HttpEngineError {
315 let status = response.status();
316 if let Err(error) = validate_content_type(&response, JSON_CONTENT_TYPE) {
317 return error;
318 }
319 let header_request_id = match response_request_id(&response) {
320 Ok(request_id) => request_id,
321 Err(error) => return error,
322 };
323 let body = match read_bounded(response, MAX_ERROR_RESPONSE_BYTES).await {
324 Ok(body) => body,
325 Err(error) => return error,
326 };
327 let Ok(envelope) = serde_json::from_slice::<ServerErrorEnvelope>(&body) else {
328 return HttpEngineError::Server {
329 status,
330 code: "HTTP_ERROR".to_owned(),
331 message: "UQA returned a non-success response".to_owned(),
332 request_id: Some(header_request_id),
333 };
334 };
335 if header_request_id != envelope.request_id {
336 return HttpEngineError::ResponseRequestIdMismatch;
337 }
338 HttpEngineError::Server {
339 status,
340 code: envelope.error.code,
341 message: envelope.error.message,
342 request_id: Some(envelope.request_id),
343 }
344}
345
346async fn read_bounded(
347 mut response: Response,
348 maximum_bytes: usize,
349) -> Result<Vec<u8>, HttpEngineError> {
350 if response
351 .content_length()
352 .is_some_and(|length| length > maximum_bytes as u64)
353 {
354 return Err(HttpEngineError::ResponseTooLarge);
355 }
356 let mut body = Vec::new();
357 while let Some(chunk) = response.chunk().await.map_err(HttpEngineError::transport)? {
358 if body.len().saturating_add(chunk.len()) > maximum_bytes {
359 return Err(HttpEngineError::ResponseTooLarge);
360 }
361 body.extend_from_slice(&chunk);
362 }
363 Ok(body)
364}
365
366fn validate_content_type(response: &Response, expected: &str) -> Result<(), HttpEngineError> {
367 let valid = response
368 .headers()
369 .get(CONTENT_TYPE)
370 .and_then(|value| value.to_str().ok())
371 .and_then(|value| value.split(';').next())
372 .is_some_and(|value| value.trim().eq_ignore_ascii_case(expected));
373 if valid {
374 Ok(())
375 } else {
376 Err(HttpEngineError::UnexpectedContentType)
377 }
378}
379
380fn response_request_id(response: &Response) -> Result<String, HttpEngineError> {
381 response
382 .headers()
383 .get(REQUEST_ID_HEADER)
384 .and_then(|value| value.to_str().ok())
385 .filter(|value| !value.is_empty())
386 .map(str::to_owned)
387 .ok_or(HttpEngineError::MissingRequestId)
388}
389
390fn validate_request_id(header: &str, body: &str) -> Result<(), HttpEngineError> {
391 if header == body {
392 Ok(())
393 } else {
394 Err(HttpEngineError::ResponseRequestIdMismatch)
395 }
396}
397
398#[cfg(test)]
399mod tests {
400 use super::*;
401
402 #[test]
403 fn base_url_rejects_credentials_paths_and_remote_plain_http() {
404 for source in [
405 "http://user:secret@127.0.0.1:8432/",
406 "http://127.0.0.1:8432/v1",
407 "http://example.com/",
408 "ftp://127.0.0.1/",
409 ] {
410 assert!(parse_base_url(source).is_err(), "accepted {source}");
411 }
412 assert!(parse_base_url("http://127.0.0.1:8432/").is_ok());
413 assert!(parse_base_url("http://[::1]:8432/").is_ok());
414 assert!(parse_base_url("https://example.com/").is_ok());
415 }
416
417 #[test]
418 fn debug_output_redacts_endpoint_and_credential() {
419 let credential = "uqa_db_customer-secret";
420 let client =
421 HttpEngine::new("http://127.0.0.1:8432/", SecretString::from(credential)).unwrap();
422 let debug = format!("{client:?}");
423 assert!(!debug.contains("127.0.0.1"));
424 assert!(!debug.contains(credential));
425 }
426}