use std::collections::HashSet;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::error::{ApiError, ValidationDetails};
use crate::ids::UserId;
use crate::time::UnixMillis;
pub const MAX_NAME_BYTES: usize = 128;
pub const DEFAULT_MAX_OBJECT_BYTES: usize = 256 * 1024;
pub const DEFAULT_MAX_OBJECTS_PER_USER: u32 = 1000;
pub const MAX_BATCH: usize = 16;
pub const MAX_BATCH_BYTES: usize = 4 * 1024 * 1024;
pub const PUT_BODY_LIMIT_BYTES: usize = DEFAULT_MAX_OBJECT_BYTES + 16 * 1024;
pub const BATCH_BODY_LIMIT_BYTES: usize = MAX_BATCH_BYTES + 64 * 1024;
const _: () = assert!(MAX_BATCH * DEFAULT_MAX_OBJECT_BYTES <= MAX_BATCH_BYTES);
const _: () = assert!(BATCH_BODY_LIMIT_BYTES < 5 * 1024 * 1024);
pub fn is_valid_name(name: &str) -> bool {
let bytes = name.as_bytes();
bytes.len() <= MAX_NAME_BYTES
&& bytes.first().is_some_and(u8::is_ascii_alphanumeric)
&& bytes.iter().all(|b| b.is_ascii_alphanumeric() || matches!(b, b'_' | b'-' | b'.'))
}
pub fn value_bytes(value: &Value) -> usize {
serde_json::to_vec(value).map_or(usize::MAX, |json| json.len())
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
#[serde(transparent)]
pub struct ObjectVersion(pub i64);
impl ObjectVersion {
pub const ABSENT: ObjectVersion = ObjectVersion(0);
pub const fn new(value: i64) -> Self {
Self(value)
}
pub const fn get(self) -> i64 {
self.0
}
pub fn etag(self) -> String {
format!("\"{}\"", self.0)
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum WriteAccess {
#[default]
Owner,
Server,
#[serde(other)]
Unknown,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct VersionConflict {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub index: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub current_version: Option<ObjectVersion>,
}
impl VersionConflict {
pub fn new(current_version: Option<ObjectVersion>) -> Self {
Self { index: None, current_version }
}
pub fn at_index(mut self, index: u32) -> Self {
self.index = Some(index);
self
}
pub fn into_error(self) -> ApiError {
let details = serde_json::to_value(self).unwrap_or(Value::Null);
ApiError::new(crate::codes::VERSION_CONFLICT, "the object's version is not the expected one").with_details(details)
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct PutObject {
pub value: Value,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub if_version: Option<ObjectVersion>,
}
impl PutObject {
pub fn new(value: Value) -> Self {
Self { value, if_version: None }
}
pub fn from_serialize<T: Serialize + ?Sized>(value: &T) -> Result<Self, serde_json::Error> {
Ok(Self::new(serde_json::to_value(value)?))
}
pub fn if_version(mut self, version: ObjectVersion) -> Self {
self.if_version = Some(version);
self
}
pub fn if_absent(self) -> Self {
self.if_version(ObjectVersion::ABSENT)
}
pub fn validate(&self, max_bytes: usize) -> Result<(), ApiError> {
let mut details = ValidationDetails::new();
check_value(&self.value, max_bytes, "value", &mut details);
details.into_result()
}
}
fn check_value(value: &Value, max_bytes: usize, field: &str, details: &mut ValidationDetails) -> usize {
let size = value_bytes(value);
if size > max_bytes {
details.add(field, format!("is larger than {max_bytes} bytes"));
}
size
}
fn check_names(collection: &str, key: &str, prefix: &str, details: &mut ValidationDetails) {
if !is_valid_name(collection) {
details.add(format!("{prefix}collection"), "is not a valid storage name");
}
if !is_valid_name(key) {
details.add(format!("{prefix}key"), "is not a valid storage name");
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct DeleteObject {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub if_version: Option<ObjectVersion>,
}
impl DeleteObject {
pub fn new() -> Self {
Self::default()
}
pub fn if_version(mut self, version: ObjectVersion) -> Self {
self.if_version = Some(version);
self
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct StorageObject {
pub collection: String,
pub key: String,
pub owner: UserId,
pub value: Value,
pub version: ObjectVersion,
#[serde(default)]
pub write: WriteAccess,
pub updated_at: UnixMillis,
}
impl StorageObject {
pub fn new(collection: impl Into<String>, key: impl Into<String>, owner: UserId, value: Value, version: ObjectVersion, updated_at: UnixMillis) -> Self {
Self { collection: collection.into(), key: key.into(), owner, value, version, write: WriteAccess::Owner, updated_at }
}
pub fn with_write(mut self, write: WriteAccess) -> Self {
self.write = write;
self
}
pub fn value_as<T: serde::de::DeserializeOwned>(&self) -> Result<T, serde_json::Error> {
T::deserialize(&self.value)
}
pub fn info(&self) -> StorageObjectInfo {
StorageObjectInfo {
collection: self.collection.clone(),
key: self.key.clone(),
version: self.version,
write: self.write,
size_bytes: u64::try_from(value_bytes(&self.value)).unwrap_or(u64::MAX),
updated_at: self.updated_at,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct StorageObjectInfo {
pub collection: String,
pub key: String,
pub version: ObjectVersion,
#[serde(default)]
pub write: WriteAccess,
pub size_bytes: u64,
pub updated_at: UnixMillis,
}
impl StorageObjectInfo {
pub fn new(collection: impl Into<String>, key: impl Into<String>, version: ObjectVersion, size_bytes: u64, updated_at: UnixMillis) -> Self {
Self { collection: collection.into(), key: key.into(), version, write: WriteAccess::Owner, size_bytes, updated_at }
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ObjectAck {
pub collection: String,
pub key: String,
pub version: ObjectVersion,
pub updated_at: UnixMillis,
}
impl ObjectAck {
pub fn new(collection: impl Into<String>, key: impl Into<String>, version: ObjectVersion, updated_at: UnixMillis) -> Self {
Self { collection: collection.into(), key: key.into(), version, updated_at }
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ObjectRef {
pub collection: String,
pub key: String,
}
impl ObjectRef {
pub fn new(collection: impl Into<String>, key: impl Into<String>) -> Self {
Self { collection: collection.into(), key: key.into() }
}
}
fn check_batch_len(len: usize, details: &mut ValidationDetails) {
if len == 0 {
details.add("objects", "is empty");
}
if len > MAX_BATCH {
details.add("objects", format!("has more than {MAX_BATCH} entries"));
}
}
fn check_duplicates<'a>(names: impl Iterator<Item = (&'a str, &'a str)>, details: &mut ValidationDetails) {
let mut seen = HashSet::new();
for (i, name) in names.enumerate() {
if !seen.insert(name) {
details.add(format!("objects.{i}"), "names the same object as an earlier entry");
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct BatchGet {
pub objects: Vec<ObjectRef>,
}
impl BatchGet {
pub fn new(objects: Vec<ObjectRef>) -> Self {
Self { objects }
}
pub fn validate(&self) -> Result<(), ApiError> {
let mut details = ValidationDetails::new();
check_batch_len(self.objects.len(), &mut details);
for (i, object) in self.objects.iter().enumerate() {
check_names(&object.collection, &object.key, &format!("objects.{i}."), &mut details);
}
check_duplicates(self.objects.iter().map(|o| (o.collection.as_str(), o.key.as_str())), &mut details);
details.into_result()
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct BatchObjects {
pub objects: Vec<StorageObject>,
}
impl BatchObjects {
pub fn new(objects: Vec<StorageObject>) -> Self {
Self { objects }
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct BatchPutItem {
pub collection: String,
pub key: String,
#[serde(flatten)]
pub put: PutObject,
}
impl BatchPutItem {
pub fn new(collection: impl Into<String>, key: impl Into<String>, put: PutObject) -> Self {
Self { collection: collection.into(), key: key.into(), put }
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct BatchPut {
pub objects: Vec<BatchPutItem>,
}
impl BatchPut {
pub fn new(objects: Vec<BatchPutItem>) -> Self {
Self { objects }
}
pub fn validate(&self, max_object_bytes: usize) -> Result<(), ApiError> {
let mut details = ValidationDetails::new();
check_batch_len(self.objects.len(), &mut details);
let mut total = 0usize;
for (i, item) in self.objects.iter().enumerate() {
let prefix = format!("objects.{i}.");
check_names(&item.collection, &item.key, &prefix, &mut details);
total = total.saturating_add(check_value(&item.put.value, max_object_bytes, &format!("{prefix}value"), &mut details));
}
if total > MAX_BATCH_BYTES {
details.add("objects", format!("carry more than {MAX_BATCH_BYTES} bytes of values"));
}
check_duplicates(self.objects.iter().map(|o| (o.collection.as_str(), o.key.as_str())), &mut details);
details.into_result()
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct BatchAcks {
pub objects: Vec<ObjectAck>,
}
impl BatchAcks {
pub fn new(objects: Vec<ObjectAck>) -> Self {
Self { objects }
}
}
mod calls {
use super::*;
use crate::envelope::Ack;
use crate::http_call::{payload_call, HttpCall, NoPayload, PathParams, PayloadKind, NO_PAYLOAD};
use crate::page::{Page, PageRequest};
use crate::routes::{self, HttpMethod, Route};
payload_call!(BatchGet, Post, routes::storage::BATCH_GET, true, Json, BatchObjects);
payload_call!(BatchPut, Post, routes::storage::BATCH_PUT, true, Json, BatchAcks);
const NOT_A_NAME: &str = "is not a valid storage name";
fn names(params: &PathParams) -> Result<(String, String), ApiError> {
Ok((params.checked("collection", is_valid_name, NOT_A_NAME)?, params.checked("key", is_valid_name, NOT_A_NAME)?))
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct ListObjects {
pub collection: String,
pub page: PageRequest,
}
impl ListObjects {
pub fn new(collection: impl Into<String>) -> Self {
Self { collection: collection.into(), page: PageRequest::first() }
}
pub fn with_page(mut self, page: PageRequest) -> Self {
self.page = page;
self
}
}
impl HttpCall for ListObjects {
type Payload = PageRequest;
type Response = Page<StorageObjectInfo>;
const ROUTE: Route = Route::new(HttpMethod::Get, routes::storage::COLLECTION, true);
const PAYLOAD: PayloadKind = PayloadKind::Query;
fn payload(&self) -> &PageRequest {
&self.page
}
fn path_params(&self) -> PathParams {
PathParams::new().with("collection", &self.collection)
}
fn from_parts(params: &PathParams, page: PageRequest) -> Result<Self, ApiError> {
Ok(Self::new(params.checked("collection", is_valid_name, NOT_A_NAME)?).with_page(page))
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct GetObject {
pub collection: String,
pub key: String,
}
impl GetObject {
pub fn new(collection: impl Into<String>, key: impl Into<String>) -> Self {
Self { collection: collection.into(), key: key.into() }
}
}
impl HttpCall for GetObject {
type Payload = NoPayload;
type Response = StorageObject;
const ROUTE: Route = Route::new(HttpMethod::Get, routes::storage::OBJECT, true);
const PAYLOAD: PayloadKind = PayloadKind::Empty;
fn payload(&self) -> &NoPayload {
&NO_PAYLOAD
}
fn path_params(&self) -> PathParams {
PathParams::new().with("collection", &self.collection).with("key", &self.key)
}
fn from_parts(params: &PathParams, _payload: NoPayload) -> Result<Self, ApiError> {
let (collection, key) = names(params)?;
Ok(Self::new(collection, key))
}
}
#[derive(Clone, Debug, PartialEq)]
#[non_exhaustive]
pub struct WriteObject {
pub collection: String,
pub key: String,
pub put: PutObject,
}
impl WriteObject {
pub fn new(collection: impl Into<String>, key: impl Into<String>, put: PutObject) -> Self {
Self { collection: collection.into(), key: key.into(), put }
}
}
impl HttpCall for WriteObject {
type Payload = PutObject;
type Response = ObjectAck;
const ROUTE: Route = Route::new(HttpMethod::Put, routes::storage::OBJECT, true);
const PAYLOAD: PayloadKind = PayloadKind::Json;
fn payload(&self) -> &PutObject {
&self.put
}
fn path_params(&self) -> PathParams {
PathParams::new().with("collection", &self.collection).with("key", &self.key)
}
fn from_parts(params: &PathParams, put: PutObject) -> Result<Self, ApiError> {
let (collection, key) = names(params)?;
Ok(Self::new(collection, key, put))
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct RemoveObject {
pub collection: String,
pub key: String,
pub delete: DeleteObject,
}
impl RemoveObject {
pub fn new(collection: impl Into<String>, key: impl Into<String>) -> Self {
Self { collection: collection.into(), key: key.into(), delete: DeleteObject::new() }
}
pub fn if_version(mut self, version: ObjectVersion) -> Self {
self.delete = self.delete.if_version(version);
self
}
}
impl HttpCall for RemoveObject {
type Payload = DeleteObject;
type Response = Ack;
const ROUTE: Route = Route::new(HttpMethod::Delete, routes::storage::OBJECT, true);
const PAYLOAD: PayloadKind = PayloadKind::Query;
fn payload(&self) -> &DeleteObject {
&self.delete
}
fn path_params(&self) -> PathParams {
PathParams::new().with("collection", &self.collection).with("key", &self.key)
}
fn from_parts(params: &PathParams, delete: DeleteObject) -> Result<Self, ApiError> {
let (collection, key) = names(params)?;
Ok(Self { collection, key, delete })
}
}
}
pub use calls::{GetObject, ListObjects, RemoveObject, WriteObject};
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn names() {
for good in ["saves", "slot-1", "a.b_c", "0", &"x".repeat(MAX_NAME_BYTES)] {
assert!(is_valid_name(good), "{good}");
}
for bad in ["", ".", "..", "_batch", "-x", "a/b", "a b", "ä", "a%2F", "a\0", "a\u{200B}", &"x".repeat(MAX_NAME_BYTES + 1)] {
assert!(!is_valid_name(bad), "{bad}");
}
}
#[test]
fn versions_and_rules() {
assert_eq!(ObjectVersion(3).etag(), "\"3\"");
assert_eq!(PutObject::new(Value::Null).if_absent().if_version, Some(ObjectVersion::ABSENT));
assert!(PutObject::new(Value::String("x".repeat(100))).validate(50).is_err());
assert!(PutObject::new(Value::Bool(true)).validate(DEFAULT_MAX_OBJECT_BYTES).is_ok());
assert!(BatchGet::new(vec![]).validate().is_err());
assert!(BatchGet::new(vec![ObjectRef::new("saves", "a")]).validate().is_ok());
assert!(BatchGet::new(vec![ObjectRef::new("saves", "../a")]).validate().is_err());
assert!(BatchGet::new(vec![ObjectRef::new("s", "k"); 2]).validate().is_err());
let many: Vec<ObjectRef> = (0..=MAX_BATCH).map(|i| ObjectRef::new("s", format!("k{i}"))).collect();
assert!(BatchGet::new(many).validate().is_err());
let put = BatchPut::new(vec![BatchPutItem::new("saves", "a", PutObject::new(Value::from("y".repeat(100))))]);
assert!(put.validate(DEFAULT_MAX_OBJECT_BYTES).is_ok());
assert!(put.validate(10).is_err());
}
#[test]
fn batch_duplicates_and_budget() {
let item = |key: &str, size: usize| BatchPutItem::new("saves", key, PutObject::new(Value::from("x".repeat(size))));
let duplicate = BatchPut::new(vec![item("a", 1), item("b", 1), item("a", 1)]);
let error = duplicate.validate(DEFAULT_MAX_OBJECT_BYTES).err().and_then(|e| e.details_as::<ValidationDetails>()).unwrap_or_default();
assert!(error.fields.contains_key("objects.2"), "{error:?}");
let heavy: Vec<BatchPutItem> = (0..MAX_BATCH).map(|i| item(&format!("k{i}"), 300 * 1024)).collect();
assert!(BatchPut::new(heavy).validate(512 * 1024).is_err());
}
#[test]
fn conflict_details_and_info() {
let error = VersionConflict::new(Some(ObjectVersion(3))).at_index(2).into_error();
assert_eq!(error.details, Some(serde_json::json!({"index": 2, "current_version": 3})));
assert_eq!(error.http_status(), 409);
assert_eq!(serde_json::to_string(&VersionConflict::new(None)).ok().as_deref(), Some("{}"));
let object = StorageObject::new("s", "k", UserId(1), serde_json::json!({"level": 3}), ObjectVersion(1), UnixMillis(5));
assert_eq!(object.info().size_bytes, 11);
#[derive(Deserialize)]
struct Save {
level: u32,
}
assert_eq!(object.value_as::<Save>().map(|s| s.level).ok(), Some(3));
assert_eq!(serde_json::from_str::<WriteAccess>("\"moderators\"").ok(), Some(WriteAccess::Unknown));
}
}