use sha2::{Digest, Sha256};
use std::marker::PhantomData;
pub struct PageCursor<const N: usize>([u8; N]);
mod sealed {
pub trait Cursor {}
impl<const N: usize> Cursor for super::PageCursor<N> {}
}
#[doc(hidden)]
pub trait CursorWire: sealed::Cursor + Send {
fn zero() -> Self;
fn bytes(&self) -> &[u8];
fn bytes_mut(&mut self) -> &mut [u8];
}
impl<const N: usize> CursorWire for PageCursor<N> {
fn zero() -> Self {
Self([0; N])
}
fn bytes(&self) -> &[u8] {
&self.0
}
fn bytes_mut(&mut self) -> &mut [u8] {
&mut self.0
}
}
#[cfg(test)]
mod cursor_tests {
type PageCursor = super::PageCursor<98>;
#[test]
fn cursor_storage_is_fixed_and_noncanonical_text_is_rejected() {
assert_eq!(std::mem::size_of::<PageCursor>(), 98);
for text in ["", "0", &"0".repeat(97), &"0".repeat(99), &"A".repeat(98)] {
assert!(PageCursor::parse(text).is_err());
}
let cursor = PageCursor::parse(&"0".repeat(98)).unwrap();
assert_eq!(cursor.as_str().len(), 98);
}
}
impl<const N: usize> PageCursor<N> {
pub fn as_str(&self) -> &str {
std::str::from_utf8(&self.0).expect("hex cursor")
}
pub fn parse(value: &str) -> Result<Self, crate::database::DatabaseFailure> {
if value.len() != N
|| !value
.bytes()
.all(|b| b.is_ascii_digit() || (b'a'..=b'f').contains(&b))
{
return Err(crate::database::DatabaseFailure::State);
}
let mut bytes = [0; N];
bytes.copy_from_slice(value.as_bytes());
Ok(Self(bytes))
}
}
impl<const N: usize> serde::Serialize for PageCursor<N> {
fn serialize<S: serde::Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
s.serialize_str(self.as_str())
}
}
impl<'de, const N: usize> serde::Deserialize<'de> for PageCursor<N> {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
struct Visitor<const N: usize>;
impl<const N: usize> serde::de::Visitor<'_> for Visitor<N> {
type Value = PageCursor<N>;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("canonical Saddle page cursor")
}
fn visit_str<E: serde::de::Error>(self, v: &str) -> Result<PageCursor<N>, E> {
PageCursor::parse(v).map_err(|_| E::custom("invalid page cursor"))
}
}
d.deserialize_str(Visitor::<N>)
}
}
#[doc(hidden)]
pub fn page_filter_fingerprint<O: crate::database_capability::GeneratedOwnedQuery>(
fields: &[(&str, &[u8])],
) -> [u8; 32] {
let mut hash = Sha256::new();
fn part(hash: &mut Sha256, part: &[u8]) {
hash.update((part.len() as u64).to_le_bytes());
hash.update(part);
}
part(&mut hash, O::OPERATION.as_bytes());
part(&mut hash, O::SQL.as_bytes());
for r in O::RELATIONS {
part(&mut hash, r.alias.as_bytes());
part(&mut hash, r.table.as_bytes());
for c in r.columns {
part(&mut hash, c.as_bytes());
}
}
part(&mut hash, std::any::type_name::<O::Parameters>().as_bytes());
for (name, value) in fields {
part(&mut hash, name.as_bytes());
part(&mut hash, value);
}
hash.finalize().into()
}
fn cursor_digest<O: crate::database_capability::GeneratedPageQuery>(
params: &O::NamedParameters,
after: &O::Key,
size: u64,
) -> [u8; 32] {
let mut h = Sha256::new();
h.update(b"saddle-keyset-v2");
h.update(O::filter_fingerprint(params));
for value in after.as_ref() {
h.update(value.to_le_bytes());
}
h.update(size.to_le_bytes());
h.finalize().into()
}
pub struct PageRequest<O: crate::database_capability::GeneratedPageQuery> {
pub(crate) parameters: O::NamedParameters,
pub(crate) after: Option<O::Key>,
pub(crate) offset: Option<u64>,
pub(crate) size: usize,
}
impl<O: crate::database_capability::GeneratedPageQuery> PageRequest<O> {
pub fn cursor(&self) -> Option<O::Cursor> {
if O::IS_OFFSET {
return None;
}
let after = self.after?;
let size = u64::try_from(self.size).ok()?;
let mut cursor = O::Cursor::zero();
let hex = cursor.bytes_mut();
if hex.len() != 2 * (1 + 8 * (after.as_ref().len() + 1) + 32) {
return None;
}
fn put(hex: &mut [u8], index: usize, value: u8) {
let chars = b"0123456789abcdef";
hex[index * 2] = chars[(value >> 4) as usize];
hex[index * 2 + 1] = chars[(value & 15) as usize];
}
put(hex, 0, 2);
for (n, value) in after
.as_ref()
.iter()
.chain(std::iter::once(&size))
.enumerate()
{
for (i, b) in value.to_le_bytes().iter().enumerate() {
put(hex, 1 + n * 8 + i, *b);
}
}
let start = 1 + (after.as_ref().len() + 1) * 8;
for (i, b) in cursor_digest::<O>(&self.parameters, &after, size)
.iter()
.enumerate()
{
put(hex, start + i, *b);
}
Some(cursor)
}
pub fn resume(
parameters: O::NamedParameters,
size: usize,
cursor: &O::Cursor,
) -> Result<Self, crate::database::DatabaseFailure> {
let mut request = Self::first(parameters, size)?;
if O::IS_OFFSET {
return Err(crate::database::DatabaseFailure::State);
}
let wire = cursor.bytes();
if wire.len() != 2 * (1 + 8 * (O::empty_key().as_ref().len() + 1) + 32)
|| !wire
.iter()
.all(|b| b.is_ascii_digit() || (b'a'..=b'f').contains(b))
{
return Err(crate::database::DatabaseFailure::State);
}
fn get(wire: &[u8], i: usize) -> u8 {
fn digit(b: u8) -> u8 {
if b <= b'9' { b - b'0' } else { b - b'a' + 10 }
}
digit(wire[i * 2]) * 16 + digit(wire[i * 2 + 1])
}
fn number(wire: &[u8], start: usize) -> u64 {
let mut b = [0; 8];
for (i, v) in b.iter_mut().enumerate() {
*v = get(wire, start + i);
}
u64::from_le_bytes(b)
}
let mut after = O::empty_key();
for (n, v) in after.as_mut().iter_mut().enumerate() {
*v = number(wire, 1 + n * 8);
}
let start = 1 + after.as_ref().len() * 8;
let encoded_size = number(wire, start);
let digest = cursor_digest::<O>(&request.parameters, &after, encoded_size);
if get(wire, 0) != 2
|| u64::try_from(size).ok() != Some(encoded_size)
|| digest
.iter()
.enumerate()
.any(|(i, b)| get(wire, start + 8 + i) != *b)
{
return Err(crate::database::DatabaseFailure::State);
}
request.after = Some(after);
Ok(request)
}
pub fn first(
parameters: O::NamedParameters,
size: usize,
) -> Result<Self, crate::database::DatabaseFailure> {
if size == 0 || size.checked_add(1).is_none() {
return Err(crate::database::DatabaseFailure::State);
}
Ok(Self {
parameters,
after: None,
offset: if O::IS_OFFSET { Some(0) } else { None },
size,
})
}
}
impl<O: crate::database_capability::GeneratedOffsetPageQuery> PageRequest<O> {
pub fn at(
parameters: O::NamedParameters,
size: usize,
offset: u64,
) -> Result<Self, crate::database::DatabaseFailure> {
let mut request = Self::first(parameters, size)?;
let limit = u64::try_from(size)
.ok()
.and_then(|size| size.checked_add(1));
if limit.and_then(|limit| offset.checked_add(limit)).is_none() {
return Err(crate::database::DatabaseFailure::ResultLimit);
}
request.offset = Some(offset);
Ok(request)
}
}
pub struct Page<O: crate::database_capability::GeneratedPageQuery> {
pub(crate) rows: Rows<O::Row, O::NamedRow>,
pub(crate) next: Option<PageRequest<O>>,
}
impl<O: crate::database_capability::GeneratedPageQuery> Page<O> {
pub fn into_parts(self) -> (Rows<O::Row, O::NamedRow>, Option<PageRequest<O>>) {
(self.rows, self.next)
}
}
pub struct Rows<Raw, Named> {
rows: saddle_db::internal::OwnedDbRows<Raw>,
named: PhantomData<fn() -> Named>,
}
impl<Raw, Named> Rows<Raw, Named> {
pub fn into_values(self) -> saddle_db::internal::OwnedDbList<Raw> {
self.rows.into()
}
pub(crate) fn from_database(rows: saddle_db::internal::OwnedDbRows<Raw>) -> Self {
Self {
rows,
named: PhantomData,
}
}
pub fn len(&self) -> usize {
self.rows.as_slice().len()
}
pub fn is_empty(&self) -> bool {
self.rows.as_slice().is_empty()
}
}
pub struct IntoRows<Raw, Named> {
rows: saddle_admission::ManagedVecIntoIter<Raw>,
named: PhantomData<fn() -> Named>,
}
impl<Raw, Named: From<Raw>> Iterator for IntoRows<Raw, Named> {
type Item = Named;
fn next(&mut self) -> Option<Named> {
self.rows.next().map(Named::from)
}
fn size_hint(&self) -> (usize, Option<usize>) {
self.rows.size_hint()
}
}
impl<Raw, Named: From<Raw>> IntoIterator for Rows<Raw, Named> {
type Item = Named;
type IntoIter = IntoRows<Raw, Named>;
fn into_iter(self) -> Self::IntoIter {
IntoRows {
rows: self.rows.into_iter(),
named: PhantomData,
}
}
}