1use std::collections::HashMap;
20use std::sync::Arc;
21
22use axum::extract::State;
23use axum::http::{HeaderMap, StatusCode};
24use axum::response::{IntoResponse, Response};
25use axum::routing::post;
26use axum::{Json, Router};
27use base64::Engine as _;
28use boatramp_core::compute::{BindingKind, ComputeBinding, ComputeBindingResolver};
29use boatramp_core::sql::{SqlBackend, SqlBackends, SqlValue};
30use hmac::{Hmac, Mac};
31use serde::{Deserialize, Serialize};
32use sha2::Sha256;
33use tokio::sync::RwLock;
34
35#[derive(Clone, Default)]
37pub struct SqlShim {
38 registry: Arc<RwLock<HashMap<String, Arc<dyn SqlBackend>>>>,
39}
40
41impl SqlShim {
42 pub fn new() -> Self {
44 Self::default()
45 }
46
47 pub async fn register(&self, token: String, backend: Arc<dyn SqlBackend>) {
50 self.registry.write().await.insert(token, backend);
51 }
52
53 pub async fn deregister(&self, token: &str) {
55 self.registry.write().await.remove(token);
56 }
57
58 async fn lookup(&self, token: &str) -> Option<Arc<dyn SqlBackend>> {
59 self.registry.read().await.get(token).cloned()
60 }
61
62 pub fn router(&self) -> Router {
65 Router::new()
66 .route("/v2/pipeline", post(pipeline))
67 .route("/v3/pipeline", post(pipeline))
68 .with_state(self.clone())
69 }
70}
71
72pub struct SqlShimResolver {
78 provider: Arc<dyn SqlBackends>,
79 shim: SqlShim,
80 base_url: String,
81 secret: [u8; 32],
82}
83
84impl SqlShimResolver {
85 pub fn new(
88 provider: Arc<dyn SqlBackends>,
89 shim: SqlShim,
90 base_url: String,
91 secret: [u8; 32],
92 ) -> Self {
93 Self {
94 provider,
95 shim,
96 base_url,
97 secret,
98 }
99 }
100
101 fn token(
104 &self,
105 project: &str,
106 workload: &str,
107 replica: u32,
108 binding: &ComputeBinding,
109 ) -> String {
110 let mut mac =
111 Hmac::<Sha256>::new_from_slice(&self.secret).expect("hmac accepts any key len");
112 for part in [
113 project.as_bytes(),
114 workload.as_bytes(),
115 binding.name.as_bytes(),
116 ] {
117 mac.update(part);
118 mac.update(&[0]);
119 }
120 mac.update(&replica.to_le_bytes());
121 mac.update(format!("{:?}", binding.kind).as_bytes());
122 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes())
123 }
124}
125
126#[async_trait::async_trait]
127impl ComputeBindingResolver for SqlShimResolver {
128 async fn resolve(
129 &self,
130 project: &str,
131 workload: &str,
132 replica: u32,
133 bindings: &[ComputeBinding],
134 ) -> Vec<(String, String)> {
135 let mut env = Vec::new();
136 for binding in bindings {
137 if binding.kind != BindingKind::Sql {
139 continue;
140 }
141 let backend = match self
144 .provider
145 .database(project, workload, &binding.name)
146 .await
147 {
148 Ok(backend) => backend,
149 Err(err) => {
150 tracing::warn!(%project, %workload, error = %err, "sql binding: resolve failed");
151 continue;
152 }
153 };
154 let token = self.token(project, workload, replica, binding);
155 self.shim.register(token.clone(), backend).await;
156 let url_env = binding.url_env();
157 env.push((url_env.clone(), self.base_url.clone()));
158 env.push((format!("{url_env}_AUTH_TOKEN"), token));
159 }
160 env
161 }
162
163 async fn release(
164 &self,
165 project: &str,
166 workload: &str,
167 replica: u32,
168 bindings: &[ComputeBinding],
169 ) {
170 for binding in bindings {
171 if binding.kind == BindingKind::Sql {
172 self.shim
173 .deregister(&self.token(project, workload, replica, binding))
174 .await;
175 }
176 }
177 }
178}
179
180pub async fn spawn_sql_shim(
186 sql: Option<Arc<dyn SqlBackends>>,
187 shim_url: Option<String>,
188) -> Option<Arc<dyn ComputeBindingResolver>> {
189 let sql = sql?;
190 let base_url = shim_url?;
191 let Some(port) = base_url
192 .rsplit_once(':')
193 .and_then(|(_, p)| p.trim_end_matches('/').parse::<u16>().ok())
194 else {
195 tracing::warn!(%base_url, "compute.sql_shim_url has no :port; sql bindings disabled");
196 return None;
197 };
198 let mut secret = [0u8; 32];
199 if getrandom::getrandom(&mut secret).is_err() {
200 tracing::error!("getrandom failed; sql bindings disabled");
201 return None;
202 }
203 let shim = SqlShim::new();
204 let router = shim.router();
205 let bind = std::net::SocketAddr::from(([0, 0, 0, 0], port));
206 match tokio::net::TcpListener::bind(bind).await {
207 Ok(listener) => {
208 use axum::serve::ListenerExt;
212 let listener = listener.tap_io(crate::disable_nagle);
213 tokio::spawn(async move {
214 if let Err(err) = axum::serve(listener, router).await {
215 tracing::error!(error = %err, "compute sql-shim listener exited");
216 }
217 });
218 tracing::info!(%bind, %base_url, "compute sql-shim listening");
219 }
220 Err(err) => {
221 tracing::warn!(%bind, error = %err, "compute sql-shim bind failed; sql bindings disabled");
222 return None;
223 }
224 }
225 Some(Arc::new(SqlShimResolver::new(sql, shim, base_url, secret)))
226}
227
228fn bearer(headers: &HeaderMap) -> Option<String> {
230 headers
231 .get(axum::http::header::AUTHORIZATION)?
232 .to_str()
233 .ok()?
234 .strip_prefix("Bearer ")
235 .map(|s| s.trim().to_string())
236}
237
238#[derive(Deserialize)]
241struct PipelineReq {
242 #[serde(default)]
243 requests: Vec<StreamRequest>,
244}
245
246#[derive(Deserialize)]
247#[serde(tag = "type", rename_all = "snake_case")]
248enum StreamRequest {
249 Execute {
250 stmt: Stmt,
251 },
252 Close,
253 #[serde(other)]
256 Unsupported,
257}
258
259#[derive(Deserialize)]
260struct Stmt {
261 sql: Option<String>,
262 #[serde(default)]
263 args: Vec<Value>,
264 #[serde(default)]
265 want_rows: bool,
266}
267
268#[derive(Serialize)]
269struct PipelineResp {
270 baton: Option<String>,
271 base_url: Option<String>,
272 results: Vec<StreamResult>,
273}
274
275#[derive(Serialize)]
276#[serde(tag = "type", rename_all = "snake_case")]
277enum StreamResult {
278 Ok { response: HranaResponse },
279 Error { error: HranaError },
280}
281
282#[derive(Serialize)]
283#[serde(tag = "type", rename_all = "snake_case")]
284enum HranaResponse {
285 Execute { result: StmtResult },
286 Close,
287}
288
289#[derive(Serialize)]
290struct StmtResult {
291 cols: Vec<Col>,
292 rows: Vec<Vec<Value>>,
293 affected_row_count: u64,
294 last_insert_rowid: Option<String>,
295}
296
297#[derive(Serialize)]
298struct Col {
299 name: Option<String>,
300 decltype: Option<String>,
301}
302
303#[derive(Serialize)]
304struct HranaError {
305 message: String,
306}
307
308#[derive(Serialize, Deserialize)]
311#[serde(tag = "type", rename_all = "snake_case")]
312enum Value {
313 Null,
314 Integer { value: String },
315 Float { value: f64 },
316 Text { value: String },
317 Blob { base64: String },
318}
319
320impl Value {
321 fn from_sql(v: &SqlValue) -> Self {
322 match v {
323 SqlValue::Null => Self::Null,
324 SqlValue::Boolean(b) => Self::Integer {
325 value: (i64::from(*b)).to_string(),
326 },
327 SqlValue::Integer(i) => Self::Integer {
328 value: i.to_string(),
329 },
330 SqlValue::Real(f) => Self::Float { value: *f },
331 SqlValue::Text(s) => Self::Text { value: s.clone() },
332 SqlValue::Json(s) => Self::Text { value: s.clone() },
335 SqlValue::Blob(b) => Self::Blob {
336 base64: base64::engine::general_purpose::STANDARD.encode(b),
337 },
338 }
339 }
340
341 fn to_sql(&self) -> Result<SqlValue, String> {
342 Ok(match self {
343 Self::Null => SqlValue::Null,
344 Self::Integer { value } => SqlValue::Integer(
345 value
346 .parse()
347 .map_err(|_| "invalid integer arg".to_string())?,
348 ),
349 Self::Float { value } => SqlValue::Real(*value),
350 Self::Text { value } => SqlValue::Text(value.clone()),
351 Self::Blob { base64 } => SqlValue::Blob(
352 base64::engine::general_purpose::STANDARD
353 .decode(base64)
354 .map_err(|_| "invalid base64 blob arg".to_string())?,
355 ),
356 })
357 }
358}
359
360async fn pipeline(
364 State(shim): State<SqlShim>,
365 headers: HeaderMap,
366 Json(req): Json<PipelineReq>,
367) -> Response {
368 let Some(token) = bearer(&headers) else {
369 return (StatusCode::UNAUTHORIZED, "missing bearer token\n").into_response();
370 };
371 let Some(backend) = shim.lookup(&token).await else {
372 return (StatusCode::UNAUTHORIZED, "unknown token\n").into_response();
373 };
374
375 let mut results = Vec::with_capacity(req.requests.len());
376 for request in req.requests {
377 let result = match request {
378 StreamRequest::Close => StreamResult::Ok {
379 response: HranaResponse::Close,
380 },
381 StreamRequest::Unsupported => StreamResult::Error {
382 error: HranaError {
383 message: "unsupported request type in the stateless pipeline".to_string(),
384 },
385 },
386 StreamRequest::Execute { stmt } => match run_stmt(backend.as_ref(), stmt).await {
387 Ok(result) => StreamResult::Ok {
388 response: HranaResponse::Execute { result },
389 },
390 Err(message) => StreamResult::Error {
391 error: HranaError { message },
392 },
393 },
394 };
395 results.push(result);
396 }
397
398 Json(PipelineResp {
399 baton: None,
400 base_url: None,
401 results,
402 })
403 .into_response()
404}
405
406async fn run_stmt(backend: &dyn SqlBackend, stmt: Stmt) -> Result<StmtResult, String> {
410 let sql = stmt.sql.ok_or_else(|| "statement has no sql".to_string())?;
411 let params: Vec<SqlValue> = stmt
412 .args
413 .iter()
414 .map(Value::to_sql)
415 .collect::<Result<_, _>>()?;
416
417 let mut tx = backend.begin().await.map_err(|e| e.to_string())?;
418 if stmt.want_rows {
419 let rows = tx.query(&sql, ¶ms).await.map_err(|e| e.to_string());
420 let rows = match rows {
421 Ok(rows) => rows,
422 Err(e) => return Err(e),
423 };
424 tx.commit().await.map_err(|e| e.to_string())?;
425 Ok(StmtResult {
426 cols: rows
427 .columns
428 .iter()
429 .map(|c| Col {
430 name: Some(c.clone()),
431 decltype: None,
432 })
433 .collect(),
434 rows: rows
435 .rows
436 .iter()
437 .map(|row| row.iter().map(Value::from_sql).collect())
438 .collect(),
439 affected_row_count: 0,
440 last_insert_rowid: None,
441 })
442 } else {
443 let affected = match tx.execute(&sql, ¶ms).await.map_err(|e| e.to_string()) {
444 Ok(n) => n,
445 Err(e) => return Err(e),
446 };
447 tx.commit().await.map_err(|e| e.to_string())?;
448 Ok(StmtResult {
449 cols: vec![],
450 rows: vec![],
451 affected_row_count: affected,
452 last_insert_rowid: None,
453 })
454 }
455}
456
457#[cfg(test)]
458mod tests {
459 use super::*;
460 use async_trait::async_trait;
461 use axum::body::Body;
462 use axum::http::Request;
463 use boatramp_core::sql::{SqlError, SqlRows, SqlTransaction};
464 use std::sync::Mutex;
465 use tower::ServiceExt as _;
466
467 type Seen = Arc<Mutex<Vec<(String, Vec<SqlValue>)>>>;
469
470 #[derive(Default)]
474 struct FakeBackend {
475 seen: Seen,
476 }
477 struct FakeTx {
478 seen: Seen,
479 }
480
481 #[async_trait]
482 impl SqlBackend for FakeBackend {
483 async fn begin(&self) -> Result<Box<dyn SqlTransaction>, SqlError> {
484 Ok(Box::new(FakeTx {
485 seen: self.seen.clone(),
486 }))
487 }
488 }
489
490 #[async_trait]
491 impl SqlTransaction for FakeTx {
492 async fn query(&mut self, sql: &str, params: &[SqlValue]) -> Result<SqlRows, SqlError> {
493 self.seen
494 .lock()
495 .unwrap()
496 .push((sql.to_string(), params.to_vec()));
497 Ok(SqlRows {
498 columns: vec!["n".to_string()],
499 rows: vec![vec![SqlValue::Integer(42)]],
500 })
501 }
502 async fn execute(&mut self, sql: &str, params: &[SqlValue]) -> Result<u64, SqlError> {
503 self.seen
504 .lock()
505 .unwrap()
506 .push((sql.to_string(), params.to_vec()));
507 Ok(7)
508 }
509 async fn commit(self: Box<Self>) -> Result<(), SqlError> {
510 Ok(())
511 }
512 async fn rollback(self: Box<Self>) -> Result<(), SqlError> {
513 Ok(())
514 }
515 }
516
517 struct FakeProvider;
519 #[async_trait]
520 impl SqlBackends for FakeProvider {
521 async fn database(
522 &self,
523 _project: &str,
524 _site: &str,
525 _name: &str,
526 ) -> Result<Arc<dyn SqlBackend>, SqlError> {
527 Ok(Arc::new(FakeBackend::default()))
528 }
529 }
530
531 async fn post(
532 shim: &SqlShim,
533 token: Option<&str>,
534 body: serde_json::Value,
535 ) -> (StatusCode, serde_json::Value) {
536 let mut builder = Request::builder()
537 .method("POST")
538 .uri("/v2/pipeline")
539 .header("content-type", "application/json");
540 if let Some(t) = token {
541 builder = builder.header("authorization", format!("Bearer {t}"));
542 }
543 let req = builder.body(Body::from(body.to_string())).unwrap();
544 let resp = shim.router().oneshot(req).await.unwrap();
545 let status = resp.status();
546 let bytes = axum::body::to_bytes(resp.into_body(), 1 << 20)
547 .await
548 .unwrap();
549 let json = serde_json::from_slice(&bytes).unwrap_or(serde_json::Value::Null);
550 (status, json)
551 }
552
553 #[tokio::test]
554 async fn unknown_or_missing_token_is_unauthorized() {
555 let shim = SqlShim::new();
556 let body = serde_json::json!({ "requests": [{ "type": "close" }] });
557 assert_eq!(
558 post(&shim, None, body.clone()).await.0,
559 StatusCode::UNAUTHORIZED
560 );
561 assert_eq!(
562 post(&shim, Some("nope"), body).await.0,
563 StatusCode::UNAUTHORIZED
564 );
565 }
566
567 #[tokio::test]
568 async fn a_registered_token_runs_a_query_and_maps_values() {
569 let shim = SqlShim::new();
570 let backend = Arc::new(FakeBackend::default());
571 shim.register("tok".to_string(), backend.clone()).await;
572
573 let body = serde_json::json!({
575 "requests": [
576 { "type": "execute", "stmt": {
577 "sql": "SELECT n WHERE x = ?",
578 "args": [{ "type": "text", "value": "hi" }],
579 "want_rows": true }},
580 { "type": "close" }
581 ]
582 });
583 let (status, json) = post(&shim, Some("tok"), body).await;
584 assert_eq!(status, StatusCode::OK);
585 let result = &json["results"][0]["response"]["result"];
586 assert_eq!(result["cols"][0]["name"], "n");
587 assert_eq!(result["rows"][0][0]["type"], "integer");
588 assert_eq!(result["rows"][0][0]["value"], "42");
589 assert_eq!(json["results"][1]["type"], "ok"); let seen = backend.seen.lock().unwrap();
593 assert_eq!(seen[0].0, "SELECT n WHERE x = ?");
594 assert_eq!(seen[0].1, vec![SqlValue::Text("hi".to_string())]);
595 }
596
597 #[tokio::test]
598 async fn want_rows_false_routes_to_execute_and_returns_affected() {
599 let shim = SqlShim::new();
600 shim.register("tok".to_string(), Arc::new(FakeBackend::default()))
601 .await;
602 let body = serde_json::json!({
603 "requests": [
604 { "type": "execute", "stmt": { "sql": "INSERT INTO t VALUES (1)", "want_rows": false }}
605 ]
606 });
607 let (status, json) = post(&shim, Some("tok"), body).await;
608 assert_eq!(status, StatusCode::OK);
609 assert_eq!(
610 json["results"][0]["response"]["result"]["affected_row_count"],
611 7
612 );
613 }
614
615 #[tokio::test]
616 async fn resolver_registers_a_working_token_and_release_revokes() {
617 let shim = SqlShim::new();
618 let resolver = SqlShimResolver::new(
619 Arc::new(FakeProvider),
620 shim.clone(),
621 "http://10.0.0.1:9999".to_string(),
622 [7u8; 32],
623 );
624 let bindings = vec![ComputeBinding {
625 kind: BindingKind::Sql,
626 name: String::new(),
627 url_env: None,
628 }];
629
630 let env = resolver.resolve("acme", "api", 0, &bindings).await;
631 let get = |k: &str| env.iter().find(|(key, _)| key == k).map(|(_, v)| v.clone());
632 assert_eq!(
633 get("BOATRAMP_SQL_URL").as_deref(),
634 Some("http://10.0.0.1:9999")
635 );
636 let token = get("BOATRAMP_SQL_URL_AUTH_TOKEN").expect("token env is injected");
637
638 let body = serde_json::json!({ "requests": [{ "type": "close" }] });
640 assert_eq!(
641 post(&shim, Some(&token), body.clone()).await.0,
642 StatusCode::OK
643 );
644 assert_eq!(
646 resolver.resolve("acme", "api", 0, &bindings).await,
647 env,
648 "token derivation is deterministic"
649 );
650 resolver.release("acme", "api", 0, &bindings).await;
652 assert_eq!(
653 post(&shim, Some(&token), body).await.0,
654 StatusCode::UNAUTHORIZED
655 );
656 }
657
658 #[tokio::test]
659 async fn deregister_revokes_access() {
660 let shim = SqlShim::new();
661 shim.register("tok".to_string(), Arc::new(FakeBackend::default()))
662 .await;
663 shim.deregister("tok").await;
664 let body = serde_json::json!({ "requests": [{ "type": "close" }] });
665 assert_eq!(
666 post(&shim, Some("tok"), body).await.0,
667 StatusCode::UNAUTHORIZED
668 );
669 }
670}