use sha2::{Digest, Sha256};
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum ApqRequest {
None,
Lookup { hash: String },
Register { hash: String, query: String },
}
pub(crate) fn parse(body: &serde_json::Value) -> ApqRequest {
let hash = body
.pointer("/extensions/persistedQuery/sha256Hash")
.and_then(|h| h.as_str());
let Some(hash) = hash else {
return ApqRequest::None;
};
match body.get("query").and_then(|q| q.as_str()) {
Some(query) => ApqRequest::Register {
hash: hash.to_string(),
query: query.to_string(),
},
None => ApqRequest::Lookup {
hash: hash.to_string(),
},
}
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum ApqOutcome {
Run {
query: String,
store: Option<(String, String)>,
},
Error(String),
}
pub(crate) const NOT_FOUND: &str = "PersistedQueryNotFound";
pub(crate) fn resolve(
req: ApqRequest,
stored: Option<String>,
safelist: bool,
) -> Option<ApqOutcome> {
match req {
ApqRequest::None => None,
ApqRequest::Lookup { .. } => Some(match stored {
Some(query) => ApqOutcome::Run { query, store: None },
None => ApqOutcome::Error(NOT_FOUND.to_string()),
}),
ApqRequest::Register { hash, query } => {
if sha256_hex(&query) != hash {
return Some(ApqOutcome::Error(
"provided sha256Hash does not match the query".to_string(),
));
}
if let Some(existing) = stored {
return Some(ApqOutcome::Run {
query: existing,
store: None,
});
}
if safelist {
return Some(ApqOutcome::Error(NOT_FOUND.to_string()));
}
Some(ApqOutcome::Run {
query: query.clone(),
store: Some((hash, query)),
})
}
}
}
pub(crate) fn sha256_hex(s: &str) -> String {
hex::encode(Sha256::digest(s.as_bytes()))
}
fn apq_key(scope: &str, hash: &str) -> String {
format!("hapq/{scope}/{hash}")
}
pub(crate) async fn safelisted_query(
kv: &dyn boatramp_core::kv::KvStore,
scope: &str,
hash: &str,
) -> Option<String> {
kv.get(&apq_key(scope, hash))
.await
.ok()
.flatten()
.and_then(|b| String::from_utf8(b).ok())
}
fn apq_prefix(scope: &str) -> String {
format!("hapq/{scope}/")
}
pub(crate) async fn register(
kv: &dyn boatramp_core::kv::KvStore,
scope: &str,
query: &str,
) -> Result<String, String> {
let hash = sha256_hex(query);
kv.put(&apq_key(scope, &hash), query.as_bytes().to_vec())
.await
.map_err(|e| e.to_string())?;
Ok(hash)
}
pub(crate) async fn list(
kv: &dyn boatramp_core::kv::KvStore,
scope: &str,
) -> Vec<(String, String)> {
let prefix = apq_prefix(scope);
let mut out = Vec::new();
for key in kv.list_prefix(&prefix).await.unwrap_or_default() {
if let Ok(Some(bytes)) = kv.get(&key).await {
if let Ok(query) = String::from_utf8(bytes) {
let hash = key.strip_prefix(&prefix).unwrap_or(&key).to_string();
out.push((hash, query));
}
}
}
out
}
pub(crate) async fn unregister(
kv: &dyn boatramp_core::kv::KvStore,
scope: &str,
hash: &str,
) -> Result<(), String> {
kv.delete(&apq_key(scope, hash))
.await
.map_err(|e| e.to_string())
}
pub(crate) enum Resolved {
Query(String),
Error(String),
Passthrough,
}
pub(crate) async fn resolve_stored(
kv: &dyn boatramp_core::kv::KvStore,
scope: &str,
body: &serde_json::Value,
safelist: bool,
) -> Resolved {
let req = parse(body);
let hash = match &req {
ApqRequest::None => return Resolved::Passthrough,
ApqRequest::Lookup { hash } | ApqRequest::Register { hash, .. } => hash.clone(),
};
let stored = kv
.get(&apq_key(scope, &hash))
.await
.ok()
.flatten()
.and_then(|b| String::from_utf8(b).ok());
match resolve(req, stored, safelist) {
None => Resolved::Passthrough,
Some(ApqOutcome::Error(msg)) => Resolved::Error(msg),
Some(ApqOutcome::Run { query, store }) => {
if let Some((h, q)) = store {
let _ = kv.put(&apq_key(scope, &h), q.into_bytes()).await;
}
Resolved::Query(query)
}
}
}
pub(crate) fn error_response(message: &str) -> axum::response::Response {
use axum::response::IntoResponse;
let body = serde_json::json!({ "errors": [ { "message": message } ] }).to_string();
(
axum::http::StatusCode::OK,
[(
axum::http::header::CONTENT_TYPE,
axum::http::HeaderValue::from_static("application/json"),
)],
body,
)
.into_response()
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
const Q: &str = "{ hello }";
fn hash_of(q: &str) -> String {
sha256_hex(q)
}
#[test]
fn parse_distinguishes_plain_lookup_and_register() {
assert_eq!(parse(&json!({ "query": Q })), ApqRequest::None);
let ext =
json!({ "extensions": { "persistedQuery": { "version": 1, "sha256Hash": "abc" } } });
assert_eq!(parse(&ext), ApqRequest::Lookup { hash: "abc".into() });
let reg = json!({
"query": Q,
"extensions": { "persistedQuery": { "version": 1, "sha256Hash": "abc" } }
});
assert_eq!(
parse(®),
ApqRequest::Register {
hash: "abc".into(),
query: Q.into()
}
);
}
#[test]
fn lookup_hit_runs_and_miss_reports_not_found() {
let req = ApqRequest::Lookup { hash: hash_of(Q) };
assert_eq!(
resolve(req, Some(Q.to_string()), false),
Some(ApqOutcome::Run {
query: Q.into(),
store: None
})
);
let miss = ApqRequest::Lookup { hash: hash_of(Q) };
assert_eq!(
resolve(miss, None, false),
Some(ApqOutcome::Error(NOT_FOUND.into()))
);
}
#[test]
fn register_persists_when_hash_matches() {
let h = hash_of(Q);
let req = ApqRequest::Register {
hash: h.clone(),
query: Q.into(),
};
assert_eq!(
resolve(req, None, false),
Some(ApqOutcome::Run {
query: Q.into(),
store: Some((h, Q.into()))
})
);
}
#[test]
fn register_with_wrong_hash_is_rejected() {
let req = ApqRequest::Register {
hash: "deadbeef".into(),
query: Q.into(),
};
assert!(matches!(
resolve(req, None, false),
Some(ApqOutcome::Error(m)) if m.contains("does not match")
));
}
#[test]
fn safelist_never_registers_a_new_query() {
let h = hash_of(Q);
let req = ApqRequest::Register {
hash: h.clone(),
query: Q.into(),
};
assert_eq!(
resolve(req, None, true),
Some(ApqOutcome::Error(NOT_FOUND.into()))
);
assert_eq!(
resolve(ApqRequest::Lookup { hash: h }, Some(Q.to_string()), true),
Some(ApqOutcome::Run {
query: Q.into(),
store: None
})
);
}
#[test]
fn non_apq_request_is_left_alone() {
assert_eq!(resolve(ApqRequest::None, None, false), None);
}
#[tokio::test]
async fn safelisted_returns_a_registered_op_and_none_otherwise() {
use boatramp_core::kv::{KvStore, MemoryKv};
let kv = MemoryKv::new();
let h = hash_of(Q);
assert_eq!(safelisted_query(&kv, "acme", &h).await, None);
kv.put(&apq_key("acme", &h), Q.as_bytes().to_vec())
.await
.unwrap();
assert_eq!(safelisted_query(&kv, "acme", &h).await, Some(Q.to_string()));
assert_eq!(safelisted_query(&kv, "other", &h).await, None);
}
#[tokio::test]
async fn register_list_and_unregister_round_trip() {
use boatramp_core::kv::MemoryKv;
let kv = MemoryKv::new();
let hash = register(&kv, "acme", Q).await.unwrap();
assert_eq!(hash, hash_of(Q));
assert_eq!(
safelisted_query(&kv, "acme", &hash).await,
Some(Q.to_string())
);
assert_eq!(list(&kv, "acme").await, vec![(hash.clone(), Q.to_string())]);
register(&kv, "acme", Q).await.unwrap();
assert_eq!(list(&kv, "acme").await.len(), 1);
unregister(&kv, "acme", &hash).await.unwrap();
assert!(safelisted_query(&kv, "acme", &hash).await.is_none());
assert!(list(&kv, "acme").await.is_empty());
unregister(&kv, "acme", &hash).await.unwrap();
}
}