Import upstream v0.16.22, stripped

Upstream commit: 474dd0229cb20cf513036619781ed97bd8073c3f
Enterprise-only files removed or emptied: 63
Enterprise-only snippets removed: 117 in 50 files
Dangling module declarations removed: 5
Cargo edits turning enterprise off: 14
Verification: clean
Enterprise feature gates left for rebuilt features: 19 in 18 files

Produced by tools/fork/strip.py. The full report is in docs/fork/strip-reports/ on main.
This commit is contained in:
2026-09-18 10:21:56 -07:00
commit 7dae9b29fd
1650 changed files with 485521 additions and 0 deletions
+169
View File
@@ -0,0 +1,169 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use azure_core::error::ErrorKind;
use azure_core::{ExponentialRetryOptions, RetryOptions, StatusCode};
use azure_storage::StorageCredentials;
use azure_storage_blobs::prelude::{ClientBuilder, ContainerClient};
use futures::stream::StreamExt;
use registry::schema::structs::{self};
use std::sync::Arc;
use std::{fmt::Display, io::Write, ops::Range};
use utils::codec::base32_custom::Base32Writer;
use crate::BlobStore;
pub struct AzureStore {
client: ContainerClient,
prefix: Option<String>,
}
impl AzureStore {
pub async fn open(config: structs::AzureStore) -> Result<BlobStore, String> {
let credentials = match (
config.access_key.secret().await?.map(|v| v.into_owned()),
config.sas_token.secret().await?.map(|v| v.into_owned()),
) {
(Some(access_key), None) => {
StorageCredentials::access_key(config.storage_account.clone(), access_key)
}
(None, Some(sas_token)) => match StorageCredentials::sas_token(sas_token) {
Ok(cred) => cred,
Err(err) => {
return Err(format!("Failed to create credentials: {err:?}"));
}
},
_ => {
return Err(concat!(
"Failed to create credentials: exactly one of ",
"'azure-access-key' and 'sas-token' must be specified"
)
.to_string());
}
};
Ok(BlobStore::Azure(Arc::new(AzureStore {
client: ClientBuilder::new(config.storage_account, credentials)
.retry(RetryOptions::exponential(
ExponentialRetryOptions::default().max_retries(config.max_retries as u32 * 2),
))
.container_client(config.container),
prefix: config.key_prefix,
})))
}
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let blob_client = self.client.blob_client(self.build_key(key));
let mut stream = blob_client.get();
let mut buf = if range.end == usize::MAX {
// Let's turn this into a proper RangeFrom.
stream = stream.range(range.start..);
// We don't know how big to expect the result to be.
Vec::new()
} else {
stream = stream.range(range.clone());
Vec::with_capacity(range.end - range.start)
};
let mut stream = stream.into_stream();
while let Some(response) = stream.next().await {
let err = match response {
Ok(chunks) => {
let mut chunks = chunks.data;
let mut err = None;
while let Some(chunk) = chunks.next().await {
match chunk {
Ok(ref data) => {
buf.extend(data);
}
Err(e) => {
err = Some(e);
break;
}
}
}
err
}
Err(e) => Some(e),
};
if let Some(e) = err {
return if matches!(
e.kind(),
ErrorKind::HttpResponse {
status: StatusCode::NotFound,
..
}
) {
Ok(None)
} else {
Err(trc::StoreEvent::AzureError.reason(e))
};
}
}
Ok(Some(buf))
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
let blob_client = self.client.blob_client(self.build_key(key));
// We unfortunately have to make a copy of `data`. This is because the Azure SDK wants to
// coerce the body into a value of type azure_core::Body, which doesn't have a lifetime
// parameter and so cannot hold any non-static references (directly or indirectly).
let data = data.to_vec();
blob_client
.put_block_blob(data)
.into_future()
.await
.map_err(into_error)?;
Ok(())
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
let blob_client = self.client.blob_client(self.build_key(key));
if let Err(e) = blob_client.delete().into_future().await {
if matches!(
e.kind(),
ErrorKind::HttpResponse {
status: StatusCode::NotFound,
..
}
) {
Ok(false)
} else {
Err(trc::StoreEvent::AzureError.reason(e))
}
} else {
Ok(true)
}
}
fn build_key(&self, key: &[u8]) -> String {
if let Some(prefix) = &self.prefix {
let mut writer =
Base32Writer::with_raw_capacity(prefix.len() + (key.len().div_ceil(4) * 5));
writer.push_string(prefix);
writer.write_all(key).unwrap();
writer.finalize()
} else {
Base32Writer::from_bytes(key).finalize()
}
}
}
#[inline(always)]
fn into_error(err: impl Display) -> trc::Error {
trc::StoreEvent::AzureError.reason(err)
}
+154
View File
@@ -0,0 +1,154 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use std::sync::Arc;
use crate::{
SearchStore,
backend::elastic::ElasticSearchStore,
search::{
CalendarSearchField, ContactSearchField, EmailSearchField, SearchableField,
TracingSearchField,
},
};
use registry::schema::structs;
use reqwest::{Error, Response, Url};
use serde_json::{Value, json};
impl ElasticSearchStore {
pub async fn open(config: structs::ElasticSearchStore) -> Result<SearchStore, String> {
Url::parse(&config.url).map_err(|e| format!("Invalid URL: {e}",))?;
Ok(SearchStore::ElasticSearch(Arc::new(Self {
client: config
.http_auth
.build_http_client(
config.http_headers,
"application/json".into(),
config.timeout,
config.allow_invalid_certs,
)
.await?,
url: config.url,
num_replicas: config.num_replicas as usize,
num_shards: config.num_shards as usize,
include_source: config.include_source,
})))
}
pub async fn create_indexes(&self) -> trc::Result<()> {
self.create_index::<EmailSearchField>().await?;
self.create_index::<CalendarSearchField>().await?;
self.create_index::<ContactSearchField>().await?;
self.create_index::<TracingSearchField>().await?;
Ok(())
}
async fn create_index<T: SearchableField>(&self) -> trc::Result<()> {
let mut mappings = serde_json::Map::new();
mappings.insert(
"properties".to_string(),
Value::Object(
T::primary_keys()
.iter()
.chain(T::all_fields())
.map(|field| (field.field_name().to_string(), field.es_schema()))
.collect::<serde_json::Map<String, Value>>(),
),
);
if !self.include_source {
mappings.insert("_source".to_string(), json!({ "enabled": false }));
}
let body = json!({
"mappings": mappings,
"settings": {
"index.number_of_shards": self.num_shards,
"index.number_of_replicas": self.num_replicas,
"analysis": {
"analyzer": {
"default": {
"type": "custom",
"tokenizer": "standard",
"filter": ["lowercase", "stemmer"]
}
}
}
}
});
let response = self
.client
.put(format!("{}/{}", self.url, T::index().index_name()))
.body(body.to_string())
.send()
.await
.map_err(|err| {
trc::StoreEvent::ElasticsearchError
.reason(err)
.details("Failed to create index")
})?;
match response.status().as_u16() {
200..300 => Ok(()),
status @ (400..500) => {
let text = response.text().await.unwrap_or_default();
if text.contains("resource_already_exists_exception") {
// Index already exists, ignore
Ok(())
} else {
Err(trc::StoreEvent::ElasticsearchError
.reason(text)
.ctx(trc::Key::Code, status))
}
}
status => {
let text = response.text().await.unwrap_or_default();
Err(trc::StoreEvent::ElasticsearchError
.reason(text)
.ctx(trc::Key::Code, status))
}
}
}
#[cfg(feature = "test_mode")]
pub async fn drop_indexes(&self) -> trc::Result<()> {
use crate::write::SearchIndex;
for index in &[
SearchIndex::Email,
SearchIndex::Calendar,
SearchIndex::Contacts,
SearchIndex::Tracing,
] {
assert_success(
self.client
.delete(format!("{}/{}", self.url, index.index_name()))
.send()
.await,
)
.await
.map(|_| ())?;
}
Ok(())
}
}
pub(crate) async fn assert_success(response: Result<Response, Error>) -> trc::Result<Response> {
match response {
Ok(response) => {
let status = response.status();
if status.is_success() {
Ok(response)
} else {
Err(trc::StoreEvent::ElasticsearchError
.reason(response.text().await.unwrap_or_default())
.ctx(trc::Key::Code, status.as_u16()))
}
}
Err(err) => Err(trc::StoreEvent::ElasticsearchError.reason(err)),
}
}
+142
View File
@@ -0,0 +1,142 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::search::*;
use reqwest::Client;
use serde::{Deserialize, Deserializer};
use serde_json::{Value, json};
pub mod main;
pub mod search;
pub struct ElasticSearchStore {
client: Client,
url: String,
num_shards: usize,
num_replicas: usize,
include_source: bool,
}
#[derive(Debug, Deserialize)]
pub struct SearchResponse {
pub hits: Hits,
}
#[derive(Debug, Deserialize)]
pub struct Hits {
pub total: Total,
pub hits: Vec<Hit>,
}
#[derive(Debug, Deserialize)]
pub struct Total {
pub value: u64,
}
#[derive(Debug, Deserialize)]
pub struct Hit {
#[serde(rename = "_id", deserialize_with = "deserialize_string_to_u64")]
pub id: u64,
pub sort: Option<Value>,
}
#[derive(Debug, Deserialize)]
pub struct DeleteByQueryResponse {
pub deleted: u64,
}
impl SearchField {
pub fn es_schema(&self) -> Value {
match self {
SearchField::AccountId
| SearchField::DocumentId
| SearchField::Email(EmailSearchField::Size) => json!({
"type": "integer"
}),
SearchField::Id
| SearchField::Email(EmailSearchField::SentAt | EmailSearchField::ReceivedAt)
| SearchField::Calendar(CalendarSearchField::Start)
| SearchField::Tracing(TracingSearchField::QueueId | TracingSearchField::EventType) => {
json!({
"type": "long"
})
}
SearchField::Email(EmailSearchField::HasAttachment) => json!({
"type": "boolean"
}),
SearchField::Calendar(CalendarSearchField::Uid)
| SearchField::Contact(ContactSearchField::Uid) => json!({
"type": "keyword",
}),
SearchField::Email(
EmailSearchField::From | EmailSearchField::To | EmailSearchField::Subject,
) => json!({
"type": "text",
"fields": {
"keyword": {
"type": "keyword"
}
}
}),
SearchField::Email(EmailSearchField::Headers) => {
json!({
"type": "object",
"enabled": true
})
}
#[cfg(feature = "test_mode")]
SearchField::Email(EmailSearchField::Bcc | EmailSearchField::Cc) => {
json!({
"type": "text",
"fields": {
"keyword": {
"type": "keyword"
}
}
})
}
#[cfg(not(feature = "test_mode"))]
SearchField::Email(EmailSearchField::Bcc | EmailSearchField::Cc) => {
json!({
"type": "text"
})
}
SearchField::Email(EmailSearchField::Body | EmailSearchField::Attachment)
| SearchField::Calendar(
CalendarSearchField::Title
| CalendarSearchField::Description
| CalendarSearchField::Location
| CalendarSearchField::Owner
| CalendarSearchField::Attendee,
)
| SearchField::Contact(
ContactSearchField::Member
| ContactSearchField::Kind
| ContactSearchField::Name
| ContactSearchField::Nickname
| ContactSearchField::Organization
| ContactSearchField::Email
| ContactSearchField::Phone
| ContactSearchField::OnlineService
| ContactSearchField::Address
| ContactSearchField::Note,
)
| SearchField::File(FileSearchField::Name | FileSearchField::Content)
| SearchField::Tracing(TracingSearchField::Keywords) => json!({
"type": "text"
}),
}
}
}
fn deserialize_string_to_u64<'de, D>(deserializer: D) -> Result<u64, D::Error>
where
D: Deserializer<'de>,
{
<&str>::deserialize(deserializer)?
.parse::<u64>()
.map_err(serde::de::Error::custom)
}
+377
View File
@@ -0,0 +1,377 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{
backend::elastic::{
DeleteByQueryResponse, ElasticSearchStore, SearchResponse, main::assert_success,
},
search::{
IndexDocument, SearchComparator, SearchDocumentId, SearchField, SearchFilter,
SearchOperator, SearchQuery, SearchValue,
},
write::SearchIndex,
};
use serde_json::{Map, Value, json};
use std::fmt::Write;
impl ElasticSearchStore {
pub async fn index(&self, documents: Vec<IndexDocument>) -> trc::Result<()> {
let mut request = String::with_capacity(512);
for document in documents {
let id = if let (Some(SearchValue::Uint(account_id)), Some(SearchValue::Uint(doc_id))) = (
document.fields.get(&SearchField::AccountId),
document.fields.get(&SearchField::DocumentId),
) {
*account_id << 32 | *doc_id
} else if let Some(SearchValue::Uint(id)) = document.fields.get(&SearchField::Id) {
*id
} else {
debug_assert!(false, "Document is missing required ID fields");
continue;
};
let _ = writeln!(
&mut request,
"{{\"index\":{{\"_index\":\"{}\",\"_id\":{id}}}}}",
document.index.index_name()
);
json_serialize(&mut request, &document);
request.push('\n');
}
assert_success(
self.client
.post(format!("{}/_bulk", self.url))
.body(request)
.send()
.await,
)
.await
.map(|_| ())
}
pub async fn query<R: SearchDocumentId>(
&self,
index: SearchIndex,
filters: &[SearchFilter],
sort: &[SearchComparator],
) -> trc::Result<Vec<R>> {
let mut search_after: Option<Value> = None;
let mut results = Vec::new();
let mut has_more = true;
while has_more {
let query = Map::from_iter(
[
Some(("query".to_string(), build_query(filters))),
Some(("size".to_string(), Value::from(10_000))),
Some(("_source".to_string(), Value::from(false))),
Some((
"sort".to_string(),
build_sort(sort, R::field().field_name()),
)),
search_after
.take()
.map(|sa| ("search_after".to_string(), sa)),
]
.into_iter()
.flatten(),
);
let response = assert_success(
self.client
.post(format!("{}/{}/_search", self.url, index.index_name()))
.body(serde_json::to_string(&query).unwrap_or_default())
.send()
.await,
)
.await?;
let text = response
.text()
.await
.map_err(|err| trc::StoreEvent::ElasticsearchError.reason(err))?;
let response = serde_json::from_str::<SearchResponse>(&text).map_err(|err| {
trc::StoreEvent::ElasticsearchError
.reason(err)
.details(text)
})?;
has_more = response.hits.hits.len() == 10_000
&& response.hits.hits.last().unwrap().sort.is_some();
for hit in response.hits.hits {
search_after = hit.sort;
results.push(R::from_u64(hit.id));
}
}
Ok(results)
}
pub async fn unindex(&self, filter: SearchQuery) -> trc::Result<u64> {
if filter.filters.is_empty() {
return Err(trc::StoreEvent::ElasticsearchError
.reason("Unindex operation requires at least one filter"));
}
let query = json!({
"query": build_query(&filter.filters),
});
let response = assert_success(
self.client
.post(format!(
"{}/{}/_delete_by_query",
self.url,
filter.index.index_name()
))
.body(serde_json::to_string(&query).unwrap_or_default())
.send()
.await,
)
.await?;
let response_body = response
.text()
.await
.map_err(|err| trc::StoreEvent::ElasticsearchError.reason(err))?;
serde_json::from_str::<DeleteByQueryResponse>(&response_body)
.map(|delete_response| delete_response.deleted)
.map_err(|err| trc::StoreEvent::ElasticsearchError.reason(err))
}
pub async fn refresh_index(&self, index: SearchIndex) -> trc::Result<()> {
let url = format!("{}/{}/_refresh", self.url, index.index_name());
assert_success(self.client.post(url).send().await)
.await
.map(|_| ())
}
}
fn build_query(filters: &[SearchFilter]) -> Value {
if filters.is_empty() {
return json!({ "match_all": {} });
}
let mut stack = Vec::new();
let mut conditions = Vec::new();
let mut logical_op = &SearchFilter::And;
for filter in filters {
match filter {
SearchFilter::Operator { field, op, value } => {
if field.is_text() && matches!(op, SearchOperator::Equal | SearchOperator::Contains)
{
let SearchValue::Text { value, .. } = value else {
debug_assert!(false, "Invalid value type for text field");
continue;
};
if op != &SearchOperator::Equal {
conditions.push(json!({
"match": { field.field_name(): {
"query": value,
"operator": "and"
} }
}));
} else {
conditions.push(json!({
"match_phrase": { field.field_name(): value }
}));
}
} else {
let value = match value {
SearchValue::Text { value, .. } => json!(value),
SearchValue::Int(value) => json!(value),
SearchValue::Uint(value) => json!(value),
SearchValue::Boolean(value) => json!(value),
SearchValue::KeyValues(kv) => {
let (key, value) = kv.iter().next().unwrap();
let cond = if !value.is_empty() {
if op == &SearchOperator::Equal {
json!({
"term": {
format!("{}.{}.keyword", field.field_name(), key): value
}
})
} else {
json!({
"match": {
format!("{}.{}", field.field_name(), key): value
}
})
}
} else {
json!({
"exists": { "field": format!("{}.{}", field.field_name(), key) }
})
};
conditions.push(cond);
continue;
}
};
let cond = match op {
SearchOperator::Equal | SearchOperator::Contains => json!({
"term": { field.field_name(): value }
}),
op => {
let op = match op {
SearchOperator::LowerThan => "lt",
SearchOperator::LowerEqualThan => "lte",
SearchOperator::GreaterThan => "gt",
SearchOperator::GreaterEqualThan => "gte",
_ => unreachable!(),
};
json!({
"range": { field.field_name(): { op: value } }
})
}
};
conditions.push(cond);
}
}
SearchFilter::And | SearchFilter::Or | SearchFilter::Not => {
stack.push((logical_op, conditions));
logical_op = filter;
conditions = Vec::new();
}
SearchFilter::End => {
if let Some((prev_logical_op, mut prev_conditions)) = stack.pop() {
if !conditions.is_empty() {
match logical_op {
SearchFilter::And => {
prev_conditions.push(json!({ "bool": { "must": conditions } }));
}
SearchFilter::Or => {
prev_conditions.push(json!({ "bool": { "should": conditions } }));
}
SearchFilter::Not => {
prev_conditions.push(json!({ "bool": { "must_not": conditions } }));
}
_ => unreachable!(),
}
}
logical_op = prev_logical_op;
conditions = prev_conditions;
}
}
SearchFilter::DocumentSet(_) => {
debug_assert!(
false,
"DocumentSet filters are not supported in this backend"
);
continue;
}
}
}
debug_assert!(
!conditions.is_empty(),
"No conditions were built for the query"
);
if conditions.len() == 1 {
conditions.pop().unwrap()
} else {
json!({ "bool": { "must": conditions } })
}
}
fn build_sort(sort: &[SearchComparator], tie_breaker: &str) -> Value {
Value::Array(
sort.iter()
.filter_map(|comp| match comp {
SearchComparator::Field { field, ascending } => {
let field = if field.is_text() {
format!("{}.keyword", field.field_name())
} else {
field.field_name().to_string()
};
Some(json!({
field: if *ascending { "asc" } else { "desc" }
}))
}
_ => None,
})
.chain([json!({
tie_breaker: "asc"
})])
.collect(),
)
}
fn json_serialize(request: &mut String, document: &IndexDocument) {
request.push('{');
for (idx, (k, v)) in document.fields.iter().enumerate() {
if idx > 0 {
request.push(',');
}
let _ = write!(request, "{:?}:", k.field_name());
match v {
SearchValue::Text { value, .. } => {
json_serialize_str(request, value);
}
SearchValue::KeyValues(map) => {
request.push('{');
for (i, (key, value)) in map.iter().enumerate() {
if i > 0 {
request.push(',');
}
json_serialize_str(request, key);
request.push(':');
json_serialize_str(request, value);
}
request.push('}');
}
SearchValue::Int(v) => {
let _ = write!(request, "{}", v);
}
SearchValue::Uint(v) => {
let _ = write!(request, "{}", v);
}
SearchValue::Boolean(v) => {
let _ = write!(request, "{}", v);
}
}
}
request.push('}');
}
fn json_serialize_str(request: &mut String, value: &str) {
request.push('"');
for c in value.chars() {
match c {
'"' => request.push_str("\\\""),
'\\' => request.push_str("\\\\"),
'\n' => request.push_str("\\n"),
'\r' => request.push_str("\\r"),
'\t' => request.push_str("\\t"),
'\u{0008}' => request.push_str("\\b"), // backspace
'\u{000C}' => request.push_str("\\f"), // form feed
_ => {
if !c.is_control() {
request.push(c);
} else {
let _ = write!(request, "\\u{:04x}", c as u32);
}
}
}
}
request.push('"');
}
@@ -0,0 +1,51 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::EphemeralStore;
use crate::SUBSPACE_BLOBS;
use std::ops::Range;
impl EphemeralStore {
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let state = self.state.read();
Ok(state
.subspaces
.get(&SUBSPACE_BLOBS)
.and_then(|m| m.get(key))
.map(|bytes| {
if range.start == 0 && range.end == usize::MAX {
bytes.clone()
} else {
bytes
.get(range.start..std::cmp::min(bytes.len(), range.end))
.unwrap_or_default()
.to_vec()
}
}))
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
let mut state = self.state.write();
state
.subspaces
.entry(SUBSPACE_BLOBS)
.or_default()
.insert(key.to_vec(), data.to_vec());
Ok(())
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
let mut state = self.state.write();
if let Some(map) = state.subspaces.get_mut(&SUBSPACE_BLOBS) {
map.remove(key);
}
Ok(true)
}
}
@@ -0,0 +1,21 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{EphemeralState, EphemeralStore};
use crate::Store;
use ahash::AHashMap;
use parking_lot::RwLock;
use std::sync::Arc;
impl EphemeralStore {
pub fn open() -> Store {
Store::Ephemeral(Arc::new(EphemeralStore {
state: RwLock::new(EphemeralState {
subspaces: AHashMap::new(),
}),
}))
}
}
+22
View File
@@ -0,0 +1,22 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
pub mod blob;
pub mod main;
pub mod read;
pub mod write;
use ahash::AHashMap;
use parking_lot::RwLock;
use std::collections::BTreeMap;
pub struct EphemeralStore {
pub(crate) state: RwLock<EphemeralState>,
}
pub(crate) struct EphemeralState {
pub(crate) subspaces: AHashMap<u8, BTreeMap<Vec<u8>, Vec<u8>>>,
}
@@ -0,0 +1,86 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::EphemeralStore;
use crate::{Deserialize, IterateParams, Key, ValueKey, write::ValueClass};
impl EphemeralStore {
pub(crate) async fn get_value<U>(&self, key: impl Key) -> trc::Result<Option<U>>
where
U: Deserialize + 'static,
{
let subspace = key.subspace();
let key_bytes = key.serialize(0);
let state = self.state.read();
match state
.subspaces
.get(&subspace)
.and_then(|m| m.get(&key_bytes))
{
Some(value) => U::deserialize_with_key(&key_bytes, value).map(Some),
None => Ok(None),
}
}
pub(crate) async fn key_exists(&self, key: impl Key) -> trc::Result<bool> {
let subspace = key.subspace();
let key_bytes = key.serialize(0);
let state = self.state.read();
Ok(state
.subspaces
.get(&subspace)
.is_some_and(|m| m.contains_key(&key_bytes)))
}
pub(crate) async fn iterate<T: Key>(
&self,
params: IterateParams<T>,
mut cb: impl for<'x> FnMut(&'x [u8], &'x [u8]) -> trc::Result<bool> + Sync + Send,
) -> trc::Result<()> {
let subspace = params.begin.subspace();
let begin = params.begin.serialize(0);
let end = params.end.serialize(0);
let state = self.state.read();
let Some(map) = state.subspaces.get(&subspace) else {
return Ok(());
};
if params.ascending {
for (k, v) in map.range(begin..=end) {
if !cb(k.as_slice(), v.as_slice())? || params.first {
break;
}
}
} else {
for (k, v) in map.range(begin..=end).rev() {
if !cb(k.as_slice(), v.as_slice())? || params.first {
break;
}
}
}
Ok(())
}
pub(crate) async fn get_counter(
&self,
key: impl Into<ValueKey<ValueClass>> + Sync + Send,
) -> trc::Result<i64> {
let key = key.into();
let subspace = key.subspace();
let key_bytes = key.serialize(0);
let state = self.state.read();
match state
.subspaces
.get(&subspace)
.and_then(|m| m.get(&key_bytes))
{
Some(bytes) => Ok(i64::from_le_bytes(bytes[..].try_into().map_err(|_| {
trc::Error::corrupted_key(&key_bytes, Some(bytes.as_slice()), trc::location!())
})?)),
None => Ok(0),
}
}
}
+200
View File
@@ -0,0 +1,200 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::EphemeralStore;
use crate::{
IndexKey, Key, LogKey, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER, SUBSPACE_INDEXES,
SUBSPACE_LOGS, SUBSPACE_QUOTA,
backend::deserialize_i64_le,
write::{AssignedIds, Batch, MergeResult, Operation, ValueClass, ValueOp},
};
impl EphemeralStore {
pub(crate) async fn write(&self, batch: Batch<'_>) -> trc::Result<AssignedIds> {
let mut account_id = u32::MAX;
let mut collection = u8::MAX;
let mut document_id = u32::MAX;
let mut change_id = 0u64;
let mut result = AssignedIds::default();
let has_changes = !batch.changes.is_empty();
let mut state = self.state.write();
if has_changes {
let map = state.subspaces.entry(SUBSPACE_COUNTER).or_default();
for &account_id in batch.changes.keys() {
let key = ValueClass::ChangeId.serialize(account_id, 0, 0, 0);
let next = match map.get(&key) {
Some(bytes) => deserialize_i64_le(&key, bytes)? + 1,
None => 1,
};
map.insert(key, next.to_le_bytes().to_vec());
result.push_change_id(account_id, next as u64);
}
}
for op in batch.ops.iter_mut() {
match op {
Operation::AccountId {
account_id: account_id_,
} => {
account_id = *account_id_;
if has_changes {
change_id = result.set_current_change_id(account_id)?;
}
}
Operation::Collection {
collection: collection_,
} => {
collection = u8::from(*collection_);
}
Operation::DocumentId {
document_id: document_id_,
} => {
document_id = *document_id_;
}
Operation::Value { class, op } => {
let subspace = class.subspace(collection);
let key = class.serialize(account_id, collection, document_id, 0);
let map = state.subspaces.entry(subspace).or_default();
match op {
ValueOp::Set(value) => {
map.insert(key, std::mem::take(value));
}
ValueOp::SetFnc(set_op) => {
let value = (set_op.fnc)(&set_op.params, &result)?;
map.insert(key, value);
}
ValueOp::MergeFnc(merge_op) => {
let merge_result = (merge_op.fnc)(
&merge_op.params,
&result,
map.get(&key).map(|v| v.as_slice()),
)?;
match merge_result {
MergeResult::Update(value) => {
map.insert(key, value);
}
MergeResult::Delete => {
map.remove(&key);
}
MergeResult::Skip => (),
}
}
ValueOp::AtomicAdd(by) => {
let current = match map.get(&key) {
Some(bytes) => deserialize_i64_le(&key, bytes)?,
None => 0,
};
let next = current + *by;
map.insert(key, next.to_le_bytes().to_vec());
}
ValueOp::AddAndGet(by) => {
let current = match map.get(&key) {
Some(bytes) => deserialize_i64_le(&key, bytes)?,
None => 0,
};
let next = current + *by;
map.insert(key, next.to_le_bytes().to_vec());
result.push_counter_id(next);
}
ValueOp::Clear => {
map.remove(&key);
}
}
}
Operation::Index { field, key, set } => {
let index_key = IndexKey {
account_id,
collection,
document_id,
field: *field,
key: key.as_slice(),
}
.serialize(0);
let map = state.subspaces.entry(SUBSPACE_INDEXES).or_default();
if *set {
map.insert(index_key, Vec::new());
} else {
map.remove(&index_key);
}
}
Operation::Log { collection, set } => {
let log_key = LogKey {
account_id,
collection: u8::from(*collection),
change_id,
}
.serialize(0);
let map = state.subspaces.entry(SUBSPACE_LOGS).or_default();
map.insert(log_key, std::mem::take(set));
}
Operation::AssertValue {
class,
assert_value,
} => {
let subspace = class.subspace(collection);
let key = class.serialize(account_id, collection, document_id, 0);
let matches = state
.subspaces
.get(&subspace)
.and_then(|m| m.get(&key))
.map(|v| assert_value.matches(v.as_slice()))
.unwrap_or_else(|| assert_value.is_none());
if !matches {
return Err(trc::StoreEvent::AssertValueFailed.into());
}
}
}
}
Ok(result)
}
pub(crate) async fn delete_range(&self, from: impl Key, to: impl Key) -> trc::Result<()> {
let subspace = from.subspace();
let from_key = from.serialize(0);
let to_key = to.serialize(0);
let mut state = self.state.write();
if let Some(map) = state.subspaces.get_mut(&subspace) {
let keys: Vec<Vec<u8>> = map
.range(from_key..to_key)
.map(|(k, _)| k.clone())
.collect();
for k in keys {
map.remove(&k);
}
}
Ok(())
}
pub(crate) async fn purge_store(&self) -> trc::Result<()> {
let mut state = self.state.write();
for subspace in [SUBSPACE_QUOTA, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER] {
if let Some(map) = state.subspaces.get_mut(&subspace) {
let keys: Vec<Vec<u8>> = map
.iter()
.filter_map(|(k, v)| {
if v.len() == std::mem::size_of::<i64>()
&& i64::from_le_bytes(v[..].try_into().unwrap()) == 0
{
Some(k.clone())
} else {
None
}
})
.collect();
for k in keys {
map.remove(&k);
}
}
}
Ok(())
}
}
@@ -0,0 +1,156 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{FdbStore, MAX_VALUE_SIZE};
use crate::{
IterateParams, SUBSPACE_BLOBS,
backend::foundationdb::into_error,
write::{AnyKey, key::KeySerializer},
};
use std::ops::Range;
use trc::AddContext;
use types::blob_hash::BLOB_HASH_LEN;
impl FdbStore {
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let block_start = range.start / MAX_VALUE_SIZE;
let bytes_start = range.start % MAX_VALUE_SIZE;
let block_end = (range.end / MAX_VALUE_SIZE) + 1;
let begin = KeySerializer::new(key.len() + 2)
.write(key)
.write(block_start as u16)
.finalize();
let end = KeySerializer::new(key.len() + 2)
.write(key)
.write(block_end as u16)
.finalize();
let key_len = begin.len();
let mut blob_data: Option<Vec<u8>> = None;
let blob_range = range.end - range.start;
self.iterate(
IterateParams::new(
AnyKey {
subspace: SUBSPACE_BLOBS,
key: begin,
},
AnyKey {
subspace: SUBSPACE_BLOBS,
key: end,
},
),
|key, value| {
if key.len() == key_len {
if let Some(blob_data) = &mut blob_data {
blob_data.extend_from_slice(
value
.get(
..std::cmp::min(
blob_range.saturating_sub(blob_data.len()),
value.len(),
),
)
.unwrap_or(&[]),
);
if blob_data.len() == blob_range {
return Ok(false);
}
} else {
let blob_size = if blob_range <= (5 * (1 << 20)) {
blob_range
} else if value.len() == MAX_VALUE_SIZE {
MAX_VALUE_SIZE * 2
} else {
value.len()
};
let mut blob_data_ = Vec::with_capacity(blob_size);
blob_data_.extend_from_slice(
value
.get(
bytes_start
..std::cmp::min(bytes_start + blob_range, value.len()),
)
.unwrap_or(&[]),
);
let is_done = blob_data_.len() == blob_range;
blob_data = blob_data_.into();
if is_done {
return Ok(false);
}
}
}
Ok(true)
},
)
.await
.caused_by(trc::location!())?;
Ok(blob_data)
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
const N_CHUNKS: usize = (1 << 5) - 1;
let last_chunk = std::cmp::max(
(data.len() / MAX_VALUE_SIZE)
+ if !data.len().is_multiple_of(MAX_VALUE_SIZE) {
1
} else {
0
},
1,
) - 1;
let mut trx = self.db.create_trx().map_err(into_error)?;
for (chunk_pos, chunk_bytes) in data.chunks(MAX_VALUE_SIZE).enumerate() {
trx.set(
&KeySerializer::new(key.len() + 3)
.write(SUBSPACE_BLOBS)
.write(key)
.write(chunk_pos as u16)
.finalize(),
chunk_bytes,
);
if chunk_pos == last_chunk || (chunk_pos > 0 && chunk_pos % N_CHUNKS == 0) {
self.commit(trx, false).await?;
if chunk_pos < last_chunk {
trx = self.db.create_trx().map_err(into_error)?;
} else {
break;
}
}
}
Ok(())
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
if key.len() < BLOB_HASH_LEN {
return Ok(false);
}
let trx = self.db.create_trx().map_err(into_error)?;
trx.clear_range(
&KeySerializer::new(key.len() + 3)
.write(SUBSPACE_BLOBS)
.write(key)
.write(0u16)
.finalize(),
&KeySerializer::new(key.len() + 3)
.write(SUBSPACE_BLOBS)
.write(key)
.write(u16::MAX)
.finalize(),
);
self.commit(trx, false).await
}
}
@@ -0,0 +1,65 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::FdbStore;
use crate::Store;
use foundationdb::{Database, api, api::NetworkAutoStop, options::DatabaseOption};
use parking_lot::Mutex;
use registry::schema::structs;
use std::sync::Arc;
static FDB_NETWORK: Mutex<Option<NetworkAutoStop>> = Mutex::new(None);
impl FdbStore {
pub async fn open(config: structs::FoundationDbStore) -> Result<Store, String> {
{
let mut guard = FDB_NETWORK.lock();
if guard.is_none() {
let network = unsafe {
api::FdbApiBuilder::default()
.build()
.map_err(|err| format!("Failed to boot FoundationDB: {err:?}"))?
.boot()
.map_err(|err| format!("Failed to boot FoundationDB: {err:?}"))?
};
*guard = Some(network);
}
}
let db = Database::new(config.cluster_file.as_deref())
.map_err(|err| format!("Failed to create FoundationDB database: {err:?}"))?;
if let Some(value) = config.transaction_timeout {
db.set_option(DatabaseOption::TransactionTimeout(
value.into_inner().as_millis() as i32,
))
.map_err(|err| format!("Failed to set option: {err:?}"))?;
}
if let Some(value) = config.transaction_retry_limit {
db.set_option(DatabaseOption::TransactionRetryLimit(value as i32))
.map_err(|err| format!("Failed to set option: {err:?}"))?;
}
if let Some(value) = config.transaction_retry_delay {
db.set_option(DatabaseOption::TransactionMaxRetryDelay(
value.into_inner().as_millis() as i32,
))
.map_err(|err| format!("Failed to set option: {err:?}"))?;
}
if let Some(value) = config.machine_id {
db.set_option(DatabaseOption::MachineId(value))
.map_err(|err| format!("Failed to set option: {err:?}"))?;
}
if let Some(value) = config.datacenter_id {
db.set_option(DatabaseOption::DatacenterId(value))
.map_err(|err| format!("Failed to set option: {err:?}"))?;
}
Ok(Store::FoundationDb(Arc::new(Self {
db,
version: Default::default(),
})))
}
}
@@ -0,0 +1,114 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use foundationdb::{Database, FdbError};
use std::{
sync::atomic::{AtomicBool, AtomicI64, AtomicU64, Ordering},
time::{Duration, Instant},
};
pub mod blob;
pub mod main;
pub mod read;
pub mod write;
const MAX_VALUE_SIZE: usize = 100000;
const REFRESH_READ_VERSION_AFTER: Duration = Duration::from_secs(1);
const MAX_READ_VERSION_AGE: Duration = Duration::from_secs(4);
pub struct FdbStore {
db: Database,
version: ReadVersion,
}
pub(crate) struct ReadVersion {
base: Instant,
version: AtomicI64,
obtained: AtomicU64,
refreshing: AtomicBool,
}
impl ReadVersion {
fn now(&self) -> u64 {
self.base.elapsed().as_nanos() as u64
}
fn current(&self) -> i64 {
self.version.load(Ordering::Acquire)
}
fn age(&self) -> u64 {
self.now()
.saturating_sub(self.obtained.load(Ordering::Acquire))
}
fn store_max(&self, version: i64) {
let mut current = self.version.load(Ordering::Relaxed);
while version > current {
match self.version.compare_exchange_weak(
current,
version,
Ordering::Release,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(actual) => current = actual,
}
}
}
fn refreshed(&self, version: i64) {
self.store_max(version);
self.obtained.store(self.now(), Ordering::Release);
}
fn raise_floor(&self, version: i64) {
self.store_max(version);
}
fn expire(&self) {
self.obtained.store(0, Ordering::Release);
}
fn try_begin_refresh(&self) -> Option<RefreshGuard<'_>> {
if self
.refreshing
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Relaxed)
.is_ok()
{
Some(RefreshGuard(&self.refreshing))
} else {
None
}
}
}
impl Default for ReadVersion {
fn default() -> Self {
Self {
base: Instant::now(),
version: AtomicI64::new(0),
obtained: AtomicU64::new(0),
refreshing: AtomicBool::new(false),
}
}
}
pub(crate) struct RefreshGuard<'a>(&'a AtomicBool);
impl Drop for RefreshGuard<'_> {
fn drop(&mut self) {
self.0.store(false, Ordering::Release);
}
}
#[inline(always)]
fn into_error(error: FdbError) -> trc::Error {
trc::StoreEvent::FoundationdbError
.reason(error.message())
.ctx(trc::Key::Code, error.code())
}
@@ -0,0 +1,334 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{
FdbStore, MAX_READ_VERSION_AGE, MAX_VALUE_SIZE, REFRESH_READ_VERSION_AFTER, into_error,
};
use crate::{
Deserialize, IterateParams, Key, ValueKey, WITH_SUBSPACE,
backend::deserialize_i64_le,
write::{MAX_COMMIT_ATTEMPTS, MAX_COMMIT_TIME, ValueClass, key::KeySerializer},
};
use foundationdb::{
FdbError, KeySelector, RangeOption, Transaction,
future::FdbSlice,
options::{self},
};
use futures::TryStreamExt;
use std::time::Instant;
#[allow(dead_code)]
pub(crate) enum ChunkedValue {
Single(FdbSlice),
Chunked { n_chunks: u8, bytes: Vec<u8> },
None,
}
struct ChunkedValueCollector {
key: Vec<u8>,
bytes: Vec<u8>,
}
impl FdbStore {
pub(crate) async fn get_value<U>(&self, key: impl Key) -> trc::Result<Option<U>>
where
U: Deserialize,
{
let key = key.serialize(WITH_SUBSPACE);
let mut retry_count = 0;
let start = Instant::now();
loop {
let trx = self.read_trx().await?;
match read_chunked_value(&key, &trx, true).await {
Ok(ChunkedValue::Single(bytes)) => {
return U::deserialize_with_key(key.get(1..).unwrap_or_default(), &bytes)
.map(Some);
}
Ok(ChunkedValue::Chunked { bytes, .. }) => {
return U::deserialize_owned_with_key(key.get(1..).unwrap_or_default(), bytes)
.map(Some);
}
Ok(ChunkedValue::None) => return Ok(None),
Err(err) => {
self.on_read_error(trx, err, &mut retry_count, start)
.await?;
}
}
}
}
pub(crate) async fn key_exists(&self, key: impl Key) -> trc::Result<bool> {
let key = key.serialize(WITH_SUBSPACE);
let mut retry_count = 0;
let start = Instant::now();
loop {
let trx = self.read_trx().await?;
match read_chunked_value(&key, &trx, true).await {
Ok(ChunkedValue::Single(_) | ChunkedValue::Chunked { .. }) => return Ok(true),
Ok(ChunkedValue::None) => return Ok(false),
Err(err) => {
self.on_read_error(trx, err, &mut retry_count, start)
.await?;
}
}
}
}
pub(crate) async fn iterate<T: Key>(
&self,
params: IterateParams<T>,
mut cb: impl for<'x> FnMut(&'x [u8], &'x [u8]) -> trc::Result<bool> + Sync + Send,
) -> trc::Result<()> {
let begin = params.begin.serialize(WITH_SUBSPACE);
let end = params.end.serialize(WITH_SUBSPACE);
let mut retry_count = 0;
let start = Instant::now();
if !params.first {
let mut last_key = vec![];
let mut chunked_key: Option<ChunkedValueCollector> = None;
'outer: loop {
let begin_selector = if last_key.is_empty() {
KeySelector::first_greater_or_equal(&begin)
} else {
KeySelector::first_greater_than(&last_key)
};
let trx = self.read_trx().await?;
let mut values = trx.get_ranges(
RangeOption {
begin: begin_selector,
end: KeySelector::first_greater_than(&end),
mode: options::StreamingMode::WantAll,
reverse: !params.ascending,
..Default::default()
},
true,
);
let mut last_key_ = vec![];
loop {
match values.try_next().await {
Ok(Some(values)) => {
let mut key = &[] as &[u8];
for value in values.iter() {
key = value.key();
// Check whether we are collecting a chunked value
let cb_key = key.get(1..).unwrap_or_default();
let cb_value = value.value();
if let Some(chunk) = &mut chunked_key {
if chunk.key.len() + 1 == cb_key.len()
&& cb_key[..chunk.key.len()] == chunk.key[..]
{
// This is a chunk of the current value
if params.values {
chunk.bytes.extend_from_slice(cb_value);
}
continue;
} else {
// Return collected chunked value
if !cb(&chunk.key, &chunk.bytes)? {
return Ok(());
}
// Reset collector
chunked_key = None;
}
}
if cb_value.len() < MAX_VALUE_SIZE {
if !cb(cb_key, cb_value)? {
return Ok(());
}
} else {
// Start collecting chunked value
chunked_key = Some(ChunkedValueCollector {
key: cb_key.to_vec(),
bytes: if params.values {
cb_value.to_vec()
} else {
Vec::new()
},
});
}
}
if values.more() {
last_key_ = key.to_vec();
}
}
Ok(None) => {
// Return any chunked value collected
if let Some(chunked_key) = chunked_key.take() {
cb(&chunked_key.key, &chunked_key.bytes)?;
}
break 'outer;
}
Err(e) => {
drop(values);
if e.code() == 1007 && !last_key_.is_empty() {
// Transaction is too old to perform reads or be committed
last_key = last_key_;
continue 'outer;
} else if e.is_retryable()
&& retry_count < MAX_COMMIT_ATTEMPTS
&& start.elapsed() < MAX_COMMIT_TIME
{
// Transient error such as a cached read version ahead of lagging
// storage servers (code 1009); resume from the last key read,
// refresh the read version and back off before retrying.
if !last_key_.is_empty() {
last_key = last_key_;
}
self.version.expire();
trx.on_error(e).await.map_err(into_error)?;
retry_count += 1;
continue 'outer;
} else {
return Err(into_error(e));
}
}
}
}
}
} else {
loop {
let trx = self.read_trx().await?;
let mut values = trx.get_ranges_keyvalues(
RangeOption {
begin: KeySelector::first_greater_or_equal(&begin),
end: KeySelector::first_greater_than(&end),
mode: options::StreamingMode::Small,
reverse: !params.ascending,
..Default::default()
},
true,
);
match values.try_next().await {
Ok(Some(value)) => {
cb(value.key().get(1..).unwrap_or_default(), value.value())?;
break;
}
Ok(None) => break,
Err(e) => {
drop(values);
self.on_read_error(trx, e, &mut retry_count, start).await?;
}
}
}
}
Ok(())
}
pub(crate) async fn get_counter(
&self,
key: impl Into<ValueKey<ValueClass>> + Sync + Send,
) -> trc::Result<i64> {
let key = key.into().serialize(WITH_SUBSPACE);
let mut retry_count = 0;
let start = Instant::now();
loop {
let trx = self.read_trx().await?;
match trx.get(&key, true).await {
Ok(Some(bytes)) => return deserialize_i64_le(&key, &bytes),
Ok(None) => return Ok(0),
Err(e) => {
self.on_read_error(trx, e, &mut retry_count, start).await?;
}
}
}
}
async fn on_read_error(
&self,
trx: Transaction,
err: FdbError,
retry_count: &mut u32,
start: Instant,
) -> trc::Result<()> {
if err.is_retryable()
&& *retry_count < MAX_COMMIT_ATTEMPTS
&& start.elapsed() < MAX_COMMIT_TIME
{
// The cached read version may be ahead of lagging storage servers under heavy write
// load (code 1009); expire it so the retry obtains a fresh read version, then let
// FoundationDB back off before retrying.
self.version.expire();
trx.on_error(err).await.map_err(into_error)?;
*retry_count += 1;
Ok(())
} else {
Err(into_error(err))
}
}
pub(crate) async fn read_trx(&self) -> trc::Result<Transaction> {
let trx = self.db.create_trx().map_err(into_error)?;
let version = self.version.current();
let age = self.version.age();
if version != 0 && age < MAX_READ_VERSION_AGE.as_nanos() as u64 {
if age >= REFRESH_READ_VERSION_AFTER.as_nanos() as u64
&& let Some(_guard) = self.version.try_begin_refresh()
{
let read_version = trx.get_read_version().await.map_err(into_error)?;
self.version.refreshed(read_version);
} else {
trx.set_read_version(version);
}
} else {
let read_version = trx.get_read_version().await.map_err(into_error)?;
self.version.refreshed(read_version);
}
Ok(trx)
}
pub(crate) fn invalidate_read_snapshot(&self) {
self.version.expire();
}
}
pub(crate) async fn read_chunked_value(
key: &[u8],
trx: &Transaction,
snapshot: bool,
) -> Result<ChunkedValue, FdbError> {
if let Some(bytes) = trx.get(key, snapshot).await? {
if bytes.len() < MAX_VALUE_SIZE {
Ok(ChunkedValue::Single(bytes))
} else {
let mut value = Vec::with_capacity(bytes.len() * 2);
value.extend_from_slice(&bytes);
let mut key = KeySerializer::new(key.len() + 1)
.write(key)
.write(0u8)
.finalize();
while let Some(bytes) = trx.get(&key, snapshot).await? {
value.extend_from_slice(&bytes);
*key.last_mut().unwrap() += 1;
}
Ok(ChunkedValue::Chunked {
bytes: value,
n_chunks: *key.last().unwrap(),
})
}
} else {
Ok(ChunkedValue::None)
}
}
@@ -0,0 +1,436 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{
FdbStore, MAX_VALUE_SIZE, into_error,
read::{ChunkedValue, read_chunked_value},
};
use crate::{
backend::deserialize_i64_le,
write::{
AssignedIds, Batch, IndexPropertyClass, MAX_COMMIT_ATTEMPTS, MAX_COMMIT_TIME, MergeResult,
Operation, QueueClass, RegistryClass, SearchIndexType, TaskQueueClass, TelemetryClass,
ValueClass, ValueOp, key::KeySerializer,
},
*,
};
use foundationdb::{
FdbError, KeySelector, RangeOption, Transaction,
options::{self, MutationType},
};
use futures::TryStreamExt;
use rand::RngExt;
use std::{
borrow::Cow,
cmp::Ordering,
time::{Duration, Instant},
};
use trc::AddContext;
impl FdbStore {
pub(crate) async fn write(&self, batch: Batch<'_>) -> trc::Result<AssignedIds> {
let start = Instant::now();
let mut retry_count = 0;
let has_changes = !batch.changes.is_empty();
loop {
let mut account_id = u32::MAX;
let mut collection = u8::MAX;
let mut document_id = u32::MAX;
let mut change_id = 0u64;
let mut result = AssignedIds::default();
let trx = self.db.create_trx().map_err(into_error)?;
if has_changes {
for &account_id in batch.changes.keys() {
debug_assert!(account_id != u32::MAX);
let key = ValueClass::ChangeId.serialize(account_id, 0, 0, WITH_SUBSPACE);
let change_id =
if let Some(bytes) = trx.get(&key, false).await.map_err(into_error)? {
deserialize_i64_le(&key, &bytes)? + 1
} else {
1
};
trx.set(&key, &change_id.to_le_bytes()[..]);
result.push_change_id(account_id, change_id as u64);
}
}
for op in batch.ops.iter_mut() {
match op {
Operation::AccountId {
account_id: account_id_,
} => {
account_id = *account_id_;
if has_changes {
change_id = result.set_current_change_id(account_id)?;
}
}
Operation::Collection {
collection: collection_,
} => {
collection = u8::from(*collection_);
}
Operation::DocumentId {
document_id: document_id_,
} => {
document_id = *document_id_;
}
Operation::Value { class, op } => {
let mut key =
class.serialize(account_id, collection, document_id, WITH_SUBSPACE);
match op {
ValueOp::Set(value) => {
if !chunk_value(&trx, &mut key, value, class, None).await {
trx.cancel();
return Err(trc::StoreEvent::FoundationdbError
.ctx(trc::Key::Reason, "Value is too large"));
}
}
ValueOp::SetFnc(set_op) => {
let value = (set_op.fnc)(&set_op.params, &result)?;
if !chunk_value(&trx, &mut key, &value, class, None).await {
trx.cancel();
return Err(trc::StoreEvent::FoundationdbError
.ctx(trc::Key::Reason, "Value is too large"));
}
}
ValueOp::MergeFnc(merge_op) => {
let (merge_result, prev_num_chunks) =
match read_chunked_value(&key, &trx, false)
.await
.map_err(into_error)
.caused_by(trc::location!())?
{
ChunkedValue::Single(slice) => (
(merge_op.fnc)(
&merge_op.params,
&result,
Some(slice.as_ref()),
)?,
1,
),
ChunkedValue::Chunked { bytes, n_chunks } => (
(merge_op.fnc)(
&merge_op.params,
&result,
Some(bytes.as_ref()),
)?,
n_chunks as usize + 1,
),
ChunkedValue::None => {
((merge_op.fnc)(&merge_op.params, &result, None)?, 0)
}
};
match merge_result {
MergeResult::Update(value) => {
if !chunk_value(
&trx,
&mut key,
&value,
class,
Some(prev_num_chunks),
)
.await
{
trx.cancel();
return Err(trc::StoreEvent::FoundationdbError
.ctx(trc::Key::Reason, "Value is too large"));
}
}
MergeResult::Delete => {
if prev_num_chunks > 1 {
clear_chunks(&trx, &key, None).await;
} else {
trx.clear(&key);
}
}
MergeResult::Skip => (),
}
}
ValueOp::AtomicAdd(by) => {
trx.atomic_op(&key, &by.to_le_bytes()[..], MutationType::Add);
}
ValueOp::AddAndGet(by) => {
let num = if let Some(bytes) =
trx.get(&key, false).await.map_err(into_error)?
{
deserialize_i64_le(&key, &bytes)? + *by
} else {
*by
};
trx.set(&key, &num.to_le_bytes()[..]);
result.push_counter_id(num);
}
ValueOp::Clear => {
if is_chunked_value(key[0], class) {
clear_chunks(&trx, &key, None).await;
} else {
trx.clear(&key);
}
}
}
}
Operation::Index { field, key, set } => {
let key = IndexKey {
account_id,
collection,
document_id,
field: *field,
key: &*key,
}
.serialize(WITH_SUBSPACE);
if *set {
trx.set(&key, &[]);
} else {
trx.clear(&key);
}
}
Operation::Log { collection, set } => {
let key = LogKey {
account_id,
collection: u8::from(*collection),
change_id,
}
.serialize(WITH_SUBSPACE);
trx.set(&key, set);
}
Operation::AssertValue {
class,
assert_value,
} => {
let key =
class.serialize(account_id, collection, document_id, WITH_SUBSPACE);
let matches = match read_chunked_value(&key, &trx, false).await {
Ok(ChunkedValue::Single(bytes)) => assert_value.matches(bytes.as_ref()),
Ok(ChunkedValue::Chunked { bytes, .. }) => {
assert_value.matches(bytes.as_ref())
}
Ok(ChunkedValue::None) => assert_value.is_none(),
Err(_) => false,
};
if !matches {
trx.cancel();
return Err(trc::StoreEvent::AssertValueFailed.into());
}
}
}
}
if self
.commit(
trx,
retry_count < MAX_COMMIT_ATTEMPTS && start.elapsed() < MAX_COMMIT_TIME,
)
.await?
{
return Ok(result);
} else {
let backoff = rand::rng().random_range(50..=100);
tokio::time::sleep(Duration::from_millis(backoff)).await;
retry_count += 1;
}
}
}
pub(crate) async fn commit(&self, trx: Transaction, will_retry: bool) -> trc::Result<bool> {
match trx.commit().await {
Ok(result) => {
let commit_version = result.committed_version().map_err(into_error)?;
self.version.raise_floor(commit_version);
Ok(true)
}
Err(err) => {
if will_retry {
err.on_error().await.map_err(into_error)?;
Ok(false)
} else {
Err(into_error(FdbError::from(err)))
}
}
}
}
pub(crate) async fn purge_store(&self) -> trc::Result<()> {
// Obtain all zero counters
let mut delete_keys = Vec::new();
for subspace in [SUBSPACE_COUNTER, SUBSPACE_QUOTA, SUBSPACE_IN_MEMORY_COUNTER] {
let trx = self.db.create_trx().map_err(into_error)?;
let from_key = [subspace, 0u8];
let to_key = [subspace, u8::MAX, u8::MAX, u8::MAX, u8::MAX, u8::MAX];
let mut values = trx.get_ranges_keyvalues(
RangeOption {
begin: KeySelector::first_greater_or_equal(&from_key[..]),
end: KeySelector::first_greater_or_equal(&to_key[..]),
mode: options::StreamingMode::WantAll,
reverse: false,
..Default::default()
},
true,
);
while let Some(value) = values.try_next().await.map_err(into_error)? {
if value.value().iter().all(|byte| *byte == 0) {
delete_keys.push(value.key().to_vec());
}
}
}
if delete_keys.is_empty() {
return Ok(());
}
// Delete keys
let integer = 0i64.to_le_bytes();
for chunk in delete_keys.chunks(1024) {
let mut retry_count = 0;
loop {
let trx = self.db.create_trx().map_err(into_error)?;
for key in chunk {
trx.atomic_op(key, &integer, MutationType::CompareAndClear);
}
if self.commit(trx, retry_count < MAX_COMMIT_ATTEMPTS).await? {
break;
} else {
retry_count += 1;
}
}
}
Ok(())
}
pub(crate) async fn delete_range(&self, from: impl Key, to: impl Key) -> trc::Result<()> {
let from = from.serialize(WITH_SUBSPACE);
let to = to.serialize(WITH_SUBSPACE);
let trx = self.db.create_trx().map_err(into_error)?;
trx.clear_range(&from, &to);
self.commit(trx, false).await.map(|_| ())
}
}
fn is_chunked_subspace(subspace: u8) -> bool {
matches!(
subspace,
crate::SUBSPACE_PROPERTY
| crate::SUBSPACE_SEARCH_INDEX
| crate::SUBSPACE_QUEUE_MESSAGE
| crate::SUBSPACE_TASK_QUEUE
| crate::SUBSPACE_DIRECTORY
| crate::SUBSPACE_REGISTRY
| crate::SUBSPACE_DELETED_ITEMS
| crate::SUBSPACE_SPAM_SAMPLES
| crate::SUBSPACE_REPORT_IN
| crate::SUBSPACE_REPORT_OUT
| crate::SUBSPACE_TELEMETRY_SPAN
)
}
fn is_chunked_value(subspace: u8, class: &ValueClass) -> bool {
is_chunked_subspace(subspace)
&& match class {
ValueClass::Property(_)
| ValueClass::IndexProperty(IndexPropertyClass::Hash { .. })
| ValueClass::Registry(RegistryClass::Item { .. })
| ValueClass::Queue(QueueClass::Message(_))
| ValueClass::TaskQueue(TaskQueueClass::Task { .. })
| ValueClass::Telemetry(TelemetryClass::Span(_)) => true,
ValueClass::SearchIndex(index) => matches!(index.typ, SearchIndexType::Document),
_ => false,
}
}
async fn clear_chunks(trx: &Transaction, key: &[u8], from_chunk: Option<u8>) {
let to = KeySerializer::new(key.len() + 1)
.write(key)
.write(u8::MAX)
.finalize();
let from = match from_chunk {
Some(from_chunk) => Cow::Owned(
KeySerializer::new(key.len() + 1)
.write(key)
.write(from_chunk)
.finalize(),
),
None => Cow::Borrowed(key),
};
#[cfg(debug_assertions)]
{
let mut chunks = trx.get_ranges_keyvalues(
RangeOption {
begin: KeySelector::first_greater_or_equal(from.as_ref()),
end: KeySelector::first_greater_or_equal(to.as_slice()),
mode: options::StreamingMode::WantAll,
..Default::default()
},
true,
);
while let Ok(Some(chunk)) = chunks.try_next().await {
let found = chunk.key();
debug_assert!(
found.len() == key.len() + 1 || found == key,
"chunk range of {key:?} holds foreign key {found:?}, clearing it would destroy data"
);
}
}
trx.clear_range(from.as_ref(), &to);
}
async fn chunk_value(
trx: &Transaction,
key: &mut Vec<u8>,
value: &[u8],
class: &ValueClass,
prev_num_chunks: Option<usize>,
) -> bool {
let num_chunks = if value.len() > MAX_VALUE_SIZE {
value.len().div_ceil(MAX_VALUE_SIZE)
} else {
1
};
if num_chunks > u8::MAX as usize {
return false;
}
if is_chunked_value(key[0], class)
&& prev_num_chunks.is_none_or(|prev_num_chunks| prev_num_chunks > num_chunks)
{
clear_chunks(trx, key, Some((num_chunks - 1) as u8)).await;
}
if value.len() > MAX_VALUE_SIZE {
for (pos, chunk) in value.chunks(MAX_VALUE_SIZE).enumerate() {
match pos.cmp(&1) {
Ordering::Less => {}
Ordering::Equal => {
key.push(0);
}
Ordering::Greater => {
*key.last_mut().unwrap() += 1;
}
}
trx.set(key, chunk);
}
} else {
trx.set(key, value);
}
true
}
+111
View File
@@ -0,0 +1,111 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::BlobStore;
use registry::schema::structs;
use std::{io::SeekFrom, ops::Range, path::PathBuf, sync::Arc};
use tokio::{
fs::{self, File},
io::{AsyncReadExt, AsyncSeekExt, AsyncWriteExt},
};
use utils::codec::base32_custom::Base32Writer;
pub struct FsStore {
path: PathBuf,
hash_levels: usize,
}
impl FsStore {
pub async fn open(config: structs::FileSystemStore) -> Result<BlobStore, String> {
let path = PathBuf::from(&config.path);
if !path.exists() {
fs::create_dir_all(&path)
.await
.map_err(|e| format!("Failed to create directory: {e}"))?;
}
Ok(BlobStore::Fs(Arc::new(FsStore {
path,
hash_levels: std::cmp::min(config.depth as usize, 5),
})))
}
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let blob_path = self.build_path(key);
let blob_size = match fs::metadata(&blob_path).await {
Ok(m) => m.len() as usize,
Err(_) => return Ok(None),
};
let mut blob = File::open(&blob_path).await.map_err(into_error)?;
Ok(Some(if range.start != 0 || range.end != usize::MAX {
let from_offset = if range.start < blob_size {
range.start
} else {
0
};
let mut buf = vec![0; (std::cmp::min(range.end, blob_size) - from_offset) as usize];
if from_offset > 0 {
blob.seek(SeekFrom::Start(from_offset as u64))
.await
.map_err(into_error)?;
}
blob.read_exact(&mut buf).await.map_err(into_error)?;
buf
} else {
let mut buf = Vec::with_capacity(blob_size as usize);
blob.read_to_end(&mut buf).await.map_err(into_error)?;
buf
}))
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
let blob_path = self.build_path(key);
if fs::metadata(&blob_path)
.await
.map_or(true, |m| m.len() as usize != data.len())
{
fs::create_dir_all(blob_path.parent().unwrap())
.await
.map_err(into_error)?;
let mut blob_file = File::create(&blob_path).await.map_err(into_error)?;
blob_file.write_all(data).await.map_err(into_error)?;
blob_file.flush().await.map_err(into_error)?;
}
Ok(())
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
let blob_path = self.build_path(key);
if fs::metadata(&blob_path).await.is_ok() {
fs::remove_file(&blob_path).await.map_err(into_error)?;
Ok(true)
} else {
Ok(false)
}
}
fn build_path(&self, key: &[u8]) -> PathBuf {
let mut path = self.path.clone();
for byte in key.iter().take(self.hash_levels) {
path.push(format!("{:x}", byte));
}
path.push(Base32Writer::from_bytes(key).finalize());
path
}
}
fn into_error(err: std::io::Error) -> trc::Error {
trc::StoreEvent::FilesystemError.reason(err)
}
+72
View File
@@ -0,0 +1,72 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{HttpStore, HttpStoreConfig, HttpStoreFormat};
use crate::{InMemoryStore, LookupStores, registry::bootstrap::Bootstrap};
use ahash::AHashMap;
use arc_swap::ArcSwap;
use registry::schema::structs::{self, HttpLookupFormat};
use std::{
collections::hash_map::Entry,
sync::atomic::{AtomicBool, AtomicU64},
};
impl LookupStores {
pub async fn parse_http(&mut self, bp: &mut Bootstrap) {
// Parse remote lists
for http in bp.list_infallible::<structs::HttpLookup>().await {
let id = http.id;
let http = http.object;
if !http.enable {
continue;
}
let http_config = HttpStoreConfig {
url: http.url,
retry: http.retry.as_secs(),
refresh: http.refresh.as_secs(),
timeout: http.timeout.into_inner(),
gzipped: http.is_gzipped,
max_size: http.max_size as usize,
max_entries: http.max_entries as usize,
max_entry_size: http.max_entry_size as usize,
format: match http.format {
HttpLookupFormat::List => HttpStoreFormat::List,
HttpLookupFormat::Csv(csv) => HttpStoreFormat::Csv {
index_key: csv.index_key as u32,
index_value: csv.index_value.map(|v| v as u32),
separator: csv.separator.chars().next().unwrap_or(','),
skip_first: csv.skip_first,
},
},
id: http.namespace,
};
match self.stores.entry(http_config.id.as_str().into()) {
Entry::Vacant(entry) => {
let store = HttpStore {
entries: ArcSwap::from_pointee(AHashMap::new()),
expires: AtomicU64::new(0),
in_flight: AtomicBool::new(false),
config: http_config,
client: utils::http::unpooled_http_client(false),
};
entry.insert(InMemoryStore::Http(store.into()));
}
Entry::Occupied(_) => {
bp.build_error(
id,
format!(
"An lookup store with the {} namespace already exists",
http_config.id
),
);
}
}
}
}
}
+232
View File
@@ -0,0 +1,232 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use std::{
io::{BufRead, BufReader},
sync::{Arc, atomic::Ordering},
time::Instant,
};
use ahash::AHashMap;
use compact_str::ToCompactString;
use rand::seq::IndexedRandom;
use utils::HttpLimitResponse;
use crate::{Value, backend::http::HttpStoreFormat, write::now};
use super::HttpStore;
const BROWSER_USER_AGENTS: [&str; 5] = [
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36",
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Edge/120.0.0.0 Safari/537.36",
"Mozilla/5.0 (Macintosh; Intel Mac OS X 14_1) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.1 Safari/605.1.15",
"Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:120.0) Gecko/20100101 Firefox/120.0",
"Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36",
];
pub(crate) trait HttpStoreGet {
fn get(&self, key: &str) -> Option<Value<'static>>;
fn contains(&self, key: &str) -> bool;
fn refresh(&self);
}
impl HttpStoreGet for Arc<HttpStore> {
fn get(&self, key: &str) -> Option<Value<'static>> {
self.refresh();
self.entries.load().get(key).cloned()
}
fn contains(&self, key: &str) -> bool {
#[cfg(feature = "test_mode")]
{
if self.config.url.contains("phishtank.com")
|| self.config.url.contains("openphish.com")
{
return (self.config.url.contains("open") && key.contains("open"))
|| (self.config.url.contains("tank") && key.contains("tank"));
} else if self.config.url.contains("disposable.github.io") {
return key.ends_with("guerrillamail.com") || key.ends_with("disposable.org");
} else if self.config.url.contains("free_email_provider_domains.txt") {
return key.ends_with("gmail.com")
|| key.ends_with("googlemail.com")
|| key.ends_with("yahoomail.com")
|| key.ends_with("outlook.com")
|| key.ends_with("freemail.org");
}
}
self.refresh();
self.entries.load().contains_key(key)
}
fn refresh(&self) {
if self.expires.load(Ordering::Relaxed) <= now() {
let in_flight = self.in_flight.swap(true, Ordering::Relaxed);
if !in_flight {
let this = self.clone();
tokio::spawn(async move {
let expires = match this.try_refresh().await {
Ok(list) => {
this.entries.store(list.into());
this.config.refresh
}
Err(err) => {
trc::error!(err);
this.config.retry
}
};
this.expires.store(now() + expires, Ordering::Relaxed);
this.in_flight.store(false, Ordering::Relaxed);
});
}
}
}
}
impl HttpStore {
async fn try_refresh(&self) -> trc::Result<AHashMap<String, Value<'static>>> {
let time = Instant::now();
let agent = BROWSER_USER_AGENTS.choose(&mut rand::rng()).unwrap();
let response = self
.client
.get(&self.config.url)
.timeout(self.config.timeout)
.header(reqwest::header::USER_AGENT, *agent)
.send()
.await
.map_err(|err| {
trc::StoreEvent::HttpStoreError
.into_err()
.reason(err)
.ctx(trc::Key::Url, self.config.url.to_compact_string())
.details("Failed to build request")
})?;
if !response.status().is_success() {
trc::bail!(
trc::StoreEvent::HttpStoreError
.into_err()
.ctx(trc::Key::Code, response.status().as_u16())
.ctx(trc::Key::Url, self.config.url.to_compact_string())
.ctx(trc::Key::Elapsed, time.elapsed())
.details("Failed to fetch HTTP list")
);
}
let bytes = response
.bytes_with_limit(self.config.max_size)
.await
.map_err(|err| {
trc::StoreEvent::HttpStoreError
.into_err()
.reason(err)
.ctx(trc::Key::Url, self.config.url.to_compact_string())
.ctx(trc::Key::Elapsed, time.elapsed())
.details("Failed to fetch resource")
})?
.ok_or_else(|| {
trc::StoreEvent::HttpStoreError
.into_err()
.ctx(trc::Key::Url, self.config.url.to_compact_string())
.ctx(trc::Key::Elapsed, time.elapsed())
.details("Resource is too large")
})?;
let reader: Box<dyn std::io::Read + Sync + Send> = if self.config.gzipped {
Box::new(flate2::read::GzDecoder::new(&bytes[..]))
} else {
Box::new(&bytes[..])
};
let mut entries = AHashMap::new();
for (pos, line) in BufReader::new(reader).lines().enumerate() {
let line_ = line.map_err(|err| {
trc::StoreEvent::HttpStoreError
.into_err()
.reason(err)
.ctx(trc::Key::Url, self.config.url.to_compact_string())
.ctx(trc::Key::Elapsed, time.elapsed())
.details("Failed to read line")
})?;
match &self.config.format {
HttpStoreFormat::List => {
let line = line_.trim();
if !line.is_empty() {
entries.insert(line.to_string(), Value::Integer(1));
}
}
HttpStoreFormat::Csv {
index_key,
index_value,
separator,
skip_first,
} if pos > 0 || !*skip_first => {
let mut in_quote = false;
let mut col_num = 0;
let mut last_ch = ' ';
let mut entry_key: String = String::new();
let mut entry_value: String = String::new();
for ch in line_.chars() {
match ch {
'"' if last_ch != '\\' => {
in_quote = !in_quote;
}
'\\' if last_ch != '\\' => (),
_ => {
if ch == *separator && !in_quote {
if col_num == *index_key && index_value.is_none() {
break;
} else {
col_num += 1;
}
} else if col_num == *index_key {
entry_key.push(ch);
if entry_key.len() > self.config.max_entry_size {
break;
}
} else if index_value.is_some_and(|v| col_num == v) {
entry_value.push(ch);
if entry_value.len() > self.config.max_entry_size {
break;
}
}
}
}
last_ch = ch;
}
if !entry_key.is_empty() {
let entry_value = if !entry_value.is_empty() {
Value::Text(entry_value.into())
} else {
Value::Integer(1)
};
entries.insert(entry_key, entry_value);
}
}
_ => (),
}
if entries.len() == self.config.max_entries {
break;
}
}
trc::event!(
Store(trc::StoreEvent::HttpStoreFetch),
Url = self.config.url.to_compact_string(),
Total = entries.len(),
Elapsed = time.elapsed(),
);
Ok(entries)
}
}
+52
View File
@@ -0,0 +1,52 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
pub mod config;
pub mod lookup;
use std::{
sync::atomic::{AtomicBool, AtomicU64},
time::Duration,
};
use ahash::AHashMap;
use arc_swap::ArcSwap;
use crate::Value;
#[derive(Debug, Clone)]
pub struct HttpStoreConfig {
pub id: String,
pub url: String,
pub retry: u64,
pub refresh: u64,
pub timeout: Duration,
pub gzipped: bool,
pub max_size: usize,
pub max_entries: usize,
pub max_entry_size: usize,
pub format: HttpStoreFormat,
}
#[derive(Debug, Clone)]
pub enum HttpStoreFormat {
List,
Csv {
index_key: u32,
index_value: Option<u32>,
separator: char,
skip_first: bool,
},
}
#[derive(Debug)]
pub struct HttpStore {
pub entries: ArcSwap<AHashMap<String, Value<'static>>>,
pub expires: AtomicU64,
pub in_flight: AtomicBool,
pub config: HttpStoreConfig,
pub client: reqwest::Client,
}
+332
View File
@@ -0,0 +1,332 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{
SearchStore,
backend::meili::{MeiliSearchStore, Task, TaskStatus, TaskUid},
search::{
CalendarSearchField, ContactSearchField, EmailSearchField, SearchField, SearchableField,
TracingSearchField,
},
write::now,
};
use registry::schema::structs;
use reqwest::{Error, Response, Url};
use serde_json::{Value, json};
use std::{sync::Arc, time::Duration};
const UNCONFIRMED_TASK_RECHECK_DELAY: u64 = 600;
pub(crate) const MAX_TOTAL_HITS: u64 = 100_000;
impl MeiliSearchStore {
pub async fn open(config: structs::MeilisearchStore) -> Result<SearchStore, String> {
let client = config
.http_auth
.build_http_client(
config.http_headers,
"application/json".into(),
config.timeout,
config.allow_invalid_certs,
)
.await?;
Url::parse(&config.url).map_err(|e| format!("Invalid URL: {e}",))?;
let ms = Self {
client,
url: config.url,
task_poll_interval: Duration::from_millis(500),
task_poll_retries: 120,
task_fail_on_timeout: true,
};
if let Err(err) = ms.create_indexes().await {
return Err(format!("Failed to create indexes: {err}"));
}
Ok(SearchStore::MeiliSearch(Arc::new(MeiliSearchStore {
client: ms.client,
url: ms.url,
task_poll_interval: config.poll_interval.into_inner(),
task_poll_retries: config.max_retries as usize,
task_fail_on_timeout: config.fail_on_timeout,
})))
}
pub async fn create_indexes(&self) -> trc::Result<()> {
self.create_index::<EmailSearchField>().await?;
self.create_index::<CalendarSearchField>().await?;
self.create_index::<ContactSearchField>().await?;
self.create_index::<TracingSearchField>().await?;
Ok(())
}
async fn index_exists(&self, index_uid: &str) -> trc::Result<bool> {
let response = self
.client
.get(format!("{}/indexes/{}", self.url, index_uid))
.send()
.await
.map_err(|err| trc::StoreEvent::MeilisearchError.reason(err))?;
match response.status().as_u16() {
200..=299 => Ok(true),
404 => Ok(false),
status => {
let text = response.text().await.unwrap_or_default();
Err(trc::StoreEvent::MeilisearchError
.reason(text)
.ctx(trc::Key::Code, status))
}
}
}
async fn create_index<T: SearchableField>(&self) -> trc::Result<()> {
let index_name = T::index().index_name();
if self.index_exists(index_name).await? {
return Ok(());
}
let response = assert_success(
self.client
.post(format!("{}/indexes", self.url))
.body(
json!({
"uid": index_name,
"primaryKey": "id",
})
.to_string(),
)
.send()
.await,
)
.await?;
if !self.wait_for_task(response).await? {
// Index already exists
return Ok(());
}
let mut searchable = Vec::new();
let mut filterable = Vec::new();
let mut sortable = Vec::new();
for field in T::all_fields() {
if field.is_indexed() {
sortable.push(Value::String(field.field_name().to_string()));
}
if field.is_text() {
searchable.push(Value::String(field.field_name().to_string()));
} else {
filterable.push(Value::String(field.field_name().to_string()));
}
}
for key in T::primary_keys() {
filterable.push(Value::String(key.field_name().to_string()));
if matches!(key, SearchField::Id) {
sortable.push(Value::String(key.field_name().to_string()));
}
}
#[cfg(feature = "test_mode")]
filterable.push(Value::String("bcc".into()));
if !searchable.is_empty() {
self.update_index_settings(
index_name,
"searchable-attributes",
Value::Array(searchable),
)
.await?;
}
if !filterable.is_empty() {
self.update_index_settings(
index_name,
"filterable-attributes",
Value::Array(filterable),
)
.await?;
}
if !sortable.is_empty() {
self.update_index_settings(index_name, "sortable-attributes", Value::Array(sortable))
.await?;
}
self.update_index_pagination(index_name).await?;
Ok(())
}
async fn update_index_pagination(&self, index_uid: &str) -> trc::Result<bool> {
let response = assert_success(
self.client
.patch(format!(
"{}/indexes/{}/settings/pagination",
self.url, index_uid
))
.body(json!({ "maxTotalHits": MAX_TOTAL_HITS }).to_string())
.send()
.await,
)
.await?;
self.wait_for_task(response).await
}
async fn update_index_settings(
&self,
index_uid: &str,
setting: &str,
value: Value,
) -> trc::Result<bool> {
let response = assert_success(
self.client
.put(format!(
"{}/indexes/{}/settings/{}",
self.url, index_uid, setting
))
.body(value.to_string())
.send()
.await,
)
.await?;
self.wait_for_task(response).await
}
#[cfg(feature = "test_mode")]
pub async fn drop_indexes(&self) -> trc::Result<()> {
use crate::write::SearchIndex;
for index in &[
SearchIndex::Email,
SearchIndex::Calendar,
SearchIndex::Contacts,
SearchIndex::Tracing,
] {
let response = self
.client
.delete(format!("{}/indexes/{}", self.url, index.index_name()))
.send()
.await
.map_err(|err| trc::StoreEvent::MeilisearchError.reason(err))?;
match response.status().as_u16() {
200..=299 => {
self.wait_for_task(response).await?;
}
400..=499 => {
// Index does not exist
return Ok(());
}
_ => {
let status = response.status();
let msg = response.text().await.unwrap_or_default();
return Err(trc::StoreEvent::MeilisearchError
.reason(msg)
.ctx(trc::Key::Code, status.as_u16()));
}
}
}
Ok(())
}
pub(crate) async fn wait_for_task(&self, response: Response) -> trc::Result<bool> {
let response_body = response.text().await.map_err(|err| {
trc::StoreEvent::MeilisearchError
.reason(err)
.details("Request failed")
})?;
let task_uid = serde_json::from_str::<TaskUid>(&response_body)
.map_err(|err| trc::StoreEvent::MeilisearchError.reason(err))?
.task_uid;
let mut loop_count = 0;
let url = format!("{}/tasks/{}", self.url, task_uid);
while loop_count < self.task_poll_retries {
let resp = assert_success(self.client.get(&url).send().await).await?;
let text = resp
.text()
.await
.map_err(|err| trc::StoreEvent::MeilisearchError.reason(err))?;
let task = serde_json::from_str::<Task>(&text).map_err(|err| {
trc::StoreEvent::MeilisearchError
.reason(err)
.details(text.clone())
})?;
match task.status {
TaskStatus::Succeeded => return Ok(true),
TaskStatus::Failed => {
let (code, message) = task
.error
.map(|e| (e.code, Some(e.message)))
.unwrap_or((None, None));
return if matches!(code.as_deref(), Some("index_already_exists")) {
Ok(false)
} else {
Err(trc::StoreEvent::MeilisearchError
.reason("Meilisearch task failed.")
.id(task_uid)
.code(code)
.details(message))
};
}
TaskStatus::Canceled => {
return Err(trc::StoreEvent::MeilisearchError
.reason("Meilisearch task was canceled")
.id(task_uid));
}
TaskStatus::Enqueued | TaskStatus::Processing => {
loop_count += 1;
tokio::time::sleep(self.task_poll_interval).await;
}
TaskStatus::Unknown => {
return Err(trc::StoreEvent::MeilisearchError
.reason("Meilisearch task returned an unknown status")
.id(task_uid)
.details(text));
}
}
}
let err = trc::StoreEvent::MeilisearchError
.reason("Timed out waiting for Meilisearch task")
.id(task_uid);
Err(if self.task_fail_on_timeout {
err
} else {
err.ctx(
trc::Key::NextRetry,
now().saturating_add(UNCONFIRMED_TASK_RECHECK_DELAY),
)
})
}
}
pub(crate) async fn assert_success(response: Result<Response, Error>) -> trc::Result<Response> {
match response {
Ok(response) => {
let status = response.status();
if status.is_success() {
Ok(response)
} else {
Err(trc::StoreEvent::MeilisearchError
.reason(response.text().await.unwrap_or_default())
.ctx(trc::Key::Code, status.as_u16()))
}
}
Err(err) => Err(trc::StoreEvent::MeilisearchError.reason(err)),
}
}
+69
View File
@@ -0,0 +1,69 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use reqwest::Client;
use serde::Deserialize;
use std::time::Duration;
pub mod main;
pub mod search;
pub struct MeiliSearchStore {
client: Client,
url: String,
task_poll_interval: Duration,
task_poll_retries: usize,
task_fail_on_timeout: bool,
}
#[derive(Debug, Deserialize)]
pub(crate) struct TaskUid {
#[serde(rename = "taskUid")]
pub task_uid: u64,
}
#[derive(Debug, Deserialize)]
struct TaskError {
message: String,
#[serde(default)]
code: Option<String>,
}
#[derive(Debug, Deserialize)]
struct Task {
//#[serde(rename = "uid")]
//uid: u64,
status: TaskStatus,
#[serde(default)]
error: Option<TaskError>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "lowercase")]
enum TaskStatus {
Enqueued,
Processing,
Succeeded,
Failed,
Canceled,
#[serde(other)]
Unknown,
}
#[derive(Debug, Deserialize)]
struct MeiliSearchResponse {
hits: Vec<MeiliHit>,
}
#[derive(Debug, Deserialize)]
struct MeiliDocumentsResponse {
results: Vec<MeiliHit>,
}
#[derive(Debug, Deserialize)]
struct MeiliHit {
id: u64,
}
+561
View File
@@ -0,0 +1,561 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{
backend::meili::{
MeiliDocumentsResponse, MeiliSearchResponse, MeiliSearchStore,
main::{MAX_TOTAL_HITS, assert_success},
},
search::*,
write::SearchIndex,
};
use ahash::AHashSet;
use serde_json::{Map, Value, json};
use std::fmt::{Display, Write};
const MAX_SEARCH_RESULTS: usize = 10_000;
impl MeiliSearchStore {
pub async fn index(&self, documents: Vec<IndexDocument>) -> trc::Result<()> {
let mut index_documents: [String; 5] = [
String::new(),
String::new(),
String::new(),
String::new(),
String::new(),
];
for document in documents {
let request = &mut index_documents[document.index.array_pos()];
if !request.is_empty() {
request.push(',');
} else {
request.reserve(1024);
request.push('[');
}
json_serialize(request, &document);
}
for (mut payload, index) in index_documents.into_iter().zip([
SearchIndex::Email,
SearchIndex::Calendar,
SearchIndex::Contacts,
SearchIndex::Tracing,
SearchIndex::File,
]) {
if payload.is_empty() {
continue;
}
payload.push(']');
let response = assert_success(
self.client
.put(format!(
"{}/indexes/{}/documents",
self.url,
index.index_name()
))
.body(payload)
.send()
.await,
)
.await?;
self.wait_for_task(response).await?;
}
Ok(())
}
pub async fn query<R: SearchDocumentId>(
&self,
index: SearchIndex,
filters: &[SearchFilter],
sort: &[SearchComparator],
) -> trc::Result<Vec<R>> {
let filter_group = build_query(filters);
if filter_group.q.is_empty() && sort.is_empty() {
return self.fetch_documents(index, &filter_group.filter).await;
}
let mut body = Map::new();
body.insert("limit".to_string(), Value::from(MAX_SEARCH_RESULTS));
body.insert(
"attributesToRetrieve".to_string(),
Value::Array(vec![Value::String("id".to_string())]),
);
if !filter_group.filter.is_empty() {
body.insert("filter".to_string(), Value::String(filter_group.filter));
}
if !filter_group.q.is_empty() {
body.insert("q".to_string(), Value::String(filter_group.q));
body.insert(
"matchingStrategy".to_string(),
Value::String("all".to_string()),
);
if !filter_group.search_on.is_empty() {
body.insert(
"attributesToSearchOn".to_string(),
Value::Array(
filter_group
.search_on
.into_iter()
.map(|field| Value::String(field.to_string()))
.collect(),
),
);
}
}
if !sort.is_empty() {
let sort_arr: Vec<Value> = sort
.iter()
.filter_map(|comp| match comp {
SearchComparator::Field { field, ascending } => Some(Value::String(format!(
"{}:{}",
field.field_name(),
if *ascending { "asc" } else { "desc" }
))),
_ => None,
})
.collect();
if !sort_arr.is_empty() {
body.insert("sort".to_string(), Value::Array(sort_arr));
}
}
let url = format!("{}/indexes/{}/search", self.url, index.index_name());
let mut results = Vec::new();
let mut offset = 0;
loop {
body.insert("offset".to_string(), Value::from(offset));
let resp = assert_success(
self.client
.post(&url)
.body(Value::Object(body.clone()).to_string())
.send()
.await,
)
.await?;
let text = resp
.text()
.await
.map_err(|err| trc::StoreEvent::MeilisearchError.reason(err))?;
let hits = serde_json::from_str::<MeiliSearchResponse>(&text)
.map_err(|err| {
trc::StoreEvent::MeilisearchError
.reason(err)
.details(text.clone())
})?
.hits;
let total = hits.len();
results.extend(hits.into_iter().map(|hit| R::from_u64(hit.id)));
if total < MAX_SEARCH_RESULTS {
break;
}
offset += total;
if offset >= MAX_TOTAL_HITS as usize {
trc::event!(
Store(trc::StoreEvent::MeilisearchError),
Reason = "Search results were truncated",
Collection = index.index_name(),
Total = offset,
);
break;
}
}
Ok(results)
}
async fn fetch_documents<R: SearchDocumentId>(
&self,
index: SearchIndex,
filter: &str,
) -> trc::Result<Vec<R>> {
let url = format!(
"{}/indexes/{}/documents/fetch",
self.url,
index.index_name()
);
let mut results = Vec::new();
let mut offset = 0;
loop {
let mut body = Map::new();
body.insert("limit".to_string(), Value::from(MAX_SEARCH_RESULTS));
body.insert("offset".to_string(), Value::from(offset));
body.insert(
"fields".to_string(),
Value::Array(vec![Value::String("id".to_string())]),
);
if !filter.is_empty() {
body.insert("filter".to_string(), Value::String(filter.to_string()));
}
let resp = assert_success(
self.client
.post(&url)
.body(Value::Object(body).to_string())
.send()
.await,
)
.await?;
let text = resp
.text()
.await
.map_err(|err| trc::StoreEvent::MeilisearchError.reason(err))?;
let documents = serde_json::from_str::<MeiliDocumentsResponse>(&text)
.map_err(|err| {
trc::StoreEvent::MeilisearchError
.reason(err)
.details(text.clone())
})?
.results;
let total = documents.len();
results.extend(documents.into_iter().map(|hit| R::from_u64(hit.id)));
if total < MAX_SEARCH_RESULTS {
break;
}
offset += total;
}
Ok(results)
}
pub async fn unindex(&self, filter: SearchQuery) -> trc::Result<u64> {
let filter_group = build_query(&filter.filters);
if filter_group.filter.is_empty() {
return Err(trc::StoreEvent::MeilisearchError.reason(
"Meilisearch delete-by-filter requires structured (non-text) filters only",
));
}
let url = format!(
"{}/indexes/{}/documents/delete",
self.url,
filter.index.index_name()
);
let response = assert_success(
self.client
.post(url)
.body(json!({ "filter": filter_group.filter }).to_string())
.send()
.await,
)
.await?;
self.wait_for_task(response).await?;
Ok(0)
}
}
#[derive(Default, Debug)]
struct FilterGroup {
q: String,
filter: String,
search_on: AHashSet<&'static str>,
}
fn build_query(filters: &[SearchFilter]) -> FilterGroup {
if filters.is_empty() {
return FilterGroup::default();
}
let mut operator_stack = Vec::new();
let mut operator = &SearchFilter::And;
let mut is_first = true;
let mut filter = String::new();
let mut queries = AHashSet::new();
let mut search_on = AHashSet::new();
for f in filters {
match f {
SearchFilter::Operator { field, op, value } => {
if field.is_text() && matches!(op, SearchOperator::Equal | SearchOperator::Contains)
{
let value = match value {
SearchValue::Text { value, .. } => value,
_ => {
debug_assert!(
false,
"Text field search with non-text value is not supported"
);
""
}
};
search_on.insert(field.field_name());
if matches!(op, SearchOperator::Equal) {
queries.insert(format!("{value:?}"));
} else {
for token in value.split_whitespace() {
queries.insert(token.to_string());
}
}
} else {
if !filter.is_empty() && !filter.ends_with('(') {
match operator {
SearchFilter::And => filter.push_str(" AND "),
SearchFilter::Or => filter.push_str(" OR "),
_ => (),
}
}
match value {
SearchValue::Text { value, .. } => {
filter.push_str(field.field_name());
filter.push(' ');
op.write_meli_op(&mut filter, format!("{value:?}"));
}
SearchValue::KeyValues(kv) => {
let (key, value) = kv.iter().next().unwrap();
filter.push_str(field.field_name());
filter.push('.');
filter.push_str(key);
filter.push(' ');
op.write_meli_op(&mut filter, format!("{value:?}"));
}
SearchValue::Int(v) => {
filter.push_str(field.field_name());
filter.push(' ');
op.write_meli_op(&mut filter, v);
}
SearchValue::Uint(v) => {
filter.push_str(field.field_name());
filter.push(' ');
op.write_meli_op(&mut filter, v);
}
SearchValue::Boolean(v) => {
filter.push_str(field.field_name());
filter.push(' ');
op.write_meli_op(&mut filter, v);
}
}
}
}
SearchFilter::And | SearchFilter::Or => {
if !filter.is_empty() && !filter.ends_with('(') {
match operator {
SearchFilter::And => filter.push_str(" AND "),
SearchFilter::Or => filter.push_str(" OR "),
_ => (),
}
}
operator_stack.push((operator, is_first));
operator = f;
is_first = true;
filter.push('(');
}
SearchFilter::Not => {
if !filter.is_empty() && !filter.ends_with('(') {
match operator {
SearchFilter::And => filter.push_str(" AND "),
SearchFilter::Or => filter.push_str(" OR "),
_ => (),
}
}
operator_stack.push((operator, is_first));
operator = &SearchFilter::And;
is_first = true;
filter.push_str("NOT (");
}
SearchFilter::End => {
let p = operator_stack.pop().unwrap_or((&SearchFilter::And, true));
operator = p.0;
is_first = p.1;
if !filter.ends_with('(') {
filter.push(')');
} else {
filter.pop();
if filter.ends_with("NOT ") {
let len = filter.len();
filter.truncate(len - 4);
}
if filter.ends_with(" AND ") {
let len = filter.len();
filter.truncate(len - 5);
is_first = true;
} else if filter.ends_with(" OR ") {
let len = filter.len();
filter.truncate(len - 4);
is_first = true;
}
}
}
SearchFilter::DocumentSet(_) => {
debug_assert!(false, "DocumentSet filters are not supported")
}
}
}
let mut q = String::new();
if !queries.is_empty() {
for (idx, term) in queries.into_iter().enumerate() {
if idx > 0 {
q.push(' ');
}
q.push_str(&term);
}
}
FilterGroup {
q,
filter,
search_on,
}
}
impl SearchOperator {
fn write_meli_op(&self, query: &mut String, value: impl Display) {
match self {
SearchOperator::LowerThan => {
let _ = write!(query, "< {value}");
}
SearchOperator::LowerEqualThan => {
let _ = write!(query, "<= {value}");
}
SearchOperator::GreaterThan => {
let _ = write!(query, "> {value}");
}
SearchOperator::GreaterEqualThan => {
let _ = write!(query, ">= {value}");
}
SearchOperator::Equal | SearchOperator::Contains => {
let _ = write!(query, "= {value}");
}
}
}
}
fn json_serialize(request: &mut String, document: &IndexDocument) {
let mut id = 0u64;
let mut is_first = true;
request.push('{');
for (k, v) in document.fields.iter() {
match k {
SearchField::AccountId => {
if let SearchValue::Uint(account_id) = v {
id |= account_id << 32;
}
}
SearchField::DocumentId => {
if let SearchValue::Uint(doc_id) = v {
id |= doc_id;
}
}
SearchField::Id => {
if let SearchValue::Uint(doc_id) = v {
id = *doc_id;
}
continue;
}
_ => {}
}
if !is_first {
request.push(',');
} else {
is_first = false;
}
let _ = write!(request, "{:?}:", k.field_name());
match v {
SearchValue::Text { value, .. } => {
json_serialize_str(request, value);
}
SearchValue::KeyValues(map) => {
request.push('{');
for (i, (key, value)) in map.iter().enumerate() {
if i > 0 {
request.push(',');
}
json_serialize_str(request, key);
request.push(':');
json_serialize_str(request, value);
}
request.push('}');
}
SearchValue::Int(v) => {
let _ = write!(request, "{}", v);
}
SearchValue::Uint(v) => {
let _ = write!(request, "{}", v);
}
SearchValue::Boolean(v) => {
let _ = write!(request, "{}", v);
}
}
}
/*if id == 0 {
debug_assert!(false, "Document is missing required ID fields");
}*/
let _ = write!(request, ",\"id\":{id}}}");
}
fn json_serialize_str(request: &mut String, value: &str) {
request.push('"');
for c in value.chars() {
match c {
'"' => request.push_str("\\\""),
'\\' => request.push_str("\\\\"),
'\n' => request.push_str("\\n"),
'\r' => request.push_str("\\r"),
'\t' => request.push_str("\\t"),
'\u{0008}' => request.push_str("\\b"), // backspace
'\u{000C}' => request.push_str("\\f"), // form feed
_ => {
if !c.is_control() {
request.push(c);
} else {
let _ = write!(request, "\\u{:04x}", c as u32);
}
}
}
}
request.push('"');
}
impl SearchIndex {
#[inline(always)]
fn array_pos(&self) -> usize {
match self {
SearchIndex::Email => 0,
SearchIndex::Calendar => 1,
SearchIndex::Contacts => 2,
SearchIndex::Tracing => 3,
SearchIndex::File => 4,
SearchIndex::InMemory => unreachable!(),
}
}
}
+65
View File
@@ -0,0 +1,65 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{InMemoryStore, LookupStores, Value, registry::bootstrap::Bootstrap};
use ahash::AHashMap;
use registry::schema::structs;
use utils::glob::{GlobMap, GlobSet};
#[derive(Debug)]
pub enum StaticMemoryStore {
Map(GlobMap<Value<'static>>),
Set(GlobSet),
}
impl LookupStores {
pub async fn parse_static(&mut self, bp: &mut Bootstrap) {
let mut lookups = AHashMap::new();
for lookup in bp.list_infallible::<structs::MemoryLookupKeyValue>().await {
if let StaticMemoryStore::Map(map) = lookups
.entry(lookup.object.namespace)
.or_insert_with(|| StaticMemoryStore::Map(Default::default()))
{
if lookup.object.is_glob_pattern {
map.insert_pattern(&lookup.object.key, Value::from(lookup.object.value));
} else {
map.insert_entry(lookup.object.key, Value::from(lookup.object.value));
}
} else {
bp.build_warning(
lookup.id,
"Memory lookup has mixed types (key-value and set)",
);
}
}
for lookup in bp.list_infallible::<structs::MemoryLookupKey>().await {
if let StaticMemoryStore::Set(set) = lookups
.entry(lookup.object.namespace)
.or_insert_with(|| StaticMemoryStore::Set(Default::default()))
{
if lookup.object.is_glob_pattern {
set.insert_pattern(&lookup.object.key);
} else {
set.insert_entry(lookup.object.key);
}
} else {
bp.build_warning(
lookup.id,
"Memory lookup has mixed types (key-value and set)",
);
}
}
for (namespace, store) in lookups {
self.stores.insert(
namespace.into_boxed_str(),
InMemoryStore::Static(store.into()),
);
}
}
}
+39
View File
@@ -0,0 +1,39 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
#[cfg(feature = "azure")]
pub mod azure;
pub mod elastic;
pub mod ephemeral;
#[cfg(feature = "foundation")]
pub mod foundationdb;
pub mod fs;
pub mod http;
pub mod meili;
pub mod memory;
#[cfg(feature = "mysql")]
pub mod mysql;
#[cfg(feature = "postgres")]
pub mod postgres;
#[cfg(feature = "redis")]
pub mod redis;
#[cfg(feature = "rocks")]
pub mod rocksdb;
#[cfg(feature = "s3")]
pub mod s3;
#[cfg(feature = "sqlite")]
pub mod sqlite;
pub const MAX_TOKEN_LENGTH: usize = (u8::MAX >> 1) as usize;
pub const MAX_TOKEN_MASK: usize = MAX_TOKEN_LENGTH - 1;
#[allow(dead_code)]
fn deserialize_i64_le(key: &[u8], bytes: &[u8]) -> trc::Result<i64> {
Ok(i64::from_le_bytes(bytes[..].try_into().map_err(|_| {
trc::Error::corrupted_key(key, bytes.into(), trc::location!())
})?))
}
+64
View File
@@ -0,0 +1,64 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use std::ops::Range;
use mysql_async::prelude::Queryable;
use super::{MysqlStore, into_error};
impl MysqlStore {
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn
.prep("SELECT v FROM t WHERE k = ?")
.await
.map_err(into_error)?;
conn.exec_first::<Vec<u8>, _, _>(&s, (key,))
.await
.map(|bytes| {
if range.start == 0 && range.end == usize::MAX {
bytes
} else {
bytes.map(|bytes| {
bytes
.get(range.start..std::cmp::min(bytes.len(), range.end))
.unwrap_or_default()
.to_vec()
})
}
})
.map_err(into_error)
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn
.prep("INSERT INTO t (k, v) VALUES (?, ?) ON DUPLICATE KEY UPDATE v = VALUES(v)")
.await
.map_err(into_error)?;
conn.exec_drop(&s, (key, data))
.await
.map_err(into_error)
.map(|_| ())
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn
.prep("DELETE FROM t WHERE k = ?")
.await
.map_err(into_error)?;
conn.exec_iter(&s, (key,))
.await
.map_err(into_error)
.map(|hits| hits.affected_rows() > 0)
}
}
+136
View File
@@ -0,0 +1,136 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use mysql_async::{Params, Row, prelude::Queryable};
use crate::{IntoRows, QueryResult, QueryType, Value};
use super::{MysqlStore, into_error};
impl MysqlStore {
pub(crate) async fn sql_query<T: QueryResult>(
&self,
query: &str,
params: &[Value<'_>],
) -> trc::Result<T> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn.prep(query).await.map_err(into_error)?;
let params = Params::Positional(params.iter().map(Into::into).collect());
match T::query_type() {
QueryType::Execute => conn.exec_drop(s, params).await.map_or_else(
|e| Err(into_error(e)),
|_| Ok(T::from_exec(conn.affected_rows() as usize)),
),
QueryType::Exists => conn
.exec_first::<Row, _, _>(s, params)
.await
.map_or_else(|e| Err(into_error(e)), |r| Ok(T::from_exists(r.is_some()))),
QueryType::QueryOne => conn
.exec_first::<Row, _, _>(s, params)
.await
.map_or_else(|e| Err(into_error(e)), |r| Ok(T::from_query_one(r))),
QueryType::QueryAll => conn
.exec::<Row, _, _>(s, params)
.await
.map_or_else(|e| Err(into_error(e)), |r| Ok(T::from_query_all(r))),
}
}
}
impl From<crate::Value<'_>> for mysql_async::Value {
fn from(value: crate::Value) -> Self {
match value {
crate::Value::Integer(i) => mysql_async::Value::Int(i),
crate::Value::Bool(b) => mysql_async::Value::Int(b as i64),
crate::Value::Float(f) => mysql_async::Value::Double(f),
crate::Value::Text(t) => mysql_async::Value::Bytes(t.into_owned().into_bytes()),
crate::Value::Blob(b) => mysql_async::Value::Bytes(b.into_owned()),
crate::Value::Null => mysql_async::Value::NULL,
}
}
}
impl From<mysql_async::Value> for crate::Value<'static> {
fn from(value: mysql_async::Value) -> Self {
match value {
mysql_async::Value::Int(i) => Self::Integer(i),
mysql_async::Value::UInt(i) => Self::Integer(i as i64),
mysql_async::Value::Double(f) => Self::Float(f),
mysql_async::Value::Bytes(b) => String::from_utf8(b).map_or_else(
|e| Self::Blob(e.into_bytes().into()),
|s| Self::Text(s.into()),
),
mysql_async::Value::NULL => Self::Null,
mysql_async::Value::Float(f) => Self::Float(f as f64),
mysql_async::Value::Date(_, _, _, _, _, _, _)
| mysql_async::Value::Time(_, _, _, _, _, _) => Self::Text(value.as_sql(true).into()),
}
}
}
impl IntoRows for Vec<mysql_async::Row> {
fn into_rows(self) -> crate::Rows {
crate::Rows {
rows: self
.into_iter()
.map(|r| crate::Row {
values: r
.unwrap_raw()
.into_iter()
.flatten()
.map(Into::into)
.collect(),
})
.collect(),
}
}
fn into_named_rows(self) -> crate::NamedRows {
crate::NamedRows {
names: self
.first()
.map(|r| r.columns().iter().map(|c| c.name_str().into()).collect())
.unwrap_or_default(),
rows: self
.into_iter()
.map(|r| crate::Row {
values: r
.unwrap_raw()
.into_iter()
.flatten()
.map(Into::into)
.collect(),
})
.collect(),
}
}
fn into_row(self) -> Option<crate::Row> {
unreachable!()
}
}
impl IntoRows for Option<mysql_async::Row> {
fn into_row(self) -> Option<crate::Row> {
self.map(|row| crate::Row {
values: row
.unwrap_raw()
.into_iter()
.flatten()
.map(Into::into)
.collect(),
})
}
fn into_rows(self) -> crate::Rows {
unreachable!()
}
fn into_named_rows(self) -> crate::NamedRows {
unreachable!()
}
}
+212
View File
@@ -0,0 +1,212 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{MysqlStore, into_error};
use crate::{
backend::mysql::MysqlSearchField,
search::{
CalendarSearchField, ContactSearchField, EmailSearchField, SearchableField,
TracingSearchField,
},
*,
};
use ::registry::schema::structs;
use mysql_async::{
Conn, OptsBuilder, Pool, PoolConstraints, PoolOpts, SslOpts, prelude::Queryable,
};
impl MysqlStore {
pub async fn open(config: structs::MySqlStore) -> Result<Store, String> {
let mut opts = OptsBuilder::default()
.ip_or_hostname(config.host)
.user(config.auth_username)
.pass(config.auth_secret.secret().await?.map(|v| v.into_owned()))
.db_name(Some(config.database))
.max_allowed_packet(config.max_allowed_packet.map(|v| v as usize))
.wait_timeout(config.timeout.map(|t| t.as_secs() as usize))
.client_found_rows(true)
.tcp_port(config.port as u16);
if config.use_tls {
opts = opts.ssl_opts(Some(
SslOpts::default()
.with_danger_accept_invalid_certs(config.allow_invalid_certs)
.with_danger_skip_domain_validation(config.allow_invalid_certs),
));
}
// Configure connection pool
let mut pool_min = PoolConstraints::default().min();
let mut pool_max = PoolConstraints::default().max();
if let Some(n_size) = config.pool_min_connections {
pool_min = n_size as usize;
}
if let Some(n_size) = config.pool_max_connections {
pool_max = n_size as usize;
}
opts = opts.pool_opts(
PoolOpts::default().with_constraints(PoolConstraints::new(pool_min, pool_max).unwrap()),
);
let mut replicas = vec![];
for replica in config.read_replicas {
replicas.push(Store::MySQL(Arc::new(MysqlStore {
conn_pool: Pool::new(
opts.clone()
.ip_or_hostname(replica.host)
.user(replica.auth_username)
.pass(replica.auth_secret.secret().await?.map(|v| v.into_owned()))
.db_name(Some(replica.database))
.tcp_port(replica.port as u16),
),
})))
}
let primary = Store::MySQL(Arc::new(MysqlStore {
conn_pool: Pool::new(opts),
}));
Ok(primary)
}
pub(crate) async fn create_storage_tables(&self) -> trc::Result<()> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
for table in [
SUBSPACE_ACL,
SUBSPACE_TASK_QUEUE,
SUBSPACE_DELETED_ITEMS,
SUBSPACE_SPAM_SAMPLES,
SUBSPACE_BLOB_LINK,
SUBSPACE_IN_MEMORY_VALUE,
SUBSPACE_PROPERTY,
SUBSPACE_REGISTRY,
SUBSPACE_REGISTRY_PK,
SUBSPACE_DIRECTORY,
SUBSPACE_QUEUE_MESSAGE,
SUBSPACE_QUEUE_EVENT,
SUBSPACE_REPORT_OUT,
SUBSPACE_REPORT_IN,
SUBSPACE_LOGS,
SUBSPACE_TELEMETRY_SPAN,
SUBSPACE_TELEMETRY_METRIC,
] {
let table = char::from(table);
conn.query_drop(format!(
"CREATE TABLE IF NOT EXISTS {table} (
k VARBINARY(255) NOT NULL,
v MEDIUMBLOB NOT NULL,
PRIMARY KEY (k)
) ENGINE=InnoDB"
))
.await
.map_err(into_error)?;
}
conn.query_drop(format!(
"CREATE TABLE IF NOT EXISTS {} (
k VARBINARY(255) NOT NULL,
v LONGBLOB NOT NULL,
PRIMARY KEY (k)
) ENGINE=InnoDB",
char::from(SUBSPACE_BLOBS),
))
.await
.map_err(into_error)?;
for table in [SUBSPACE_INDEXES, SUBSPACE_REGISTRY_IDX] {
let table = char::from(table);
conn.query_drop(format!(
"CREATE TABLE IF NOT EXISTS {table} (
k BLOB,
PRIMARY KEY (k(400))
) ENGINE=InnoDB"
))
.await
.map_err(into_error)?;
}
for table in [SUBSPACE_COUNTER, SUBSPACE_QUOTA, SUBSPACE_IN_MEMORY_COUNTER] {
conn.query_drop(format!(
"CREATE TABLE IF NOT EXISTS {} (
k VARBINARY(255) NOT NULL,
v BIGINT NOT NULL DEFAULT 0,
PRIMARY KEY (k)
) ENGINE=InnoDB",
char::from(table)
))
.await
.map_err(into_error)?;
}
Ok(())
}
pub(crate) async fn create_search_tables(&self) -> trc::Result<()> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
create_search_tables::<EmailSearchField>(&mut conn).await?;
create_search_tables::<CalendarSearchField>(&mut conn).await?;
create_search_tables::<ContactSearchField>(&mut conn).await?;
//create_search_tables::<FileSearchField>(&mut conn).await?;
create_search_tables::<TracingSearchField>(&mut conn).await?;
Ok(())
}
}
async fn create_search_tables<T: SearchableField + MysqlSearchField + 'static>(
conn: &mut Conn,
) -> trc::Result<()> {
let table_name = T::index().mysql_table();
let mut query = format!("CREATE TABLE IF NOT EXISTS {} (", table_name);
// Add primary key columns
let pkeys = T::primary_keys();
for pkey in pkeys {
query.push_str(&format!("{} {}, ", pkey.column(), pkey.column_type()));
}
// Add other columns
for field in T::all_fields() {
query.push_str(&format!("{} {}, ", field.column(), field.column_type()));
}
// Add primary key constraint
query.push_str("PRIMARY KEY (");
for (i, pkey) in pkeys.iter().enumerate() {
if i > 0 {
query.push_str(", ");
}
query.push_str(pkey.column());
}
query.push_str(")) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci");
conn.query_drop(&query).await.map_err(into_error)?;
// Create indexes
for field in T::all_fields() {
if field.is_text() {
let column_name = field.column();
let create_index_query = format!(
"CREATE FULLTEXT INDEX fts_{table_name}_{column_name} ON {table_name}({column_name})",
);
let _ = conn.query_drop(&create_index_query).await;
}
if field.is_indexed() {
let column_name = field.column();
let create_index_query = format!(
"CREATE INDEX idx_{table_name}_{column_name} ON {table_name}({column_name})",
);
let _ = conn.query_drop(&create_index_query).await;
}
}
Ok(())
}
+208
View File
@@ -0,0 +1,208 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{
search::{
CalendarSearchField, ContactSearchField, EmailSearchField, FileSearchField, SearchField,
TracingSearchField,
},
write::SearchIndex,
};
use mysql_async::Pool;
use std::fmt::Display;
pub mod blob;
pub mod lookup;
pub mod main;
pub mod read;
pub mod search;
pub mod write;
pub struct MysqlStore {
pub(crate) conn_pool: Pool,
}
#[inline(always)]
fn into_error(err: impl Display) -> trc::Error {
trc::StoreEvent::MysqlError.reason(err)
}
const ER_LOCK_WAIT_TIMEOUT: u16 = 1205;
const ER_STATEMENT_TIMEOUT: u16 = 1969;
const ER_QUERY_TIMEOUT: u16 = 3024;
pub(crate) const DELETE_CHUNK_SIZE: usize = 1000;
pub(crate) const MIN_DELETE_CHUNK_SIZE: usize = 10;
#[inline(always)]
pub(crate) fn is_timeout_error(err: &mysql_async::Error) -> bool {
matches!(err, mysql_async::Error::Server(err)
if matches!(
err.code,
ER_LOCK_WAIT_TIMEOUT | ER_STATEMENT_TIMEOUT | ER_QUERY_TIMEOUT
)
)
}
impl SearchIndex {
pub fn mysql_table(&self) -> &'static str {
match self {
SearchIndex::Email => "s_email",
SearchIndex::Calendar => "s_cal",
SearchIndex::Contacts => "s_card",
SearchIndex::File => "s_file",
SearchIndex::Tracing => "s_trace",
SearchIndex::InMemory => "",
}
}
}
trait MysqlSearchField {
fn column(&self) -> &'static str;
fn column_type(&self) -> &'static str;
}
impl MysqlSearchField for EmailSearchField {
fn column(&self) -> &'static str {
match self {
EmailSearchField::From => "fadr",
EmailSearchField::To => "tadr",
EmailSearchField::Cc => "cc",
EmailSearchField::Bcc => "bcc",
EmailSearchField::Subject => "subj",
EmailSearchField::Body => "body",
EmailSearchField::Attachment => "atta",
EmailSearchField::ReceivedAt => "rcvd",
EmailSearchField::SentAt => "sent",
EmailSearchField::Size => "size",
EmailSearchField::HasAttachment => "hatt",
EmailSearchField::Headers => "hdrs",
}
}
fn column_type(&self) -> &'static str {
match self {
EmailSearchField::ReceivedAt | EmailSearchField::SentAt => "BIGINT",
EmailSearchField::Size => "INT",
EmailSearchField::HasAttachment => "BOOLEAN",
EmailSearchField::Headers => "JSON",
EmailSearchField::From => "TEXT",
EmailSearchField::To => "TEXT",
EmailSearchField::Cc => "TEXT",
EmailSearchField::Bcc => "TEXT",
EmailSearchField::Subject => "TEXT",
EmailSearchField::Body => "MEDIUMTEXT",
EmailSearchField::Attachment => "MEDIUMTEXT",
}
}
}
impl MysqlSearchField for CalendarSearchField {
fn column(&self) -> &'static str {
match self {
CalendarSearchField::Title => "titl",
CalendarSearchField::Description => "dscd",
CalendarSearchField::Location => "locn",
CalendarSearchField::Owner => "ownr",
CalendarSearchField::Attendee => "atnd",
CalendarSearchField::Start => "strt",
CalendarSearchField::Uid => "uid",
}
}
fn column_type(&self) -> &'static str {
match self {
CalendarSearchField::Start => "BIGINT NOT NULL",
_ => "TEXT",
}
}
}
impl MysqlSearchField for ContactSearchField {
fn column(&self) -> &'static str {
match self {
ContactSearchField::Member => "mmbr",
ContactSearchField::Name => "name",
ContactSearchField::Nickname => "nick",
ContactSearchField::Organization => "orgn",
ContactSearchField::Email => "eml",
ContactSearchField::Phone => "phon",
ContactSearchField::OnlineService => "olsv",
ContactSearchField::Address => "addr",
ContactSearchField::Note => "note",
ContactSearchField::Kind => "kind",
ContactSearchField::Uid => "uid",
}
}
fn column_type(&self) -> &'static str {
match self {
ContactSearchField::Kind | ContactSearchField::Uid => "TEXT",
_ => "TEXT",
}
}
}
impl MysqlSearchField for FileSearchField {
fn column(&self) -> &'static str {
match self {
FileSearchField::Name => "name",
FileSearchField::Content => "body",
}
}
fn column_type(&self) -> &'static str {
match self {
FileSearchField::Name => "TEXT",
FileSearchField::Content => "MEDIUMTEXT",
}
}
}
impl MysqlSearchField for TracingSearchField {
fn column(&self) -> &'static str {
match self {
TracingSearchField::QueueId => "qid",
TracingSearchField::EventType => "etyp",
TracingSearchField::Keywords => "kwds",
}
}
fn column_type(&self) -> &'static str {
match self {
TracingSearchField::EventType => "BIGINT",
TracingSearchField::QueueId => "BIGINT",
TracingSearchField::Keywords => "TEXT",
}
}
}
impl MysqlSearchField for SearchField {
fn column(&self) -> &'static str {
match self {
SearchField::AccountId => "accid",
SearchField::DocumentId => "docid",
SearchField::Id => "id",
SearchField::Email(field) => field.column(),
SearchField::Calendar(field) => field.column(),
SearchField::Contact(field) => field.column(),
SearchField::File(field) => field.column(),
SearchField::Tracing(field) => field.column(),
}
}
fn column_type(&self) -> &'static str {
match self {
SearchField::AccountId => "INT NOT NULL",
SearchField::DocumentId => "INT NOT NULL",
SearchField::Id => "BIGINT NOT NULL",
SearchField::Email(field) => field.column_type(),
SearchField::Calendar(field) => field.column_type(),
SearchField::Contact(field) => field.column_type(),
SearchField::File(field) => field.column_type(),
SearchField::Tracing(field) => field.column_type(),
}
}
}
+169
View File
@@ -0,0 +1,169 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{MysqlStore, into_error, is_timeout_error};
use crate::{Deserialize, IterateParams, Key, ValueKey, write::ValueClass};
use futures::TryStreamExt;
use mysql_async::{Row, prelude::Queryable};
impl MysqlStore {
pub(crate) async fn get_value<U>(&self, key: impl Key) -> trc::Result<Option<U>>
where
U: Deserialize + 'static,
{
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn
.prep(format!(
"SELECT v FROM {} WHERE k = ?",
char::from(key.subspace())
))
.await
.map_err(into_error)?;
let key = key.serialize(0);
conn.exec_first::<Vec<u8>, _, _>(&s, (&key,))
.await
.map_err(into_error)
.and_then(|r| {
if let Some(r) = r {
Ok(Some(U::deserialize_owned_with_key(&key, r)?))
} else {
Ok(None)
}
})
}
pub(crate) async fn key_exists(&self, key: impl Key) -> trc::Result<bool> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn
.prep(format!(
"SELECT 1 FROM {} WHERE k = ?",
char::from(key.subspace())
))
.await
.map_err(into_error)?;
let key = key.serialize(0);
conn.exec_first::<u8, _, _>(&s, (&key,))
.await
.map_err(into_error)
.map(|r| r.is_some())
}
pub(crate) async fn iterate<T: Key>(
&self,
params: IterateParams<T>,
mut cb: impl for<'x> FnMut(&'x [u8], &'x [u8]) -> trc::Result<bool> + Sync + Send,
) -> trc::Result<()> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let table = char::from(params.begin.subspace());
let begin = params.begin.serialize(0);
let end = params.end.serialize(0);
let keys = if params.values { "k, v" } else { "k" };
let s = conn
.prep(&match (params.first, params.ascending) {
(true, true) => {
format!(
"SELECT {keys} FROM {table} WHERE k >= ? AND k <= ? ORDER BY k ASC LIMIT 1"
)
}
(true, false) => {
format!(
"SELECT {keys} FROM {table} WHERE k >= ? AND k <= ? ORDER BY k DESC LIMIT 1"
)
}
(false, true) => {
format!("SELECT {keys} FROM {table} WHERE k >= ? AND k <= ? ORDER BY k ASC")
}
(false, false) => {
format!("SELECT {keys} FROM {table} WHERE k >= ? AND k <= ? ORDER BY k DESC")
}
})
.await
.map_err(into_error)?;
let mut from = begin;
let mut to = end;
let mut resume_key = None;
loop {
let mut last_key = None;
let mut timed_out = false;
{
let mut rows = conn
.exec_stream::<Row, _, _>(&s, (from.clone(), to.clone()))
.await
.map_err(into_error)?;
loop {
match rows.try_next().await {
Ok(Some(mut row)) => {
let value = if params.values {
row.take_opt::<Vec<u8>, _>(1)
.unwrap_or_else(|| Ok(vec![]))
.map_err(into_error)?
} else {
vec![]
};
let key = row
.take_opt::<Vec<u8>, _>(0)
.unwrap_or_else(|| Ok(vec![]))
.map_err(into_error)?;
if resume_key.take().is_some_and(|resumed| resumed == key) {
continue;
}
if !cb(&key, &value)? {
return Ok(());
}
last_key = Some(key);
}
Ok(None) => break,
Err(err) => {
if params.first || last_key.is_none() || !is_timeout_error(&err) {
return Err(into_error(err));
}
timed_out = true;
break;
}
}
}
}
match last_key {
Some(last_key) if timed_out => {
if params.ascending {
from.clone_from(&last_key);
} else {
to.clone_from(&last_key);
}
resume_key = Some(last_key);
}
_ => return Ok(()),
}
}
}
pub(crate) async fn get_counter(
&self,
key: impl Into<ValueKey<ValueClass>> + Sync + Send,
) -> trc::Result<i64> {
let key = key.into();
let table = char::from(key.subspace());
let key = key.serialize(0);
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn
.prep(format!("SELECT v FROM {table} WHERE k = ?"))
.await
.map_err(into_error)?;
match conn.exec_first::<i64, _, _>(&s, (key,)).await {
Ok(Some(num)) => Ok(num),
Ok(None) => Ok(0),
Err(e) => Err(into_error(e)),
}
}
}
+353
View File
@@ -0,0 +1,353 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{
backend::{
MAX_TOKEN_LENGTH,
mysql::{
DELETE_CHUNK_SIZE, MIN_DELETE_CHUNK_SIZE, MysqlSearchField, MysqlStore, into_error,
is_timeout_error,
},
},
search::{
IndexDocument, SearchComparator, SearchDocumentId, SearchFilter, SearchOperator,
SearchQuery, SearchValue,
},
write::SearchIndex,
};
use mysql_async::{IsolationLevel, TxOpts, Value, prelude::Queryable};
use nlp::tokenizers::word::WordTokenizer;
use std::fmt::Write;
impl MysqlStore {
pub async fn index(&self, documents: Vec<IndexDocument>) -> trc::Result<()> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let mut tx_opts = TxOpts::default();
tx_opts
.with_consistent_snapshot(false)
.with_isolation_level(IsolationLevel::ReadCommitted);
let mut trx = conn.start_transaction(tx_opts).await.map_err(into_error)?;
for document in documents {
let index = document.index;
let primary_keys = index.primary_keys();
let all_fields = index.all_fields();
let mut fields = document.fields;
let mut values = Vec::with_capacity(fields.len() + 2);
let mut query = format!("INSERT INTO {} (", index.mysql_table());
for (i, field) in primary_keys.iter().chain(all_fields).enumerate() {
if i > 0 {
query.push(',');
}
query.push_str(field.column());
}
query.push_str(") VALUES (");
for (i, field) in primary_keys.iter().chain(all_fields).enumerate() {
if i > 0 {
query.push(',');
}
if let Some(value) = fields.remove(field) {
query.push('?');
values.push(value);
} else {
query.push_str("NULL");
}
}
query.push_str(") ON DUPLICATE KEY UPDATE ");
for (i, field) in all_fields.iter().enumerate() {
if i > 0 {
query.push(',');
}
let column = field.column();
let _ = write!(&mut query, "{column} = VALUES({column})");
}
let s = trx.prep(&query).await.map_err(into_error)?;
trx.exec_drop(&s, values).await.map_err(into_error)?;
}
trx.commit().await.map_err(into_error)
}
pub async fn query<R: SearchDocumentId>(
&self,
index: SearchIndex,
filters: &[SearchFilter],
sort: &[SearchComparator],
) -> trc::Result<Vec<R>> {
let mut query = format!(
"SELECT {} FROM {}",
R::field().column(),
index.mysql_table()
);
let params = build_filter(&mut query, filters);
if !sort.is_empty() {
build_sort(&mut query, sort);
}
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn.prep(query).await.map_err(into_error)?;
conn.exec::<i64, _, _>(s, params)
.await
.map(|r| r.into_iter().map(|r| R::from_u64(r as u64)).collect())
.map_err(into_error)
}
pub async fn unindex(&self, filter: SearchQuery) -> trc::Result<u64> {
let table = filter.index.mysql_table();
let mut query = format!("DELETE FROM {table} ");
let params = build_filter(&mut query, &filter.filters);
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let s = conn.prep(&query).await.map_err(into_error)?;
match conn.exec_drop(s, params.clone()).await {
Ok(_) => return Ok(conn.affected_rows()),
Err(err) if is_timeout_error(&err) => (),
Err(err) => return Err(into_error(err)),
}
let mut chunk_size = DELETE_CHUNK_SIZE;
let mut deleted = 0;
loop {
let s = conn
.prep(format!("{query} LIMIT {chunk_size}"))
.await
.map_err(into_error)?;
loop {
match conn.exec_drop(&s, params.clone()).await {
Ok(_) => {
let affected = conn.affected_rows();
if affected == 0 {
return Ok(deleted);
}
deleted += affected;
}
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
}
}
}
}
}
fn build_filter(query: &mut String, filters: &[SearchFilter]) -> Vec<Value> {
if filters.is_empty() {
return Vec::new();
}
query.push_str(" WHERE ");
let mut operator_stack = Vec::new();
let mut operator = &SearchFilter::And;
let mut is_first = true;
let mut values: Vec<Value> = Vec::new();
for filter in filters {
match filter {
SearchFilter::Operator { field, op, value } => {
if !is_first {
match operator {
SearchFilter::And => query.push_str(" AND "),
SearchFilter::Or => query.push_str(" OR "),
_ => (),
}
} else {
is_first = false;
}
if field.is_text() && matches!(op, SearchOperator::Equal | SearchOperator::Contains)
{
let (value, mode) = match (value, op) {
(SearchValue::Text { value, .. }, SearchOperator::Equal) => {
(Value::Bytes(format!("{value:?}").into_bytes()), "BOOLEAN")
}
(SearchValue::Text { value, .. }, ..) => {
let mut text_query = String::with_capacity(value.len() + 1);
for item in WordTokenizer::new(value, MAX_TOKEN_LENGTH) {
if !text_query.is_empty() {
text_query.push(' ');
}
text_query.push('+');
text_query.push_str(&item.word);
}
(Value::Bytes(text_query.into_bytes()), "BOOLEAN")
}
_ => {
debug_assert!(false, "Invalid search value for text field");
continue;
}
};
let _ = write!(query, "MATCH({}) AGAINST(? IN {mode} MODE)", field.column());
values.push(value);
} else if let SearchValue::KeyValues(kv) = value {
let (key, value) = kv.iter().next().unwrap();
values.push(Value::Bytes(format!("$.{key:?}").into_bytes()));
if !value.is_empty() {
if op == &SearchOperator::Equal {
let _ = write!(query, "JSON_EXTRACT({}, ?) = ?", field.column());
values.push(Value::Bytes(value.as_bytes().to_vec()));
} else {
let _ = write!(query, "JSON_EXTRACT({}, ?) LIKE ?", field.column(),);
values.push(Value::Bytes(format!("%{value}%").into_bytes()));
}
} else {
let _ = write!(query, "JSON_CONTAINS_PATH({}, 'one', ?)", field.column(),);
}
} else {
query.push_str(field.column());
query.push(' ');
op.write_mysql(query);
values.push(to_mysql(value));
}
}
SearchFilter::And | SearchFilter::Or => {
if !is_first {
match operator {
SearchFilter::And => query.push_str(" AND "),
SearchFilter::Or => query.push_str(" OR "),
_ => (),
}
} else {
is_first = false;
}
operator_stack.push((operator, is_first));
operator = filter;
is_first = true;
query.push('(');
}
SearchFilter::Not => {
if !is_first {
match operator {
SearchFilter::And => query.push_str(" AND "),
SearchFilter::Or => query.push_str(" OR "),
_ => (),
}
} else {
is_first = false;
}
operator_stack.push((operator, is_first));
operator = &SearchFilter::And;
is_first = true;
query.push_str("NOT (");
}
SearchFilter::End => {
let p = operator_stack.pop().unwrap_or((&SearchFilter::And, true));
operator = p.0;
is_first = p.1;
query.push(')');
}
SearchFilter::DocumentSet(_) => {
debug_assert!(
false,
"DocumentSet filters are not supported in Postgres backend"
)
}
}
}
values
}
fn build_sort(query: &mut String, sort: &[SearchComparator]) {
query.push_str(" ORDER BY ");
for (i, comparator) in sort.iter().enumerate() {
if i > 0 {
query.push_str(", ");
}
match comparator {
SearchComparator::Field { field, ascending } => {
query.push_str(field.column());
if *ascending {
query.push_str(" ASC");
} else {
query.push_str(" DESC");
}
}
SearchComparator::DocumentSet { .. } | SearchComparator::SortedSet { .. } => {
debug_assert!(
false,
"DocumentSet and SortedSet comparators are not supported "
);
}
}
}
}
impl SearchOperator {
fn write_mysql(&self, query: &mut String) {
match self {
SearchOperator::LowerThan => {
let _ = write!(query, "< ?");
}
SearchOperator::LowerEqualThan => {
let _ = write!(query, "<= ?");
}
SearchOperator::GreaterThan => {
let _ = write!(query, "> ?");
}
SearchOperator::GreaterEqualThan => {
let _ = write!(query, ">= ?");
}
SearchOperator::Equal => {
let _ = write!(query, "= ?");
}
SearchOperator::Contains => {
let _ = write!(query, "LIKE '%' CONCAT('%', ?, '%')");
}
}
}
}
impl From<SearchValue> for Value {
fn from(value: SearchValue) -> Self {
match value {
SearchValue::Text { mut value, .. } => {
// Truncate values larger than 16MB to avoid MySQL errors
if value.len() > 16_777_214 {
let pos = value.floor_char_boundary(16_777_214);
value.truncate(pos);
}
Value::Bytes(value.into_bytes())
}
SearchValue::KeyValues(vec_map) => serde_json::to_string(&vec_map)
.map(|v| Value::Bytes(v.into_bytes()))
.unwrap_or(Value::NULL),
SearchValue::Int(i) => Value::Int(i),
SearchValue::Uint(i) => Value::Int(i as i64),
SearchValue::Boolean(b) => Value::Int(b as i64),
}
}
}
fn to_mysql(value: &SearchValue) -> Value {
match value {
SearchValue::Text { value, .. } => Value::Bytes(value.as_bytes().to_vec()),
SearchValue::KeyValues(vec_map) => serde_json::to_string(&vec_map)
.map(|v| Value::Bytes(v.into_bytes()))
.unwrap_or(Value::NULL),
SearchValue::Int(i) => Value::Int(*i),
SearchValue::Uint(i) => Value::Int(*i as i64),
SearchValue::Boolean(b) => Value::Int(*b as i64),
}
}
+529
View File
@@ -0,0 +1,529 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{DELETE_CHUNK_SIZE, MIN_DELETE_CHUNK_SIZE, MysqlStore, into_error, is_timeout_error};
use crate::{
IndexKey, Key, LogKey, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER, SUBSPACE_QUOTA,
SUBSPACE_REGISTRY_IDX,
write::{
AssignedIds, Batch, MAX_COMMIT_ATTEMPTS, MAX_COMMIT_TIME, MergeResult, Operation,
ValueClass, ValueOp,
},
};
use ahash::AHashMap;
use mysql_async::{Conn, Error, IsolationLevel, TxOpts, params, prelude::Queryable};
use rand::RngExt;
use std::time::{Duration, Instant};
#[derive(Debug)]
enum CommitError {
Mysql(mysql_async::Error),
Internal(trc::Error),
//Retry,
}
impl MysqlStore {
pub(crate) async fn write(&self, mut batch: Batch<'_>) -> trc::Result<AssignedIds> {
let start = Instant::now();
let mut retry_count = 0;
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
loop {
let err = match self.write_trx(&mut conn, &mut batch).await {
Ok(result) => {
return Ok(result);
}
Err(err) => err,
};
let _ = conn.query_drop("ROLLBACK;").await;
match err {
CommitError::Mysql(Error::Server(err))
if [1062, 1213].contains(&err.code)
&& retry_count < MAX_COMMIT_ATTEMPTS
&& start.elapsed() < MAX_COMMIT_TIME => {}
/*CommitError::Retry => {
if retry_count > MAX_COMMIT_ATTEMPTS || start.elapsed() > MAX_COMMIT_TIME {
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!()));
}
}*/
CommitError::Mysql(err) => {
return Err(into_error(err));
}
CommitError::Internal(err) => {
return Err(err);
}
}
let backoff = rand::rng().random_range(50..=300);
tokio::time::sleep(Duration::from_millis(backoff)).await;
retry_count += 1;
}
}
async fn write_trx(
&self,
conn: &mut Conn,
batch: &mut Batch<'_>,
) -> Result<AssignedIds, CommitError> {
let has_changes = !batch.changes.is_empty();
let mut account_id = u32::MAX;
let mut collection = u8::MAX;
let mut document_id = u32::MAX;
let mut change_id = 0u64;
let mut asserted_values = AHashMap::new();
let mut tx_opts = TxOpts::default();
tx_opts
.with_consistent_snapshot(false)
.with_isolation_level(IsolationLevel::ReadCommitted);
let mut trx = conn.start_transaction(tx_opts).await?;
let mut result = AssignedIds::default();
if has_changes {
for &account_id in batch.changes.keys() {
let key = ValueClass::ChangeId.serialize(account_id, 0, 0, 0);
let s = trx
.prep(concat!(
"INSERT INTO n (k, v) VALUES (:k, LAST_INSERT_ID(1)) ",
"ON DUPLICATE KEY UPDATE v = LAST_INSERT_ID(v + 1)"
))
.await?;
trx.exec_drop(&s, params! {"k" => key}).await?;
let s = trx.prep("SELECT LAST_INSERT_ID()").await?;
let change_id = trx.exec_first::<i64, _, _>(&s, ()).await?.ok_or_else(|| {
mysql_async::Error::Io(mysql_async::IoError::Io(std::io::Error::other(
"LAST_INSERT_ID() did not return a value",
)))
})?;
result.push_change_id(account_id, change_id as u64);
}
}
for op in batch.ops.iter_mut() {
match op {
Operation::AccountId {
account_id: account_id_,
} => {
account_id = *account_id_;
if has_changes {
change_id = result.set_current_change_id(account_id)?;
}
}
Operation::Collection {
collection: collection_,
} => {
collection = u8::from(*collection_);
}
Operation::DocumentId {
document_id: document_id_,
} => {
document_id = *document_id_;
}
Operation::Value { class, op } => {
let key = class.serialize(account_id, collection, document_id, 0);
let subspace = class.subspace(collection);
let table = char::from(subspace);
match op {
ValueOp::Set(value) => {
if subspace != SUBSPACE_REGISTRY_IDX {
let exists = asserted_values.get(&key);
let s = if let Some(exists) = exists {
if *exists {
trx.prep(format!(
"UPDATE {} SET v = :v WHERE k = :k",
table
))
.await?
} else {
trx.prep(format!(
"INSERT INTO {} (k, v) VALUES (:k, :v)",
table
))
.await?
}
} else {
trx
.prep(
format!("INSERT INTO {} (k, v) VALUES (:k, :v) ON DUPLICATE KEY UPDATE v = VALUES(v)", table),
)
.await?
};
match trx
.exec_drop(&s, params! {"k" => key, "v" => &*value})
.await
{
Ok(_) => {
if trx.affected_rows() == 0 {
trx.rollback().await?;
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!())
.into());
}
}
Err(err) => {
trx.rollback().await?;
return Err(err.into());
}
}
} else {
let s = trx.prep("INSERT IGNORE INTO b (k) VALUES (?)").await?;
trx.exec_drop(&s, (key,)).await?;
}
}
ValueOp::SetFnc(set_op) => {
let value = (set_op.fnc)(&set_op.params, &result)?;
let exists = asserted_values.get(&key);
let s = if let Some(exists) = exists {
if *exists {
trx.prep(format!("UPDATE {} SET v = :v WHERE k = :k", table))
.await?
} else {
trx.prep(format!(
"INSERT INTO {} (k, v) VALUES (:k, :v)",
table
))
.await?
}
} else {
trx
.prep(
format!("INSERT INTO {} (k, v) VALUES (:k, :v) ON DUPLICATE KEY UPDATE v = VALUES(v)", table),
)
.await?
};
match trx.exec_drop(&s, params! {"k" => key, "v" => &value}).await {
Ok(_) => {
if trx.affected_rows() == 0 {
trx.rollback().await?;
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!())
.into());
}
}
Err(err) => {
trx.rollback().await?;
return Err(err.into());
}
}
}
ValueOp::MergeFnc(merge_op) => {
let s = trx
.prep(format!("SELECT v FROM {} WHERE k = ? FOR UPDATE", table))
.await?;
let (exists, merge_result) = trx
.exec_first::<Vec<u8>, _, _>(&s, (&key,))
.await?
.map(|bytes| {
(merge_op.fnc)(&merge_op.params, &result, Some(bytes.as_ref()))
.map(|v| (true, v))
.map_err(CommitError::from)
})
.unwrap_or_else(|| {
(merge_op.fnc)(&merge_op.params, &result, None)
.map(|v| (false, v))
.map_err(CommitError::from)
})?;
let s = if exists {
trx.prep(format!("UPDATE {} SET v = :v WHERE k = :k", table))
.await?
} else {
trx.prep(format!("INSERT INTO {} (k, v) VALUES (:k, :v)", table))
.await?
};
match merge_result {
MergeResult::Update(value) => {
if let Err(err) =
trx.exec_drop(&s, params! {"k" => key, "v" => &value}).await
{
trx.rollback().await?;
return Err(err.into());
}
}
MergeResult::Delete if exists => {
// Update asserted value
if let Some(exists) = asserted_values.get_mut(&key) {
*exists = false;
}
let s = trx
.prep(format!("DELETE FROM {} WHERE k = ?", table))
.await?;
trx.exec_drop(&s, (key,)).await?;
}
_ => (),
}
}
ValueOp::AtomicAdd(by) => {
if *by >= 0 {
let s = trx
.prep(format!(
concat!(
"INSERT INTO {} (k, v) VALUES (?, ?) ",
"ON DUPLICATE KEY UPDATE v = v + VALUES(v)"
),
table
))
.await?;
trx.exec_drop(&s, (key, &*by)).await?;
} else {
let s = trx
.prep(format!("UPDATE {table} SET v = v + ? WHERE k = ?"))
.await?;
trx.exec_drop(&s, (&*by, key)).await?;
}
}
ValueOp::AddAndGet(by) => {
let s = trx
.prep(format!(
concat!(
"INSERT INTO {} (k, v) VALUES (:k, LAST_INSERT_ID(:v)) ",
"ON DUPLICATE KEY UPDATE v = LAST_INSERT_ID(v + :v)"
),
table
))
.await?;
trx.exec_drop(&s, params! {"k" => key, "v" => &*by}).await?;
let s = trx.prep("SELECT LAST_INSERT_ID()").await?;
result.push_counter_id(
trx.exec_first::<i64, _, _>(&s, ()).await?.ok_or_else(|| {
mysql_async::Error::Io(mysql_async::IoError::Io(
std::io::Error::other(
"LAST_INSERT_ID() did not return a value",
),
))
})?,
);
}
ValueOp::Clear => {
// Update asserted value
if let Some(exists) = asserted_values.get_mut(&key) {
*exists = false;
}
let s = trx
.prep(format!("DELETE FROM {} WHERE k = ?", table))
.await?;
trx.exec_drop(&s, (key,)).await?;
}
}
}
Operation::Index { field, key, set } => {
let key = IndexKey {
account_id,
collection,
document_id,
field: *field,
key: &*key,
}
.serialize(0);
let s = if *set {
trx.prep("INSERT IGNORE INTO i (k) VALUES (?)").await?
} else {
trx.prep("DELETE FROM i WHERE k = ?").await?
};
trx.exec_drop(&s, (key,)).await?;
}
Operation::Log { collection, set } => {
let key = LogKey {
account_id,
collection: u8::from(*collection),
change_id,
}
.serialize(0);
let s = trx
.prep("INSERT INTO l (k, v) VALUES (?, ?) ON DUPLICATE KEY UPDATE v = VALUES(v)")
.await?;
trx.exec_drop(&s, (key, &*set)).await?;
}
Operation::AssertValue {
class,
assert_value,
} => {
let key = class.serialize(account_id, collection, document_id, 0);
let table = char::from(class.subspace(collection));
let s = trx
.prep(format!("SELECT v FROM {} WHERE k = ? FOR UPDATE", table))
.await?;
let (exists, matches) = trx
.exec_first::<Vec<u8>, _, _>(&s, (&key,))
.await?
.map(|bytes| (true, assert_value.matches(&bytes)))
.unwrap_or_else(|| (false, assert_value.is_none()));
if !matches {
trx.rollback().await?;
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!())
.into());
}
asserted_values.insert(key, exists);
}
}
}
trx.commit().await.map(|_| result).map_err(Into::into)
}
pub(crate) async fn purge_store(&self) -> trc::Result<()> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
for subspace in [SUBSPACE_QUOTA, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER] {
purge_table(&mut conn, char::from(subspace)).await?;
}
Ok(())
}
pub(crate) async fn delete_range(&self, from: impl Key, to: impl Key) -> trc::Result<()> {
let mut conn = self.conn_pool.get_conn().await.map_err(into_error)?;
let table = char::from(from.subspace());
let mut from = from.serialize(0);
let to = to.serialize(0);
let delete = conn
.prep(format!("DELETE FROM {table} WHERE k >= ? AND k < ?"))
.await
.map_err(into_error)?;
match conn.exec_drop(&delete, (&from, &to)).await {
Ok(_) => return Ok(()),
Err(err) if is_timeout_error(&err) => (),
Err(err) => return Err(into_error(err)),
}
let mut chunk_size = DELETE_CHUNK_SIZE;
loop {
let boundary = conn
.prep(format!(
"SELECT k FROM {table} WHERE k >= ? AND k < ? ORDER BY k ASC LIMIT 1 OFFSET {chunk_size}"
))
.await
.map_err(into_error)?;
loop {
let next = match conn
.exec_first::<Vec<u8>, _, _>(&boundary, (&from, &to))
.await
{
Ok(next) => next,
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
};
match conn
.exec_drop(&delete, (&from, next.as_ref().unwrap_or(&to)))
.await
{
Ok(_) => (),
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
}
match next {
Some(next) => from = next,
None => return Ok(()),
}
}
}
}
}
async fn purge_table(conn: &mut Conn, table: char) -> trc::Result<()> {
let s = conn
.prep(format!("DELETE FROM {table} WHERE v = 0"))
.await
.map_err(into_error)?;
match conn.exec_drop(&s, ()).await {
Ok(_) => return Ok(()),
Err(err) if is_timeout_error(&err) => (),
Err(err) => return Err(into_error(err)),
}
let purge = conn
.prep(format!(
"DELETE FROM {table} WHERE v = 0 AND k >= ? AND k < ?"
))
.await
.map_err(into_error)?;
let purge_last = conn
.prep(format!("DELETE FROM {table} WHERE v = 0 AND k >= ?"))
.await
.map_err(into_error)?;
let mut chunk_size = DELETE_CHUNK_SIZE;
let mut from = Vec::new();
loop {
let boundary = conn
.prep(format!(
"SELECT k FROM {table} WHERE k >= ? ORDER BY k ASC LIMIT 1 OFFSET {chunk_size}"
))
.await
.map_err(into_error)?;
loop {
let next = match conn.exec_first::<Vec<u8>, _, _>(&boundary, (&from,)).await {
Ok(next) => next,
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
};
let result = match &next {
Some(next) => conn.exec_drop(&purge, (&from, next)).await,
None => conn.exec_drop(&purge_last, (&from,)).await,
};
match result {
Ok(_) => (),
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
}
match next {
Some(next) => from = next,
None => return Ok(()),
}
}
}
}
impl From<trc::Error> for CommitError {
fn from(err: trc::Error) -> Self {
CommitError::Internal(err)
}
}
impl From<mysql_async::Error> for CommitError {
fn from(err: mysql_async::Error) -> Self {
CommitError::Mysql(err)
}
}
+69
View File
@@ -0,0 +1,69 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use std::ops::Range;
use crate::backend::postgres::into_pool_error;
use super::{PostgresStore, into_error};
impl PostgresStore {
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn
.prepare_cached("SELECT v FROM t WHERE k = $1")
.await
.map_err(into_error)?;
conn.query_opt(&s, &[&key])
.await
.and_then(|row| {
if let Some(row) = row {
Ok(Some(if range.start == 0 && range.end == usize::MAX {
row.try_get::<_, Vec<u8>>(0)?
} else {
let bytes = row.try_get::<_, &[u8]>(0)?;
bytes
.get(range.start..std::cmp::min(bytes.len(), range.end))
.unwrap_or_default()
.to_vec()
}))
} else {
Ok(None)
}
})
.map_err(into_error)
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn
.prepare_cached(
"INSERT INTO t (k, v) VALUES ($1, $2) ON CONFLICT (k) DO UPDATE SET v = EXCLUDED.v",
)
.await
.map_err(into_error)?;
conn.execute(&s, &[&key, &data])
.await
.map_err(into_error)
.map(|_| ())
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn
.prepare_cached("DELETE FROM t WHERE k = $1")
.await
.map_err(into_error)?;
conn.execute(&s, &[&key])
.await
.map_err(into_error)
.map(|hits| hits > 0)
}
}
+201
View File
@@ -0,0 +1,201 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{QueryResult, QueryType, backend::postgres::into_pool_error};
use bytes::BytesMut;
use futures::{TryStreamExt, pin_mut};
use tokio_postgres::types::{FromSql, ToSql, Type};
use crate::IntoRows;
use super::{PostgresStore, into_error};
impl PostgresStore {
pub(crate) async fn sql_query<T: QueryResult>(
&self,
query: &str,
params_: &[crate::Value<'_>],
) -> trc::Result<T> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn.prepare_cached(query).await.map_err(into_error)?;
let params = params_
.iter()
.map(|v| v as &(dyn tokio_postgres::types::ToSql + Sync))
.collect::<Vec<_>>();
match T::query_type() {
QueryType::Execute => conn
.execute(&s, params.as_slice())
.await
.map_or_else(|e| Err(into_error(e)), |r| Ok(T::from_exec(r as usize))),
QueryType::Exists => {
let rows = conn.query_raw(&s, params).await.map_err(into_error)?;
pin_mut!(rows);
rows.try_next()
.await
.map_or_else(|e| Err(into_error(e)), |r| Ok(T::from_exists(r.is_some())))
}
QueryType::QueryOne => conn
.query_opt(&s, params.as_slice())
.await
.map_or_else(|e| Err(into_error(e)), |r| Ok(T::from_query_one(r))),
QueryType::QueryAll => conn
.query(&s, params.as_slice())
.await
.map_or_else(|e| Err(into_error(e)), |r| Ok(T::from_query_all(r))),
}
}
}
impl ToSql for crate::Value<'_> {
fn to_sql(
&self,
ty: &tokio_postgres::types::Type,
out: &mut BytesMut,
) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>>
where
Self: Sized,
{
match self {
crate::Value::Integer(v) => match *ty {
Type::CHAR => (*v as i8).to_sql(ty, out),
Type::INT2 => (*v as i16).to_sql(ty, out),
Type::INT4 => (*v as i32).to_sql(ty, out),
_ => v.to_sql(ty, out),
},
crate::Value::Bool(v) => v.to_sql(ty, out),
crate::Value::Float(v) => {
if matches!(ty, &Type::FLOAT4) {
(*v as f32).to_sql(ty, out)
} else {
v.to_sql(ty, out)
}
}
crate::Value::Text(v) => v.to_sql(ty, out),
crate::Value::Blob(v) => v.to_sql(ty, out),
crate::Value::Null => None::<String>.to_sql(ty, out),
}
}
fn accepts(_: &tokio_postgres::types::Type) -> bool
where
Self: Sized,
{
true
}
fn to_sql_checked(
&self,
ty: &tokio_postgres::types::Type,
out: &mut BytesMut,
) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
match self {
crate::Value::Integer(v) => match *ty {
Type::CHAR => (*v as i8).to_sql_checked(ty, out),
Type::INT2 => (*v as i16).to_sql_checked(ty, out),
Type::INT4 => (*v as i32).to_sql_checked(ty, out),
_ => v.to_sql_checked(ty, out),
},
crate::Value::Bool(v) => v.to_sql_checked(ty, out),
crate::Value::Float(v) => {
if matches!(ty, &Type::FLOAT4) {
(*v as f32).to_sql_checked(ty, out)
} else {
v.to_sql_checked(ty, out)
}
}
crate::Value::Text(v) => v.to_sql_checked(ty, out),
crate::Value::Blob(v) => v.to_sql_checked(ty, out),
crate::Value::Null => None::<String>.to_sql_checked(ty, out),
}
}
}
impl IntoRows for Vec<tokio_postgres::Row> {
fn into_rows(self) -> crate::Rows {
crate::Rows {
rows: self
.into_iter()
.map(|r| crate::Row {
values: (0..r.len())
.map(|idx| r.try_get(idx).unwrap_or(crate::Value::Null))
.collect(),
})
.collect(),
}
}
fn into_named_rows(self) -> crate::NamedRows {
crate::NamedRows {
names: self
.first()
.map(|r| r.columns().iter().map(|c| c.name().to_string()).collect())
.unwrap_or_default(),
rows: self
.into_iter()
.map(|r| crate::Row {
values: (0..r.len())
.map(|idx| r.try_get(idx).unwrap_or(crate::Value::Null))
.collect(),
})
.collect(),
}
}
fn into_row(self) -> Option<crate::Row> {
unreachable!()
}
}
impl IntoRows for Option<tokio_postgres::Row> {
fn into_row(self) -> Option<crate::Row> {
self.map(|row| crate::Row {
values: (0..row.len())
.map(|idx| row.try_get(idx).unwrap_or(crate::Value::Null))
.collect(),
})
}
fn into_rows(self) -> crate::Rows {
unreachable!()
}
fn into_named_rows(self) -> crate::NamedRows {
unreachable!()
}
}
impl FromSql<'_> for crate::Value<'static> {
fn from_sql(
ty: &tokio_postgres::types::Type,
raw: &'_ [u8],
) -> Result<Self, Box<dyn std::error::Error + Sync + Send>> {
match ty {
&Type::VARCHAR | &Type::TEXT | &Type::BPCHAR | &Type::NAME | &Type::UNKNOWN => {
String::from_sql(ty, raw).map(|s| crate::Value::Text(s.into()))
}
&Type::BOOL => bool::from_sql(ty, raw).map(crate::Value::Bool),
&Type::CHAR => i8::from_sql(ty, raw).map(|v| crate::Value::Integer(v as i64)),
&Type::INT2 => i16::from_sql(ty, raw).map(|v| crate::Value::Integer(v as i64)),
&Type::INT4 => i32::from_sql(ty, raw).map(|v| crate::Value::Integer(v as i64)),
&Type::INT8 | &Type::OID => i64::from_sql(ty, raw).map(crate::Value::Integer),
&Type::FLOAT4 | &Type::FLOAT8 => f64::from_sql(ty, raw).map(crate::Value::Float),
ty if (ty.name() == "citext"
|| ty.name() == "ltree"
|| ty.name() == "lquery"
|| ty.name() == "ltxtquery") =>
{
String::from_sql(ty, raw).map(|s| crate::Value::Text(s.into()))
}
_ => Vec::<u8>::from_sql(ty, raw).map(|b| crate::Value::Blob(b.into())),
}
}
fn accepts(_: &tokio_postgres::types::Type) -> bool {
true
}
}
+265
View File
@@ -0,0 +1,265 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{PostgresStore, into_error};
use crate::{
backend::postgres::{
PsqlSearchField, into_pool_error,
search::{PG_FALLBACK_LANG, PG_LANGS, PG_UNSTEMMED_LANG},
tls::MakeRustlsConnect,
},
search::{
CalendarSearchField, ContactSearchField, EmailSearchField, SearchableField,
TracingSearchField,
},
*,
};
use ::registry::schema::{enums::PostgreSqlRecyclingMethod, structs};
use ahash::AHashSet;
use deadpool_postgres::{
Config, ManagerConfig, Object, Pool, PoolConfig, RecyclingMethod, Runtime,
};
use tokio_postgres::NoTls;
use utils::tls::rustls_client_config;
impl PostgresStore {
pub async fn open(config: structs::PostgreSqlStore) -> Result<Store, String> {
let mut cfg = Config::new();
cfg.dbname = config.database.into();
cfg.host = config.host.into();
cfg.user = config.auth_username;
cfg.password = config.auth_secret.secret().await?.map(|v| v.into_owned());
cfg.port = (config.port as u16).into();
cfg.connect_timeout = config.timeout.map(|t| t.into_inner());
cfg.options = config.options;
cfg.manager = Some(ManagerConfig {
recycling_method: match config.pool_recycling_method {
PostgreSqlRecyclingMethod::Fast => RecyclingMethod::Fast,
PostgreSqlRecyclingMethod::Verified => RecyclingMethod::Verified,
PostgreSqlRecyclingMethod::Clean => RecyclingMethod::Clean,
},
});
if let Some(max_conn) = config.pool_max_connections {
cfg.pool = PoolConfig::new(max_conn as usize).into();
}
let primary_pool = if config.use_tls {
cfg.create_pool(
Some(Runtime::Tokio1),
MakeRustlsConnect::new(rustls_client_config(config.allow_invalid_certs)?),
)
} else {
cfg.create_pool(Some(Runtime::Tokio1), NoTls)
}
.map_err(|e| format!("Failed to create connection pool: {e}"))?;
let ts_configs = discover_ts_configs(&primary_pool).await;
let mut replicas = vec![];
for replica in config.read_replicas {
let mut cfg = cfg.clone();
cfg.dbname = replica.database.into();
cfg.host = replica.host.into();
cfg.user = replica.auth_username;
cfg.password = replica.auth_secret.secret().await?.map(|v| v.into_owned());
cfg.port = (replica.port as u16).into();
cfg.options = replica.options;
replicas.push(Store::PostgreSQL(Arc::new(PostgresStore {
conn_pool: if config.use_tls {
cfg.create_pool(
Some(Runtime::Tokio1),
MakeRustlsConnect::new(rustls_client_config(config.allow_invalid_certs)?),
)
} else {
cfg.create_pool(Some(Runtime::Tokio1), NoTls)
}
.map_err(|e| format!("Failed to create connection pool: {e}"))?,
ts_configs: ts_configs.clone(),
})));
}
let primary = Store::PostgreSQL(Arc::new(PostgresStore {
conn_pool: primary_pool,
ts_configs,
}));
Ok(primary)
}
pub(crate) async fn create_storage_tables(&self) -> trc::Result<()> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
for table in [
SUBSPACE_ACL,
SUBSPACE_TASK_QUEUE,
SUBSPACE_DELETED_ITEMS,
SUBSPACE_SPAM_SAMPLES,
SUBSPACE_BLOB_LINK,
SUBSPACE_IN_MEMORY_VALUE,
SUBSPACE_PROPERTY,
SUBSPACE_REGISTRY,
SUBSPACE_REGISTRY_PK,
SUBSPACE_QUEUE_MESSAGE,
SUBSPACE_QUEUE_EVENT,
SUBSPACE_REPORT_OUT,
SUBSPACE_REPORT_IN,
SUBSPACE_LOGS,
SUBSPACE_BLOBS,
SUBSPACE_DIRECTORY,
SUBSPACE_TELEMETRY_SPAN,
SUBSPACE_TELEMETRY_METRIC,
] {
let table = char::from(table);
conn.execute(
&format!(
"CREATE TABLE IF NOT EXISTS {table} (
k BYTEA PRIMARY KEY,
v BYTEA NOT NULL
)"
),
&[],
)
.await
.map_err(into_error)?;
}
for table in [SUBSPACE_INDEXES, SUBSPACE_REGISTRY_IDX] {
let table = char::from(table);
conn.execute(
&format!(
"CREATE TABLE IF NOT EXISTS {table} (
k BYTEA PRIMARY KEY
)"
),
&[],
)
.await
.map_err(into_error)?;
}
for table in [SUBSPACE_COUNTER, SUBSPACE_QUOTA, SUBSPACE_IN_MEMORY_COUNTER] {
conn.execute(
&format!(
"CREATE TABLE IF NOT EXISTS {} (
k BYTEA PRIMARY KEY,
v BIGINT NOT NULL DEFAULT 0
)",
char::from(table)
),
&[],
)
.await
.map_err(into_error)?;
}
Ok(())
}
pub(crate) async fn create_search_tables(&self) -> trc::Result<()> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
create_search_tables::<EmailSearchField>(&conn).await?;
create_search_tables::<CalendarSearchField>(&conn).await?;
create_search_tables::<ContactSearchField>(&conn).await?;
//create_search_tables::<FileSearchField>(&conn).await?;
create_search_tables::<TracingSearchField>(&conn).await?;
Ok(())
}
}
async fn create_search_tables<T: SearchableField + PsqlSearchField + 'static>(
conn: &Object,
) -> trc::Result<()> {
let table_name = T::index().psql_table();
let mut query = format!("CREATE TABLE IF NOT EXISTS {} (", table_name);
// Add primary key columns
let pkeys = T::primary_keys();
for pkey in pkeys {
query.push_str(&format!("{} {}, ", pkey.column(), pkey.column_type()));
}
// Add other columns
for field in T::all_fields() {
query.push_str(&format!("{} {}", field.column(), field.column_type()));
if let Some(sort_type) = field.sort_column_type() {
query.push_str(&format!(", {} {}", field.sort_column().unwrap(), sort_type));
}
query.push_str(", ");
}
// Add primary key constraint
query.push_str("PRIMARY KEY (");
for (i, pkey) in pkeys.iter().enumerate() {
if i > 0 {
query.push_str(", ");
}
query.push_str(pkey.column());
}
query.push_str("))");
conn.execute(&query, &[]).await.map_err(into_error)?;
// Create indexes
for field in T::all_fields() {
if field.is_text() || field.is_json() {
let column_name = field.column();
let create_index_query = format!(
"CREATE INDEX IF NOT EXISTS gin_{table_name}_{column_name} ON {table_name} USING GIN({column_name})",
);
conn.execute(&create_index_query, &[])
.await
.map_err(into_error)?;
}
if field.is_indexed() {
let column_name = field.sort_column().unwrap_or(field.column());
let create_index_query = format!(
"CREATE INDEX IF NOT EXISTS idx_{table_name}_{column_name} ON {table_name}({column_name})",
);
conn.execute(&create_index_query, &[])
.await
.map_err(into_error)?;
}
}
Ok(())
}
async fn discover_ts_configs(pool: &Pool) -> AHashSet<&'static str> {
let mut ts_configs = AHashSet::from_iter([PG_FALLBACK_LANG, PG_UNSTEMMED_LANG]);
match probe_ts_configs(pool).await {
Ok(available) => {
for name in available {
if let Some(config) = PG_LANGS.iter().copied().find(|config| *config == name) {
ts_configs.insert(config);
}
}
}
Err(err) => {
trc::event!(
Store(trc::StoreEvent::PostgresqlError),
Details = "Failed to query pg_ts_config, assuming english only",
Reason = err.to_string(),
);
}
}
ts_configs
}
async fn probe_ts_configs(pool: &Pool) -> trc::Result<Vec<String>> {
let conn = pool.get().await.map_err(into_pool_error)?;
conn.query("SELECT cfgname::text FROM pg_ts_config", &[])
.await
.map_err(into_error)?
.into_iter()
.map(|row| row.try_get::<_, String>(0).map_err(into_error))
.collect()
}
+311
View File
@@ -0,0 +1,311 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{
search::{
CalendarSearchField, ContactSearchField, EmailSearchField, FileSearchField, SearchField,
TracingSearchField,
},
write::SearchIndex,
};
use ahash::AHashSet;
use deadpool_postgres::Pool;
use tokio_postgres::error::SqlState;
pub mod blob;
pub mod lookup;
pub mod main;
pub mod read;
pub mod search;
pub mod tls;
pub mod write;
pub struct PostgresStore {
pub(crate) conn_pool: Pool,
pub(crate) ts_configs: AHashSet<&'static str>,
}
#[inline(always)]
fn into_error(err: tokio_postgres::error::Error) -> trc::Error {
let mut local_err = trc::StoreEvent::PostgresqlError.reason(error_chain(&err));
if let Some(db_err) = err.as_db_error() {
local_err = local_err.code(db_err.code().code().to_string());
if let Some(detail) = db_err.detail() {
local_err = local_err.details(detail.to_string());
}
if let Some(hint) = db_err.hint() {
local_err = local_err.caused_by(hint.to_string());
}
}
local_err
}
fn error_chain(err: &(dyn std::error::Error + 'static)) -> String {
let mut message = err.to_string();
let mut source = err.source();
while let Some(cause) = source {
let cause_message = cause.to_string();
if !cause_message.is_empty() && !message.ends_with(&cause_message) {
message.push_str(": ");
message.push_str(&cause_message);
}
source = cause.source();
}
message
}
pub(crate) const DELETE_CHUNK_SIZE: usize = 1000;
pub(crate) const MIN_DELETE_CHUNK_SIZE: usize = 10;
#[inline(always)]
pub(crate) fn is_timeout_error(err: &tokio_postgres::Error) -> bool {
err.code().is_some_and(|code| {
*code == SqlState::QUERY_CANCELED
|| *code == SqlState::IDLE_IN_TRANSACTION_SESSION_TIMEOUT
|| *code == SqlState::LOCK_NOT_AVAILABLE
})
}
#[inline(always)]
fn into_pool_error(err: deadpool_postgres::PoolError) -> trc::Error {
match err {
deadpool_postgres::PoolError::Backend(err) => into_error(err),
err => trc::StoreEvent::PostgresqlError.reason(error_chain(&err)),
}
}
impl SearchIndex {
pub fn psql_table(&self) -> &'static str {
match self {
SearchIndex::Email => "s_email",
SearchIndex::Calendar => "s_cal",
SearchIndex::Contacts => "s_card",
SearchIndex::File => "s_file",
SearchIndex::Tracing => "s_trace",
SearchIndex::InMemory => "",
}
}
}
trait PsqlSearchField {
fn column(&self) -> &'static str;
fn column_type(&self) -> &'static str;
fn sort_column_type(&self) -> Option<&'static str>;
fn sort_column(&self) -> Option<&'static str>;
}
impl PsqlSearchField for EmailSearchField {
fn column(&self) -> &'static str {
match self {
EmailSearchField::From => "fadr",
EmailSearchField::To => "tadr",
EmailSearchField::Cc => "cc",
EmailSearchField::Bcc => "bcc",
EmailSearchField::Subject => "subj",
EmailSearchField::Body => "body",
EmailSearchField::Attachment => "atta",
EmailSearchField::ReceivedAt => "rcvd",
EmailSearchField::SentAt => "sent",
EmailSearchField::Size => "size",
EmailSearchField::HasAttachment => "hatt",
EmailSearchField::Headers => "hdrs",
}
}
fn column_type(&self) -> &'static str {
match self {
EmailSearchField::ReceivedAt | EmailSearchField::SentAt => "BIGINT",
EmailSearchField::Size => "INTEGER",
EmailSearchField::HasAttachment => "BOOLEAN",
EmailSearchField::Headers => "JSONB",
_ => "TSVECTOR",
}
}
fn sort_column_type(&self) -> Option<&'static str> {
match self {
EmailSearchField::From | EmailSearchField::To | EmailSearchField::Subject => {
Some("TEXT")
}
#[cfg(feature = "test_mode")]
EmailSearchField::Cc | EmailSearchField::Bcc => Some("TEXT"),
_ => None,
}
}
fn sort_column(&self) -> Option<&'static str> {
match self {
EmailSearchField::From => Some("s_fr"),
EmailSearchField::To => Some("s_to"),
EmailSearchField::Subject => Some("s_sj"),
#[cfg(feature = "test_mode")]
EmailSearchField::Bcc => Some("s_bc"),
#[cfg(feature = "test_mode")]
EmailSearchField::Cc => Some("s_cc"),
_ => None,
}
}
}
impl PsqlSearchField for CalendarSearchField {
fn column(&self) -> &'static str {
match self {
CalendarSearchField::Title => "titl",
CalendarSearchField::Description => "dscd",
CalendarSearchField::Location => "locn",
CalendarSearchField::Owner => "ownr",
CalendarSearchField::Attendee => "atnd",
CalendarSearchField::Start => "strt",
CalendarSearchField::Uid => "uid",
}
}
fn column_type(&self) -> &'static str {
match self {
CalendarSearchField::Start => "BIGINT",
CalendarSearchField::Uid => "TEXT",
_ => "TSVECTOR",
}
}
fn sort_column_type(&self) -> Option<&'static str> {
None
}
fn sort_column(&self) -> Option<&'static str> {
None
}
}
impl PsqlSearchField for ContactSearchField {
fn column(&self) -> &'static str {
match self {
ContactSearchField::Member => "mmbr",
ContactSearchField::Name => "name",
ContactSearchField::Nickname => "nick",
ContactSearchField::Organization => "orgn",
ContactSearchField::Email => "eml",
ContactSearchField::Phone => "phon",
ContactSearchField::OnlineService => "olsv",
ContactSearchField::Address => "addr",
ContactSearchField::Note => "note",
ContactSearchField::Kind => "kind",
ContactSearchField::Uid => "uid",
}
}
fn column_type(&self) -> &'static str {
match self {
ContactSearchField::Kind | ContactSearchField::Uid => "TEXT",
_ => "TSVECTOR",
}
}
fn sort_column_type(&self) -> Option<&'static str> {
None
}
fn sort_column(&self) -> Option<&'static str> {
None
}
}
impl PsqlSearchField for FileSearchField {
fn column(&self) -> &'static str {
match self {
FileSearchField::Name => "name",
FileSearchField::Content => "body",
}
}
fn column_type(&self) -> &'static str {
"TSVECTOR"
}
fn sort_column_type(&self) -> Option<&'static str> {
None
}
fn sort_column(&self) -> Option<&'static str> {
None
}
}
impl PsqlSearchField for TracingSearchField {
fn column(&self) -> &'static str {
match self {
TracingSearchField::QueueId => "qid",
TracingSearchField::EventType => "etyp",
TracingSearchField::Keywords => "kwds",
}
}
fn column_type(&self) -> &'static str {
match self {
TracingSearchField::EventType => "BIGINT",
TracingSearchField::QueueId => "BIGINT",
TracingSearchField::Keywords => "TSVECTOR",
}
}
fn sort_column_type(&self) -> Option<&'static str> {
None
}
fn sort_column(&self) -> Option<&'static str> {
None
}
}
impl PsqlSearchField for SearchField {
fn column(&self) -> &'static str {
match self {
SearchField::AccountId => "accid",
SearchField::DocumentId => "docid",
SearchField::Id => "id",
SearchField::Email(field) => field.column(),
SearchField::Calendar(field) => field.column(),
SearchField::Contact(field) => field.column(),
SearchField::File(field) => field.column(),
SearchField::Tracing(field) => field.column(),
}
}
fn column_type(&self) -> &'static str {
match self {
SearchField::AccountId => "INTEGER NOT NULL",
SearchField::DocumentId => "INTEGER NOT NULL",
SearchField::Id => "BIGINT NOT NULL",
SearchField::Email(field) => field.column_type(),
SearchField::Calendar(field) => field.column_type(),
SearchField::Contact(field) => field.column_type(),
SearchField::File(field) => field.column_type(),
SearchField::Tracing(field) => field.column_type(),
}
}
fn sort_column_type(&self) -> Option<&'static str> {
match self {
SearchField::Email(field) => field.sort_column_type(),
SearchField::Calendar(field) => field.sort_column_type(),
SearchField::Contact(field) => field.sort_column_type(),
SearchField::File(field) => field.sort_column_type(),
SearchField::Tracing(field) => field.sort_column_type(),
SearchField::AccountId | SearchField::DocumentId | SearchField::Id => None,
}
}
fn sort_column(&self) -> Option<&'static str> {
match self {
SearchField::Email(field) => field.sort_column(),
SearchField::Calendar(field) => field.sort_column(),
SearchField::Contact(field) => field.sort_column(),
SearchField::File(field) => field.sort_column(),
SearchField::Tracing(field) => field.sort_column(),
SearchField::AccountId | SearchField::DocumentId | SearchField::Id => None,
}
}
}
+168
View File
@@ -0,0 +1,168 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{PostgresStore, into_error, is_timeout_error};
use crate::{
Deserialize, IterateParams, Key, ValueKey, backend::postgres::into_pool_error,
write::ValueClass,
};
use futures::{TryStreamExt, pin_mut};
impl PostgresStore {
pub(crate) async fn get_value<U>(&self, key: impl Key) -> trc::Result<Option<U>>
where
U: Deserialize + 'static,
{
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn
.prepare_cached(&format!(
"SELECT v FROM {} WHERE k = $1",
char::from(key.subspace())
))
.await
.map_err(into_error)?;
let key = key.serialize(0);
conn.query_opt(&s, &[&key])
.await
.map_err(into_error)
.and_then(|r| {
if let Some(r) = r {
Ok(Some(U::deserialize_with_key(&key, r.get(0))?))
} else {
Ok(None)
}
})
}
pub(crate) async fn key_exists(&self, key: impl Key) -> trc::Result<bool> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn
.prepare_cached(&format!(
"SELECT 1 FROM {} WHERE k = $1",
char::from(key.subspace())
))
.await
.map_err(into_error)?;
let key = key.serialize(0);
conn.query_opt(&s, &[&key])
.await
.map_err(into_error)
.map(|r| r.is_some())
}
pub(crate) async fn iterate<T: Key>(
&self,
params: IterateParams<T>,
mut cb: impl for<'x> FnMut(&'x [u8], &'x [u8]) -> trc::Result<bool> + Sync + Send,
) -> trc::Result<()> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let table = char::from(params.begin.subspace());
let begin = params.begin.serialize(0);
let end = params.end.serialize(0);
let keys = if params.values { "k, v" } else { "k" };
let s = conn
.prepare_cached(&match (params.first, params.ascending) {
(true, true) => {
format!(
"SELECT {keys} FROM {table} WHERE k >= $1 AND k <= $2 ORDER BY k ASC LIMIT 1"
)
}
(true, false) => {
format!(
"SELECT {keys} FROM {table} WHERE k >= $1 AND k <= $2 ORDER BY k DESC LIMIT 1"
)
}
(false, true) => {
format!("SELECT {keys} FROM {table} WHERE k >= $1 AND k <= $2 ORDER BY k ASC")
}
(false, false) => {
format!("SELECT {keys} FROM {table} WHERE k >= $1 AND k <= $2 ORDER BY k DESC")
}
})
.await.map_err(into_error)?;
let mut from = begin;
let mut to = end;
let mut resume_key: Option<Vec<u8>> = None;
loop {
let mut last_key = None;
let mut timed_out = false;
{
let rows = conn
.query_raw(&s, &[&from, &to])
.await
.map_err(into_error)?;
pin_mut!(rows);
loop {
match rows.try_next().await {
Ok(Some(row)) => {
let key = row.try_get::<_, &[u8]>(0).map_err(into_error)?;
let value = if params.values {
row.try_get::<_, &[u8]>(1).map_err(into_error)?
} else {
b"".as_slice()
};
if resume_key.take().is_some_and(|resumed| resumed == key) {
continue;
}
if !cb(key, value)? {
return Ok(());
}
last_key = Some(key.to_vec());
}
Ok(None) => break,
Err(err) => {
if params.first || last_key.is_none() || !is_timeout_error(&err) {
return Err(into_error(err));
}
timed_out = true;
break;
}
}
}
}
match last_key {
Some(last_key) if timed_out => {
if params.ascending {
from.clone_from(&last_key);
} else {
to.clone_from(&last_key);
}
resume_key = Some(last_key);
}
_ => return Ok(()),
}
}
}
pub(crate) async fn get_counter(
&self,
key: impl Into<ValueKey<ValueClass>> + Sync + Send,
) -> trc::Result<i64> {
let key = key.into();
let table = char::from(key.subspace());
let key = key.serialize(0);
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn
.prepare_cached(&format!("SELECT v FROM {table} WHERE k = $1"))
.await
.map_err(into_error)?;
match conn.query_opt(&s, &[&key]).await {
Ok(Some(row)) => row.try_get(0).map_err(into_error),
Ok(None) => Ok(0),
Err(e) => Err(into_error(e)),
}
}
}
+571
View File
@@ -0,0 +1,571 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::{
backend::postgres::{
DELETE_CHUNK_SIZE, MIN_DELETE_CHUNK_SIZE, PostgresStore, PsqlSearchField, into_error,
into_pool_error, is_timeout_error,
},
search::{
IndexDocument, SearchComparator, SearchDocumentId, SearchFilter, SearchOperator,
SearchQuery, SearchValue,
},
write::SearchIndex,
};
use nlp::language::Language;
use std::fmt::Write;
use tokio_postgres::{
IsolationLevel,
types::{FromSql, ToSql, Type, WrongType},
};
impl PostgresStore {
fn ts_config(&self, language: &Language) -> &'static str {
pg_lang(language)
.filter(|config| self.ts_configs.contains(config))
.unwrap_or(PG_UNSTEMMED_LANG)
}
pub async fn index(&self, documents: Vec<IndexDocument>) -> trc::Result<()> {
let mut conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let trx = conn
.build_transaction()
.isolation_level(IsolationLevel::ReadCommitted)
.start()
.await
.map_err(into_error)?;
for document in documents {
let index = document.index;
let primary_keys = index.primary_keys();
let all_fields = index.all_fields();
let fields = document.fields;
let mut values = Vec::with_capacity(fields.len() + 2);
let mut query = format!("INSERT INTO {} (", index.psql_table());
for (i, field) in primary_keys.iter().chain(all_fields).enumerate() {
if i > 0 {
query.push(',');
}
query.push_str(field.column());
if let Some(sort_column) = field.sort_column() {
query.push(',');
query.push_str(sort_column);
}
}
query.push_str(") VALUES (");
for (i, field) in primary_keys.iter().chain(all_fields).enumerate() {
if i > 0 {
query.push(',');
}
if let Some(value) = fields.get(field) {
let value_ref = format!("${}", values.len() + 1);
let (text_len, language) = if let SearchValue::Text { value, language } = value
{
(value.len(), self.ts_config(language))
} else {
(0, PG_UNSTEMMED_LANG)
};
if field.is_text() {
let _ = write!(&mut query, "to_tsvector('{language}',{value_ref})");
} else if text_len > 512 {
query.push_str("left(");
query.push_str(&value_ref);
query.push_str(",512)");
} else {
query.push_str(&value_ref);
}
if field.sort_column().is_some() {
if text_len > 255 {
query.push_str(",left(");
query.push_str(&value_ref);
query.push_str(",255)");
} else {
query.push(',');
query.push_str(&value_ref);
}
}
values.push(value as &(dyn ToSql + Sync));
} else {
query.push_str("NULL");
if field.sort_column().is_some() {
query.push_str(",NULL");
}
}
}
query.push_str(") ON CONFLICT (");
for (i, pkey) in primary_keys.iter().enumerate() {
if i > 0 {
query.push(',');
}
query.push_str(pkey.column());
}
query.push_str(") DO UPDATE SET ");
for (i, field) in all_fields.iter().enumerate() {
if i > 0 {
query.push(',');
}
let column = field.column();
let _ = write!(&mut query, "{column} = EXCLUDED.{column}");
}
trx.execute(&query, &values).await.map_err(into_error)?;
}
trx.commit().await.map_err(into_error)
}
pub async fn query<R: SearchDocumentId>(
&self,
index: SearchIndex,
filters: &[SearchFilter],
sort: &[SearchComparator],
) -> trc::Result<Vec<R>> {
let mut query = format!("SELECT {} FROM {}", R::field().column(), index.psql_table());
let params = self.build_filter(&mut query, filters);
if !sort.is_empty() {
build_sort(&mut query, sort);
}
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn.prepare_cached(&query).await.map_err(into_error)?;
conn.query(&s, params.as_slice())
.await
.and_then(|rows| {
rows.into_iter()
.map(|row| row.try_get::<_, DocId>(0).map(|v| R::from_u64(v.0)))
.collect::<Result<Vec<R>, _>>()
})
.map_err(into_error)
}
pub async fn unindex(&self, filter: SearchQuery) -> trc::Result<u64> {
debug_assert!(!filter.filters.is_empty());
let table = filter.index.psql_table();
let mut where_clause = String::new();
let params = self.build_filter(&mut where_clause, &filter.filters);
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let s = conn
.prepare_cached(&format!("DELETE FROM {table}{where_clause}"))
.await
.map_err(into_error)?;
match conn.execute(&s, params.as_slice()).await {
Ok(deleted) => return Ok(deleted),
Err(err) if is_timeout_error(&err) => (),
Err(err) => return Err(into_error(err)),
}
let mut chunk_size = DELETE_CHUNK_SIZE;
let mut deleted = 0;
loop {
let s = conn
.prepare_cached(&format!(
"DELETE FROM {table} WHERE ctid IN (SELECT ctid FROM {table}{where_clause} LIMIT {chunk_size})"
))
.await
.map_err(into_error)?;
loop {
match conn.execute(&s, params.as_slice()).await {
Ok(0) => return Ok(deleted),
Ok(affected) => deleted += affected,
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
}
}
}
}
fn build_filter<'x>(
&self,
query: &mut String,
filters: &'x [SearchFilter],
) -> Vec<&'x (dyn ToSql + Sync)> {
if filters.is_empty() {
return Vec::new();
}
query.push_str(" WHERE ");
let mut operator_stack = Vec::new();
let mut operator = &SearchFilter::And;
let mut is_first = true;
let mut values = Vec::new();
for filter in filters {
match filter {
SearchFilter::Operator { field, op, value } => {
if !is_first {
match operator {
SearchFilter::And => query.push_str(" AND "),
SearchFilter::Or => query.push_str(" OR "),
_ => (),
}
} else {
is_first = false;
}
let value_pos = values.len() + 1;
if field.is_text()
&& matches!(op, SearchOperator::Equal | SearchOperator::Contains)
{
query.push_str(field.column());
query.push(' ');
let language = match &value {
SearchValue::Text { language, .. } => *language,
_ => Language::None,
};
let config = self.ts_config(&language);
let method = match op {
SearchOperator::Equal => "phraseto_tsquery",
_ => "plainto_tsquery",
};
if matches!(language, Language::None) {
let _ = write!(query, "@@ {method}('{config}', ${value_pos})");
} else {
let _ = write!(query, "@@ ({method}('{config}', ${value_pos})");
for fallback in [PG_FALLBACK_LANG, PG_UNSTEMMED_LANG] {
if fallback != config && self.ts_configs.contains(fallback) {
let _ =
write!(query, " || {method}('{fallback}', ${value_pos})");
}
}
query.push(')');
}
values.push(value as &(dyn ToSql + Sync));
} else if let SearchValue::KeyValues(kv) = value {
query.push_str(field.column());
query.push(' ');
let (key, value) = kv.iter().next().unwrap();
values.push(key as &(dyn ToSql + Sync));
if !value.is_empty() {
let _ = write!(query, "->> ${value_pos} ");
op.write_pqsql(query, values.len() + 1);
values.push(value as &(dyn ToSql + Sync));
} else {
let _ = write!(query, " ? ${value_pos}");
}
} else {
query.push_str(field.sort_column().unwrap_or(field.column()));
query.push(' ');
op.write_pqsql(query, value_pos);
values.push(value as &(dyn ToSql + Sync));
}
}
SearchFilter::And | SearchFilter::Or => {
if !is_first {
match operator {
SearchFilter::And => query.push_str(" AND "),
SearchFilter::Or => query.push_str(" OR "),
_ => (),
}
} else {
is_first = false;
}
operator_stack.push((operator, is_first));
operator = filter;
is_first = true;
query.push('(');
}
SearchFilter::Not => {
if !is_first {
match operator {
SearchFilter::And => query.push_str(" AND "),
SearchFilter::Or => query.push_str(" OR "),
_ => (),
}
} else {
is_first = false;
}
operator_stack.push((operator, is_first));
operator = &SearchFilter::And;
is_first = true;
query.push_str("NOT (");
}
SearchFilter::End => {
let p = operator_stack.pop().unwrap_or((&SearchFilter::And, true));
operator = p.0;
is_first = p.1;
query.push(')');
}
SearchFilter::DocumentSet(_) => {
debug_assert!(
false,
"DocumentSet filters are not supported in Postgres backend"
)
}
}
}
values
}
}
fn build_sort(query: &mut String, sort: &[SearchComparator]) {
query.push_str(" ORDER BY ");
for (i, comparator) in sort.iter().enumerate() {
if i > 0 {
query.push_str(", ");
}
match comparator {
SearchComparator::Field { field, ascending } => {
query.push_str(field.sort_column().unwrap_or(field.column()));
if *ascending {
query.push_str(" ASC");
} else {
query.push_str(" DESC");
}
}
SearchComparator::DocumentSet { .. } | SearchComparator::SortedSet { .. } => {
debug_assert!(
false,
"DocumentSet and SortedSet comparators are not supported "
);
}
}
}
}
impl ToSql for SearchValue {
fn to_sql(
&self,
ty: &tokio_postgres::types::Type,
out: &mut bytes::BytesMut,
) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>>
where
Self: Sized,
{
match self {
SearchValue::Text { value, .. } => {
// Truncate large text fields to avoid Postgres errors (see https://www.postgresql.org/docs/current/textsearch-limitations.html)
if value.len() > 650_000 {
(&value[..value.floor_char_boundary(650_000)]).to_sql(ty, out)
} else {
value.to_sql(ty, out)
}
}
SearchValue::Int(v) => match *ty {
Type::INT4 => (*v as i32).to_sql(ty, out),
_ => v.to_sql(ty, out),
},
SearchValue::Uint(v) => match *ty {
Type::INT4 => (*v as i32).to_sql(ty, out),
_ => (*v as i64).to_sql(ty, out),
},
SearchValue::Boolean(v) => v.to_sql(ty, out),
SearchValue::KeyValues(kv) => {
serde_json::to_value(kv).unwrap_or_default().to_sql(ty, out)
}
}
}
fn accepts(_: &tokio_postgres::types::Type) -> bool
where
Self: Sized,
{
true
}
fn to_sql_checked(
&self,
ty: &tokio_postgres::types::Type,
out: &mut bytes::BytesMut,
) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
match self {
SearchValue::Text { value, .. } => {
// Truncate large text fields to avoid Postgres errors (see https://www.postgresql.org/docs/current/textsearch-limitations.html)
if value.len() > 650_000 {
(&value[..value.floor_char_boundary(650_000)]).to_sql_checked(ty, out)
} else {
value.to_sql_checked(ty, out)
}
}
SearchValue::Int(v) => match *ty {
Type::INT4 => (*v as i32).to_sql_checked(ty, out),
_ => v.to_sql_checked(ty, out),
},
SearchValue::Uint(v) => match *ty {
Type::INT4 => (*v as i32).to_sql_checked(ty, out),
_ => (*v as i64).to_sql_checked(ty, out),
},
SearchValue::Boolean(v) => v.to_sql_checked(ty, out),
SearchValue::KeyValues(kv) => serde_json::to_value(kv)
.unwrap_or_default()
.to_sql_checked(ty, out),
}
}
}
struct DocId(u64);
impl FromSql<'_> for DocId {
fn from_sql(
ty: &tokio_postgres::types::Type,
raw: &'_ [u8],
) -> Result<Self, Box<dyn std::error::Error + Sync + Send>> {
match ty {
&Type::INT4 => i32::from_sql(ty, raw).map(|v| DocId(v as u64)),
&Type::INT8 | &Type::OID => i64::from_sql(ty, raw).map(|v| DocId(v as u64)),
_ => Err(Box::new(WrongType::new::<DocId>(ty.clone()))),
}
}
fn accepts(typ: &Type) -> bool {
matches!(typ, &Type::INT4 | &Type::INT8 | &Type::OID)
}
}
impl SearchOperator {
fn write_pqsql(&self, query: &mut String, value_pos: usize) {
match self {
SearchOperator::LowerThan => {
let _ = write!(query, "< ${value_pos}");
}
SearchOperator::LowerEqualThan => {
let _ = write!(query, "<= ${value_pos}");
}
SearchOperator::GreaterThan => {
let _ = write!(query, "> ${value_pos}");
}
SearchOperator::GreaterEqualThan => {
let _ = write!(query, ">= ${value_pos}");
}
SearchOperator::Equal => {
let _ = write!(query, "= ${value_pos}");
}
SearchOperator::Contains => {
let _ = write!(query, "LIKE '%' || ${value_pos} || '%'");
}
}
}
}
pub(super) const PG_FALLBACK_LANG: &str = "english";
pub(super) const PG_UNSTEMMED_LANG: &str = "simple";
pub(super) const PG_LANGS: &[&str] = &[
"arabic",
"armenian",
"catalan",
"danish",
"dutch",
"english",
"finnish",
"french",
"german",
"greek",
"hindi",
"hungarian",
"indonesian",
"italian",
"lithuanian",
"nepali",
"norwegian",
"portuguese",
"romanian",
"russian",
"serbian",
"spanish",
"swedish",
"tamil",
"turkish",
"yiddish",
];
#[inline(always)]
fn pg_lang(lang: &Language) -> Option<&'static str> {
match lang {
Language::Esperanto => None,
Language::English => Some("english"),
Language::Russian => Some("russian"),
Language::Mandarin => None,
Language::Spanish => Some("spanish"),
Language::Portuguese => Some("portuguese"),
Language::Italian => Some("italian"),
Language::Bengali => None,
Language::French => Some("french"),
Language::German => Some("german"),
Language::Ukrainian => None,
Language::Georgian => None,
Language::Arabic => Some("arabic"),
Language::Hindi => Some("hindi"),
Language::Japanese => None,
Language::Hebrew => None,
Language::Yiddish => Some("yiddish"),
Language::Polish => None,
Language::Amharic => None,
Language::Javanese => None,
Language::Korean => None,
Language::Bokmal => Some("norwegian"), // Norwegian covers Bokmål
Language::Danish => Some("danish"),
Language::Swedish => Some("swedish"),
Language::Finnish => Some("finnish"),
Language::Turkish => Some("turkish"),
Language::Dutch => Some("dutch"),
Language::Hungarian => Some("hungarian"),
Language::Czech => None,
Language::Greek => Some("greek"),
Language::Bulgarian => None,
Language::Belarusian => None,
Language::Marathi => None,
Language::Kannada => None,
Language::Romanian => Some("romanian"),
Language::Slovene => None,
Language::Croatian => None,
Language::Serbian => Some("serbian"),
Language::Macedonian => None,
Language::Lithuanian => Some("lithuanian"),
Language::Latvian => None,
Language::Estonian => None,
Language::Tamil => Some("tamil"),
Language::Vietnamese => None,
Language::Urdu => None,
Language::Thai => None,
Language::Gujarati => None,
Language::Uzbek => None,
Language::Punjabi => None,
Language::Azerbaijani => None,
Language::Indonesian => Some("indonesian"),
Language::Telugu => None,
Language::Persian => None,
Language::Malayalam => None,
Language::Oriya => None,
Language::Burmese => None,
Language::Nepali => Some("nepali"),
Language::Sinhalese => None,
Language::Khmer => None,
Language::Turkmen => None,
Language::Akan => None,
Language::Zulu => None,
Language::Shona => None,
Language::Afrikaans => None,
Language::Latin => None,
Language::Slovak => None,
Language::Catalan => Some("catalan"),
Language::Tagalog => None,
Language::Armenian => Some("armenian"),
Language::Unknown | Language::None => None,
}
}
+198
View File
@@ -0,0 +1,198 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
// Credits: https://github.com/jbg/tokio-postgres-rustls
use std::{
convert::TryFrom,
future::Future,
io,
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use aws_lc_rs::digest;
use futures::future::{FutureExt, TryFutureExt};
use rustls::ClientConfig;
use rustls_pki_types::ServerName;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio_postgres::tls::{ChannelBinding, MakeTlsConnect, TlsConnect};
use tokio_rustls::{TlsConnector, client::TlsStream};
use x509_parser::{
asn1_rs::oid,
oid_registry::{
OID_HASH_SHA1, OID_MD5_WITH_RSA, OID_NIST_HASH_SHA256, OID_NIST_HASH_SHA384,
OID_NIST_HASH_SHA512, OID_PKCS1_MD5WITHRSAENC, OID_PKCS1_RSASSAPSS, OID_PKCS1_SHA1WITHRSA,
OID_PKCS1_SHA224WITHRSA, OID_PKCS1_SHA256WITHRSA, OID_PKCS1_SHA384WITHRSA,
OID_PKCS1_SHA512WITHRSA, OID_SHA1_WITH_RSA, OID_SIG_DSA_WITH_SHA1,
OID_SIG_ECDSA_WITH_SHA224, OID_SIG_ECDSA_WITH_SHA256, OID_SIG_ECDSA_WITH_SHA384,
OID_SIG_ECDSA_WITH_SHA512,
},
parse_x509_certificate,
prelude::X509Certificate,
signature_algorithm::RsaSsaPssParams,
};
#[derive(Clone)]
pub struct MakeRustlsConnect {
config: Arc<ClientConfig>,
}
impl MakeRustlsConnect {
pub fn new(config: ClientConfig) -> Self {
Self {
config: Arc::new(config),
}
}
}
impl<S> MakeTlsConnect<S> for MakeRustlsConnect
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
type Stream = RustlsStream<S>;
type TlsConnect = RustlsConnect;
type Error = io::Error;
fn make_tls_connect(&mut self, hostname: &str) -> io::Result<RustlsConnect> {
ServerName::try_from(hostname.to_string())
.map(|dns_name| {
RustlsConnect(Some(RustlsConnectData {
hostname: dns_name,
connector: Arc::clone(&self.config).into(),
}))
})
.or(Ok(RustlsConnect(None)))
}
}
pub struct RustlsConnect(Option<RustlsConnectData>);
struct RustlsConnectData {
hostname: ServerName<'static>,
connector: TlsConnector,
}
impl<S> TlsConnect<S> for RustlsConnect
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
type Stream = RustlsStream<S>;
type Error = io::Error;
type Future = Pin<Box<dyn Future<Output = io::Result<RustlsStream<S>>> + Send>>;
fn connect(self, stream: S) -> Self::Future {
match self.0 {
None => Box::pin(core::future::ready(Err(io::ErrorKind::InvalidInput.into()))),
Some(c) => c
.connector
.connect(c.hostname, stream)
.map_ok(|s| RustlsStream(Box::pin(s)))
.boxed(),
}
}
}
pub struct RustlsStream<S>(Pin<Box<TlsStream<S>>>);
fn cb_digest_for_cert(cert: &X509Certificate<'_>) -> Option<&'static digest::Algorithm> {
let sig_alg = cert.signature_algorithm.oid();
// Signature algorithms that use a digest should use the same digest for channel binding:
if sig_alg == &OID_PKCS1_SHA512WITHRSA || sig_alg == &OID_SIG_ECDSA_WITH_SHA512 {
Some(&digest::SHA512)
} else if sig_alg == &OID_PKCS1_SHA384WITHRSA || sig_alg == &OID_SIG_ECDSA_WITH_SHA384 {
Some(&digest::SHA384)
} else if sig_alg == &OID_PKCS1_MD5WITHRSAENC
|| sig_alg == &OID_MD5_WITH_RSA
|| sig_alg == &OID_PKCS1_SHA1WITHRSA
|| sig_alg == &OID_SHA1_WITH_RSA
|| sig_alg == &OID_SIG_DSA_WITH_SHA1
|| sig_alg == &OID_PKCS1_SHA256WITHRSA
|| sig_alg == &OID_SIG_ECDSA_WITH_SHA256
{
// ...apart from MD5 or SHA1, which use SHA256 for channel binding, as per RFC 5929 section 4.1:
Some(&digest::SHA256)
} else if sig_alg == &OID_PKCS1_SHA224WITHRSA || sig_alg == &OID_SIG_ECDSA_WITH_SHA224 {
Some(&digest::SHA224)
} else if sig_alg == &OID_PKCS1_RSASSAPSS {
// For RSASSA-PSS, the hash algorithm is specified in the parameters of the signature algorithm:
let params_any = cert.signature_algorithm.parameters()?;
let pss = RsaSsaPssParams::try_from(params_any).ok()?;
let alg = pss.hash_algorithm_oid();
if alg == &OID_NIST_HASH_SHA512 {
Some(&digest::SHA512)
} else if alg == &OID_NIST_HASH_SHA384 {
Some(&digest::SHA384)
} else if alg == &OID_NIST_HASH_SHA256 || alg == &OID_HASH_SHA1 {
Some(&digest::SHA256)
} else if alg == &oid!(2.16.840.1.101.3.4.2.4) {
// id-sha224 from RFC 4055 ^
Some(&digest::SHA224)
} else {
None
}
} else {
None
}
}
impl<S> tokio_postgres::tls::TlsStream for RustlsStream<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
fn channel_binding(&self) -> ChannelBinding {
let (_, session) = self.0.get_ref();
match session.peer_certificates() {
Some(certs) if !certs.is_empty() => match parse_x509_certificate(certs[0].as_ref()) {
Ok((_, cert)) => {
if let Some(digest_alg) = cb_digest_for_cert(&cert) {
let dgst = digest::digest(digest_alg, certs[0].as_ref());
ChannelBinding::tls_server_end_point(dgst.as_ref().into())
} else {
ChannelBinding::none()
}
}
Err(_) => ChannelBinding::none(),
},
_ => ChannelBinding::none(),
}
}
}
impl<S> AsyncRead for RustlsStream<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context,
buf: &mut ReadBuf<'_>,
) -> Poll<tokio::io::Result<()>> {
self.0.as_mut().poll_read(cx, buf)
}
}
impl<S> AsyncWrite for RustlsStream<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context,
buf: &[u8],
) -> Poll<tokio::io::Result<usize>> {
self.0.as_mut().poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<tokio::io::Result<()>> {
self.0.as_mut().poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<tokio::io::Result<()>> {
self.0.as_mut().poll_shutdown(cx)
}
}
+543
View File
@@ -0,0 +1,543 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{PostgresStore, into_error, is_timeout_error};
use crate::{
IndexKey, Key, LogKey, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER, SUBSPACE_QUOTA,
SUBSPACE_REGISTRY_IDX,
backend::postgres::{DELETE_CHUNK_SIZE, MIN_DELETE_CHUNK_SIZE, into_pool_error},
write::{
AssignedIds, Batch, MAX_COMMIT_ATTEMPTS, MAX_COMMIT_TIME, MergeResult, Operation,
ValueClass, ValueOp,
},
};
use ahash::AHashMap;
use deadpool_postgres::Object;
use rand::RngExt;
use std::time::{Duration, Instant};
use tokio_postgres::{IsolationLevel, error::SqlState};
#[derive(Debug)]
enum CommitError {
Postgres(tokio_postgres::Error),
Internal(trc::Error),
//Retry,
}
impl PostgresStore {
pub(crate) async fn write(&self, mut batch: Batch<'_>) -> trc::Result<AssignedIds> {
let mut conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let start = Instant::now();
let mut retry_count = 0;
loop {
match self.write_trx(&mut conn, &mut batch).await {
Ok(result) => {
return Ok(result);
}
Err(err) => {
match err {
CommitError::Postgres(err) => match err.code() {
Some(
&SqlState::T_R_SERIALIZATION_FAILURE
| &SqlState::T_R_DEADLOCK_DETECTED,
) if retry_count < MAX_COMMIT_ATTEMPTS
&& start.elapsed() < MAX_COMMIT_TIME => {}
Some(&SqlState::UNIQUE_VIOLATION) => {
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.reason("Unique violation")
.caused_by(trc::location!()));
}
_ => return Err(into_error(err)),
},
CommitError::Internal(err) => return Err(err),
/*CommitError::Retry => {
if retry_count > MAX_COMMIT_ATTEMPTS
|| start.elapsed() > MAX_COMMIT_TIME
{
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!()));
}
}*/
}
let backoff = rand::rng().random_range(50..=300);
tokio::time::sleep(Duration::from_millis(backoff)).await;
retry_count += 1;
}
}
}
}
async fn write_trx(
&self,
conn: &mut Object,
batch: &mut Batch<'_>,
) -> Result<AssignedIds, CommitError> {
let mut account_id = u32::MAX;
let mut collection = u8::MAX;
let mut document_id = u32::MAX;
let mut change_id = 0u64;
let mut asserted_values = AHashMap::new();
let trx = conn
.build_transaction()
.isolation_level(IsolationLevel::ReadCommitted)
.start()
.await?;
let mut result = AssignedIds::default();
let has_changes = !batch.changes.is_empty();
if has_changes {
for &account_id in batch.changes.keys() {
let key = ValueClass::ChangeId.serialize(account_id, 0, 0, 0);
let s = trx
.prepare_cached(concat!(
"INSERT INTO n (k, v) VALUES ($1, 1) ",
"ON CONFLICT(k) DO UPDATE SET v = n.v + 1 RETURNING v"
))
.await?;
let change_id = trx
.query_one(&s, &[&key])
.await
.and_then(|row| row.try_get::<_, i64>(0))?;
result.push_change_id(account_id, change_id as u64);
}
}
for op in batch.ops.iter_mut() {
match op {
Operation::AccountId {
account_id: account_id_,
} => {
account_id = *account_id_;
if has_changes {
change_id = result.set_current_change_id(account_id)?;
}
}
Operation::Collection {
collection: collection_,
} => {
collection = u8::from(*collection_);
}
Operation::DocumentId {
document_id: document_id_,
} => {
document_id = *document_id_;
}
Operation::Value { class, op } => {
let key = class.serialize(account_id, collection, document_id, 0);
let subspace = class.subspace(collection);
let table = char::from(subspace);
match op {
ValueOp::Set(value) => {
if subspace != SUBSPACE_REGISTRY_IDX {
let s = if let Some(exists) = asserted_values.get(&key) {
if *exists {
trx.prepare_cached(&format!(
"UPDATE {} SET v = $2 WHERE k = $1",
table
))
.await?
} else {
trx.prepare_cached(&format!(
"INSERT INTO {} (k, v) VALUES ($1, $2)",
table
))
.await?
}
} else {
trx.prepare_cached(&format!(
concat!(
"INSERT INTO {} (k, v) VALUES ($1, $2) ",
"ON CONFLICT (k) DO UPDATE SET v = EXCLUDED.v"
),
table
))
.await?
};
if trx.execute(&s, &[&key, &(*value)]).await? == 0 {
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!())
.into());
}
} else {
let s = trx
.prepare_cached(
"INSERT INTO b (k) VALUES ($1) ON CONFLICT (k) DO NOTHING",
)
.await?;
trx.execute(&s, &[&key]).await?;
}
}
ValueOp::SetFnc(set_op) => {
let value = (set_op.fnc)(&set_op.params, &result)?;
let s = if let Some(exists) = asserted_values.get(&key) {
if *exists {
trx.prepare_cached(&format!(
"UPDATE {} SET v = $2 WHERE k = $1",
table
))
.await?
} else {
trx.prepare_cached(&format!(
"INSERT INTO {} (k, v) VALUES ($1, $2)",
table
))
.await?
}
} else {
trx.prepare_cached(&format!(
concat!(
"INSERT INTO {} (k, v) VALUES ($1, $2) ",
"ON CONFLICT (k) DO UPDATE SET v = EXCLUDED.v"
),
table
))
.await?
};
if trx.execute(&s, &[&key, &value]).await? == 0 {
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!())
.into());
}
}
ValueOp::MergeFnc(merge_op) => {
let s = trx
.prepare_cached(&format!(
"SELECT v FROM {} WHERE k = $1 FOR UPDATE",
table
))
.await?;
let (exists, merge_result) = trx
.query_opt(&s, &[&key])
.await?
.map(|row| {
row.try_get::<_, &[u8]>(0)
.map_err(CommitError::from)
.and_then(|v| {
(merge_op.fnc)(&merge_op.params, &result, Some(v))
.map(|v| (true, v))
.map_err(CommitError::from)
})
})
.unwrap_or_else(|| {
(merge_op.fnc)(&merge_op.params, &result, None)
.map(|v| (false, v))
.map_err(CommitError::from)
})?;
match merge_result {
MergeResult::Update(value) => {
let s = if exists {
trx.prepare_cached(&format!(
"UPDATE {} SET v = $2 WHERE k = $1",
table
))
.await?
} else {
trx.prepare_cached(&format!(
"INSERT INTO {} (k, v) VALUES ($1, $2)",
table
))
.await?
};
trx.execute(&s, &[&key, &value]).await?;
}
MergeResult::Delete if exists => {
let s = trx
.prepare_cached(&format!(
"DELETE FROM {} WHERE k = $1",
table
))
.await?;
trx.execute(&s, &[&key]).await?;
// Update asserted value
if let Some(exists) = asserted_values.get_mut(&key) {
*exists = false;
}
}
_ => (),
}
}
ValueOp::AtomicAdd(by) => {
if *by >= 0 {
let s = trx
.prepare_cached(&format!(
concat!(
"INSERT INTO {} (k, v) VALUES ($1, $2) ",
"ON CONFLICT(k) DO UPDATE SET v = {}.v + EXCLUDED.v"
),
table, table
))
.await?;
trx.execute(&s, &[&key, &*by]).await?;
} else {
let s = trx
.prepare_cached(&format!(
"UPDATE {table} SET v = v + $1 WHERE k = $2"
))
.await?;
trx.execute(&s, &[&*by, &key]).await?;
}
}
ValueOp::AddAndGet(by) => {
let s = trx
.prepare_cached(&format!(
concat!(
"INSERT INTO {} (k, v) VALUES ($1, $2) ",
"ON CONFLICT(k) DO UPDATE SET v = {}.v + EXCLUDED.v RETURNING v"
),
table, table
))
.await?;
result.push_counter_id(
trx.query_one(&s, &[&key, &*by])
.await
.and_then(|row| row.try_get::<_, i64>(0))?,
);
}
ValueOp::Clear => {
let s = trx
.prepare_cached(&format!("DELETE FROM {} WHERE k = $1", table))
.await?;
trx.execute(&s, &[&key]).await?;
// Update asserted value
if let Some(exists) = asserted_values.get_mut(&key) {
*exists = false;
}
}
}
}
Operation::Index { field, key, set } => {
let key = IndexKey {
account_id,
collection,
document_id,
field: *field,
key: &*key,
}
.serialize(0);
let s = if *set {
trx.prepare_cached(
"INSERT INTO i (k) VALUES ($1) ON CONFLICT (k) DO NOTHING",
)
.await?
} else {
trx.prepare_cached("DELETE FROM i WHERE k = $1").await?
};
trx.execute(&s, &[&key]).await?;
}
Operation::Log { collection, set } => {
let key = LogKey {
account_id,
collection: u8::from(*collection),
change_id,
}
.serialize(0);
let s = trx
.prepare_cached(concat!(
"INSERT INTO l (k, v) VALUES ($1, $2) ",
"ON CONFLICT (k) DO UPDATE SET v = EXCLUDED.v"
))
.await?;
trx.execute(&s, &[&key, &*set]).await?;
}
Operation::AssertValue {
class,
assert_value,
} => {
let key = class.serialize(account_id, collection, document_id, 0);
let table = char::from(class.subspace(collection));
let s = trx
.prepare_cached(&format!("SELECT v FROM {} WHERE k = $1 FOR UPDATE", table))
.await?;
let (exists, matches) = trx
.query_opt(&s, &[&key])
.await?
.map(|row| {
row.try_get::<_, &[u8]>(0)
.map_or((true, false), |v| (true, assert_value.matches(v)))
})
.unwrap_or_else(|| (false, assert_value.is_none()));
if !matches {
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!())
.into());
}
asserted_values.insert(key, exists);
}
}
}
trx.commit().await.map(|_| result).map_err(Into::into)
}
pub(crate) async fn purge_store(&self) -> trc::Result<()> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
for subspace in [SUBSPACE_QUOTA, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER] {
purge_table(&conn, char::from(subspace)).await?;
}
Ok(())
}
pub(crate) async fn delete_range(&self, from: impl Key, to: impl Key) -> trc::Result<()> {
let conn = self.conn_pool.get().await.map_err(into_pool_error)?;
let table = char::from(from.subspace());
let mut from = from.serialize(0);
let to = to.serialize(0);
let delete = conn
.prepare_cached(&format!("DELETE FROM {table} WHERE k >= $1 AND k < $2"))
.await
.map_err(into_error)?;
match conn.execute(&delete, &[&from, &to]).await {
Ok(_) => return Ok(()),
Err(err) if is_timeout_error(&err) => (),
Err(err) => return Err(into_error(err)),
}
let mut chunk_size = DELETE_CHUNK_SIZE;
loop {
let boundary = conn
.prepare_cached(&format!(
"SELECT k FROM {table} WHERE k >= $1 AND k < $2 ORDER BY k ASC LIMIT 1 OFFSET {chunk_size}"
))
.await
.map_err(into_error)?;
loop {
let next = match conn.query_opt(&boundary, &[&from, &to]).await {
Ok(next) => match next {
Some(row) => Some(row.try_get::<_, Vec<u8>>(0).map_err(into_error)?),
None => None,
},
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
};
match conn
.execute(&delete, &[&from, next.as_ref().unwrap_or(&to)])
.await
{
Ok(_) => (),
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
}
match next {
Some(next) => from = next,
None => return Ok(()),
}
}
}
}
}
async fn purge_table(conn: &Object, table: char) -> trc::Result<()> {
let s = conn
.prepare_cached(&format!("DELETE FROM {table} WHERE v = 0"))
.await
.map_err(into_error)?;
match conn.execute(&s, &[]).await {
Ok(_) => return Ok(()),
Err(err) if is_timeout_error(&err) => (),
Err(err) => return Err(into_error(err)),
}
let purge = conn
.prepare_cached(&format!(
"DELETE FROM {table} WHERE v = 0 AND k >= $1 AND k < $2"
))
.await
.map_err(into_error)?;
let purge_last = conn
.prepare_cached(&format!("DELETE FROM {table} WHERE v = 0 AND k >= $1"))
.await
.map_err(into_error)?;
let mut chunk_size = DELETE_CHUNK_SIZE;
let mut from = Vec::new();
loop {
let boundary = conn
.prepare_cached(&format!(
"SELECT k FROM {table} WHERE k >= $1 ORDER BY k ASC LIMIT 1 OFFSET {chunk_size}"
))
.await
.map_err(into_error)?;
loop {
let next = match conn.query_opt(&boundary, &[&from]).await {
Ok(next) => match next {
Some(row) => Some(row.try_get::<_, Vec<u8>>(0).map_err(into_error)?),
None => None,
},
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
};
let result = match &next {
Some(next) => conn.execute(&purge, &[&from, next]).await,
None => conn.execute(&purge_last, &[&from]).await,
};
match result {
Ok(_) => (),
Err(err) if is_timeout_error(&err) && chunk_size > MIN_DELETE_CHUNK_SIZE => {
chunk_size = (chunk_size / 2).max(MIN_DELETE_CHUNK_SIZE);
break;
}
Err(err) => return Err(into_error(err)),
}
match next {
Some(next) => from = next,
None => return Ok(()),
}
}
}
}
impl From<trc::Error> for CommitError {
fn from(err: trc::Error) -> Self {
CommitError::Internal(err)
}
}
impl From<tokio_postgres::Error> for CommitError {
fn from(err: tokio_postgres::Error) -> Self {
CommitError::Postgres(err)
}
}
+301
View File
@@ -0,0 +1,301 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{RedisPool, RedisStore, into_error};
use crate::{Deserialize, write::now};
use redis::AsyncCommands;
impl RedisStore {
pub async fn key_set(&self, key: &[u8], value: &[u8], expires: Option<u64>) -> trc::Result<()> {
match &self.pool {
RedisPool::Single(pool) => {
self.key_set_(
pool.get().await.map_err(into_error)?.as_mut(),
key,
value,
expires,
)
.await
}
RedisPool::Cluster(pool) => {
self.key_set_(
pool.get().await.map_err(into_error)?.as_mut(),
key,
value,
expires,
)
.await
}
RedisPool::Sentinel(pool) => {
self.key_set_(
pool.get().await.map_err(into_error)?.as_mut(),
key,
value,
expires,
)
.await
}
}
}
pub async fn key_incr(&self, key: &[u8], value: i64, expires: Option<u64>) -> trc::Result<i64> {
match &self.pool {
RedisPool::Single(pool) => {
self.key_incr_(
pool.get().await.map_err(into_error)?.as_mut(),
key,
value,
expires,
)
.await
}
RedisPool::Cluster(pool) => {
self.key_incr_(
pool.get().await.map_err(into_error)?.as_mut(),
key,
value,
expires,
)
.await
}
RedisPool::Sentinel(pool) => {
self.key_incr_(
pool.get().await.map_err(into_error)?.as_mut(),
key,
value,
expires,
)
.await
}
}
}
pub async fn try_lock(&self, key: &[u8], expires: u64) -> trc::Result<bool> {
match &self.pool {
RedisPool::Single(pool) => {
self.try_lock_(pool.get().await.map_err(into_error)?.as_mut(), key, expires)
.await
}
RedisPool::Cluster(pool) => {
self.try_lock_(pool.get().await.map_err(into_error)?.as_mut(), key, expires)
.await
}
RedisPool::Sentinel(pool) => {
self.try_lock_(pool.get().await.map_err(into_error)?.as_mut(), key, expires)
.await
}
}
}
pub async fn key_delete(&self, key: &[u8]) -> trc::Result<()> {
match &self.pool {
RedisPool::Single(pool) => {
self.key_delete_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
RedisPool::Cluster(pool) => {
self.key_delete_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
RedisPool::Sentinel(pool) => {
self.key_delete_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
}
}
pub async fn key_delete_prefix(&self, prefix: &[u8]) -> trc::Result<()> {
match &self.pool {
RedisPool::Single(pool) => {
self.key_delete_prefix_(pool.get().await.map_err(into_error)?.as_mut(), prefix)
.await
}
RedisPool::Cluster(pool) => {
self.key_delete_prefix_(pool.get().await.map_err(into_error)?.as_mut(), prefix)
.await
}
RedisPool::Sentinel(pool) => {
self.key_delete_prefix_(pool.get().await.map_err(into_error)?.as_mut(), prefix)
.await
}
}
}
pub async fn key_get<T: Deserialize + std::fmt::Debug + 'static>(
&self,
key: &[u8],
) -> trc::Result<Option<T>> {
match &self.pool {
RedisPool::Single(pool) => {
self.key_get_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
RedisPool::Cluster(pool) => {
self.key_get_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
RedisPool::Sentinel(pool) => {
self.key_get_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
}
}
pub async fn counter_get(&self, key: &[u8]) -> trc::Result<i64> {
match &self.pool {
RedisPool::Single(pool) => {
self.counter_get_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
RedisPool::Cluster(pool) => {
self.counter_get_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
RedisPool::Sentinel(pool) => {
self.counter_get_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
}
}
pub async fn key_exists(&self, key: &[u8]) -> trc::Result<bool> {
match &self.pool {
RedisPool::Single(pool) => {
self.key_exists_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
RedisPool::Cluster(pool) => {
self.key_exists_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
RedisPool::Sentinel(pool) => {
self.key_exists_(pool.get().await.map_err(into_error)?.as_mut(), key)
.await
}
}
}
async fn key_get_<T: Deserialize + std::fmt::Debug + 'static>(
&self,
conn: &mut impl AsyncCommands,
key: &[u8],
) -> trc::Result<Option<T>> {
if let Some(value) = redis::cmd("GET")
.arg(key)
.query_async::<Option<Vec<u8>>>(conn)
.await
.map_err(into_error)?
{
T::deserialize_owned(value).map(Some)
} else {
Ok(None)
}
}
async fn counter_get_(&self, conn: &mut impl AsyncCommands, key: &[u8]) -> trc::Result<i64> {
redis::cmd("GET")
.arg(key)
.query_async::<Option<i64>>(conn)
.await
.map(|x| x.unwrap_or(0))
.map_err(into_error)
}
async fn key_exists_(&self, conn: &mut impl AsyncCommands, key: &[u8]) -> trc::Result<bool> {
conn.exists(key).await.map_err(into_error)
}
async fn key_set_(
&self,
conn: &mut impl AsyncCommands,
key: &[u8],
value: &[u8],
expires: Option<u64>,
) -> trc::Result<()> {
if let Some(expires) = expires {
conn.set_ex(key, value, expires).await.map_err(into_error)
} else {
conn.set(key, value).await.map_err(into_error)
}
}
async fn key_incr_(
&self,
conn: &mut impl AsyncCommands,
key: &[u8],
value: i64,
expires: Option<u64>,
) -> trc::Result<i64> {
if let Some(expires) = expires {
redis::pipe()
.atomic()
.incr(key, value)
.expire(key, expires as i64)
.ignore()
.query_async::<Vec<i64>>(conn)
.await
.map_err(into_error)
.map(|v| v.first().copied().unwrap_or(0))
} else {
conn.incr(key, value).await.map_err(into_error)
}
}
async fn try_lock_(
&self,
conn: &mut impl AsyncCommands,
key: &[u8],
expires: u64,
) -> trc::Result<bool> {
redis::cmd("SET")
.arg(key)
.arg(now() + expires)
.arg("NX")
.arg("EX")
.arg(expires as i64)
.query_async::<Option<String>>(conn)
.await
.map(|reply| reply.is_some())
.map_err(into_error)
}
async fn key_delete_(&self, conn: &mut impl AsyncCommands, key: &[u8]) -> trc::Result<()> {
conn.del(key).await.map_err(into_error)
}
async fn key_delete_prefix_(
&self,
conn: &mut impl AsyncCommands,
prefix: &[u8],
) -> trc::Result<()> {
let mut pattern = Vec::with_capacity(prefix.len() + 1);
pattern.extend_from_slice(prefix);
pattern.push(b'*');
let mut cursor = 0;
loop {
let (new_cursor, keys): (u64, Vec<Vec<u8>>) = redis::cmd("SCAN")
.cursor_arg(cursor)
.arg("MATCH")
.arg(&pattern)
.arg("COUNT")
.arg(100)
.query_async(conn)
.await
.map_err(into_error)?;
if !keys.is_empty() {
conn.del::<_, ()>(&keys).await.map_err(into_error)?;
}
if new_cursor != 0 {
cursor = new_cursor;
} else {
return Ok(());
}
}
}
}
+215
View File
@@ -0,0 +1,215 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::InMemoryStore;
use deadpool::{
Runtime,
managed::{Manager, Pool},
};
use redis::{
Client, ConnectionAddr, IntoConnectionInfo, ProtocolVersion, TlsMode,
cluster::{ClusterClient, ClusterClientBuilder},
cluster_read_routing::RandomReplicaStrategy,
sentinel::{SentinelClient, SentinelClientBuilder, SentinelServerType},
};
use registry::{
schema::{enums::RedisProtocol, structs},
types::duration::Duration,
};
use std::{fmt::Display, sync::Arc};
pub mod lookup;
pub mod pool;
#[derive(Debug)]
pub struct RedisStore {
pub pool: RedisPool,
}
pub struct RedisConnectionManager {
pub client: Client,
timeout: std::time::Duration,
}
pub struct RedisClusterConnectionManager {
pub client: ClusterClient,
timeout: std::time::Duration,
}
pub struct RedisSentinelConnectionManager {
pub client: tokio::sync::Mutex<SentinelClient>,
timeout: std::time::Duration,
}
pub enum RedisPool {
Single(Pool<RedisConnectionManager>),
Cluster(Pool<RedisClusterConnectionManager>),
Sentinel(Pool<RedisSentinelConnectionManager>),
}
impl RedisStore {
pub async fn open_single(config: structs::RedisStore) -> Result<InMemoryStore, String> {
Ok(InMemoryStore::Redis(Arc::new(RedisStore {
pool: RedisPool::Single(build_pool(
RedisConnectionManager {
client: Client::open(config.url)
.map_err(|err| format!("Failed to open Redis client: {err:?}"))?,
timeout: config.timeout.into_inner(),
},
config.pool_max_connections,
config.pool_timeout_create,
config.pool_timeout_wait,
config.pool_timeout_recycle,
)?),
})))
}
pub async fn open_cluster(config: structs::RedisClusterStore) -> Result<InMemoryStore, String> {
let mut builder = ClusterClientBuilder::new(config.urls);
if let Some(value) = config.auth_username {
builder = builder.username(value);
}
if let Some(value) = config.auth_secret.secret().await?.map(|v| v.into_owned()) {
builder = builder.password(value);
}
if let Some(value) = config.max_retries {
builder = builder.retries(value as u32);
}
if let Some(value) = config.max_retry_wait {
builder = builder.max_retry_wait(value.as_millis());
}
if let Some(value) = config.min_retry_wait {
builder = builder.min_retry_wait(value.as_millis());
}
if config.read_from_replicas {
builder = builder.read_routing_strategy(RandomReplicaStrategy);
}
if matches!(config.protocol_version, RedisProtocol::Resp3) {
builder = builder.use_protocol(ProtocolVersion::RESP3);
}
let client = builder
.build()
.map_err(|err| format!("Failed to open Redis client: {err:?}"))?;
Ok(InMemoryStore::Redis(Arc::new(RedisStore {
pool: RedisPool::Cluster(build_pool(
RedisClusterConnectionManager {
client,
timeout: config.timeout.into_inner(),
},
config.pool_max_connections,
config.pool_timeout_create,
config.pool_timeout_wait,
config.pool_timeout_recycle,
)?),
})))
}
pub async fn open_sentinel(
config: structs::RedisSentinelStore,
) -> Result<InMemoryStore, String> {
let mut sentinels = Vec::with_capacity(config.urls.len());
let mut tls_mode = None;
for url in config.urls {
let info = url
.into_connection_info()
.map_err(|err| format!("Invalid Redis Sentinel URL: {err}"))?;
let url_tls_mode = match info.addr() {
ConnectionAddr::TcpTls { insecure: true, .. } => Some(TlsMode::Insecure),
ConnectionAddr::TcpTls {
insecure: false, ..
} => Some(TlsMode::Secure),
_ => None,
};
if sentinels.is_empty() {
tls_mode = url_tls_mode;
} else if tls_mode != url_tls_mode {
return Err(
"All Redis Sentinel URLs must use the same scheme and TLS settings".to_string(),
);
}
sentinels.push(info.addr().clone());
}
let mut builder =
SentinelClientBuilder::new(sentinels, config.service_name, SentinelServerType::Master)
.map_err(|err| format!("Failed to create Redis Sentinel client: {err:?}"))?;
if let Some(value) = config.auth_username {
builder = builder.set_client_to_redis_username(value);
}
if let Some(value) = config.auth_secret.secret().await?.map(|v| v.into_owned()) {
builder = builder.set_client_to_redis_password(value);
}
if let Some(value) = config.sentinel_username {
builder = builder.set_client_to_sentinel_username(value);
}
if let Some(value) = config
.sentinel_secret
.secret()
.await?
.map(|v| v.into_owned())
{
builder = builder.set_client_to_sentinel_password(value);
}
if matches!(config.protocol_version, RedisProtocol::Resp3) {
builder = builder.set_client_to_redis_protocol(ProtocolVersion::RESP3);
}
if let Some(tls_mode) = tls_mode {
builder = builder.set_client_to_redis_tls_mode(tls_mode);
}
let client = builder
.build()
.map_err(|err| format!("Failed to open Redis Sentinel client: {err:?}"))?;
Ok(InMemoryStore::Redis(Arc::new(RedisStore {
pool: RedisPool::Sentinel(build_pool(
RedisSentinelConnectionManager {
client: tokio::sync::Mutex::new(client),
timeout: config.timeout.into_inner(),
},
config.pool_max_connections,
config.pool_timeout_create,
config.pool_timeout_wait,
config.pool_timeout_recycle,
)?),
})))
}
}
fn build_pool<M: Manager>(
manager: M,
max_size: u64,
create_timeout: Option<Duration>,
wait_timeout: Option<Duration>,
recycle_timeout: Option<Duration>,
) -> Result<Pool<M>, String> {
Pool::builder(manager)
.runtime(Runtime::Tokio1)
.max_size(max_size as usize)
.create_timeout(create_timeout.map(|v| v.into_inner()))
.wait_timeout(wait_timeout.map(|v| v.into_inner()))
.recycle_timeout(recycle_timeout.map(|v| v.into_inner()))
.build()
.map_err(|err| format!("Failed to build pool: {err}"))
}
#[inline(always)]
fn into_error(err: impl Display) -> trc::Error {
trc::StoreEvent::RedisError.reason(err)
}
impl std::fmt::Debug for RedisPool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Single(_) => f.debug_tuple("Single").finish(),
Self::Cluster(_) => f.debug_tuple("Cluster").finish(),
Self::Sentinel(_) => f.debug_tuple("Sentinel").finish(),
}
}
}
+87
View File
@@ -0,0 +1,87 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{
RedisClusterConnectionManager, RedisConnectionManager, RedisSentinelConnectionManager,
into_error,
};
use deadpool::managed;
use redis::{
aio::{ConnectionLike, MultiplexedConnection},
cluster_async::ClusterConnection,
};
impl managed::Manager for RedisConnectionManager {
type Type = MultiplexedConnection;
type Error = trc::Error;
async fn create(&self) -> Result<MultiplexedConnection, trc::Error> {
match tokio::time::timeout(self.timeout, self.client.get_multiplexed_async_connection())
.await
{
Ok(conn) => conn.map_err(into_error),
Err(_) => Err(trc::StoreEvent::RedisError.ctx(trc::Key::Details, "Connection Timeout")),
}
}
async fn recycle(
&self,
conn: &mut MultiplexedConnection,
_: &managed::Metrics,
) -> managed::RecycleResult<trc::Error> {
conn.req_packed_command(&redis::cmd("PING"))
.await
.map(|_| ())
.map_err(|err| managed::RecycleError::Backend(into_error(err)))
}
}
impl managed::Manager for RedisClusterConnectionManager {
type Type = ClusterConnection;
type Error = trc::Error;
async fn create(&self) -> Result<ClusterConnection, trc::Error> {
match tokio::time::timeout(self.timeout, self.client.get_async_connection()).await {
Ok(conn) => conn.map_err(into_error),
Err(_) => Err(trc::StoreEvent::RedisError.ctx(trc::Key::Details, "Connection Timeout")),
}
}
async fn recycle(
&self,
conn: &mut ClusterConnection,
_: &managed::Metrics,
) -> managed::RecycleResult<trc::Error> {
conn.req_packed_command(&redis::cmd("PING"))
.await
.map(|_| ())
.map_err(|err| managed::RecycleError::Backend(into_error(err)))
}
}
impl managed::Manager for RedisSentinelConnectionManager {
type Type = MultiplexedConnection;
type Error = trc::Error;
async fn create(&self) -> Result<MultiplexedConnection, trc::Error> {
let mut client = self.client.lock().await;
match tokio::time::timeout(self.timeout, client.get_async_connection()).await {
Ok(conn) => conn.map_err(into_error),
Err(_) => Err(trc::StoreEvent::RedisError.ctx(trc::Key::Details, "Connection Timeout")),
}
}
async fn recycle(
&self,
conn: &mut MultiplexedConnection,
_: &managed::Metrics,
) -> managed::RecycleResult<trc::Error> {
conn.req_packed_command(&redis::cmd("PING"))
.await
.map(|_| ())
.map_err(|err| managed::RecycleError::Backend(into_error(err)))
}
}
+55
View File
@@ -0,0 +1,55 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use std::ops::Range;
use super::{CF_BLOBS, RocksDbStore, into_error};
impl RocksDbStore {
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let db = self.db.clone();
self.spawn_worker(move || {
db.get_pinned_cf(&db.cf_handle(CF_BLOBS).unwrap(), key)
.map(|obj| {
obj.map(|bytes| {
if range.start == 0 && range.end == usize::MAX {
bytes.to_vec()
} else {
bytes
.get(range.start..std::cmp::min(bytes.len(), range.end))
.unwrap_or_default()
.to_vec()
}
})
})
.map_err(into_error)
})
.await
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
let db = self.db.clone();
self.spawn_worker(move || {
db.put_cf(&db.cf_handle(CF_BLOBS).unwrap(), key, data)
.map_err(into_error)
})
.await
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
let db = self.db.clone();
self.spawn_worker(move || {
db.delete_cf(&db.cf_handle(CF_BLOBS).unwrap(), key)
.map_err(into_error)
.map(|_| true)
})
.await
}
}
+230
View File
@@ -0,0 +1,230 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{CF_BLOBS, RocksDbStore};
use crate::*;
use ::registry::schema::structs;
use rocksdb::{
BlockBasedOptions, Cache, ColumnFamilyDescriptor, DBCompressionType, MergeOperands,
OptimisticTransactionDB, Options,
};
use std::path::PathBuf;
use tokio::sync::oneshot;
const MIN_WRITE_BUFFER_SIZE: usize = 4 * 1024 * 1024;
const MAX_WRITE_BUFFER_SIZE: usize = 64 * 1024 * 1024;
const MIN_DB_WRITE_BUFFER_SIZE: usize = 32 * 1024 * 1024;
const BLOOM_BITS_PER_KEY: f64 = 10.0;
const SCAN_BLOCK_SIZE: usize = 16 * 1024;
const CHURN_TARGET_FILE_SIZE: u64 = 16 * 1024 * 1024;
const CHURN_DELETION_WINDOW: usize = 4096;
const CHURN_DELETION_TRIGGER: usize = 1024;
const CHURN_DELETION_RATIO: f64 = 0.5;
const BYTES_PER_SYNC: u64 = 1024 * 1024;
#[derive(Clone, Copy)]
enum CfProfile {
/// Read through `get_value` / `key_exists`, so a whole key bloom filter pays off.
PointLookup,
/// Read only through `iterate`, which never consults a whole key bloom filter.
Scan,
/// Point read and point deleted at a high rate.
Churn,
/// Scanned from the oldest key and point deleted once consumed, with empty values.
Queue,
/// Counters updated through the merge operator.
Counter,
/// Blob values held in RocksDB blob files.
Blob,
}
impl RocksDbStore {
pub async fn open(config: structs::RocksDbStore) -> Result<Store, String> {
// Create the database directory if it doesn't exist
let idx_path: PathBuf = PathBuf::from(config.path);
std::fs::create_dir_all(&idx_path).map_err(|err| {
format!(
"Failed to create database directory {}: {:?}",
idx_path.display(),
err
)
})?;
let cache = Cache::new_lru_cache(config.cache_size as usize);
let write_buffer_size =
((config.buffer_size as usize) / 4).clamp(MIN_WRITE_BUFFER_SIZE, MAX_WRITE_BUFFER_SIZE);
let mut cfs = Vec::new();
// Counters
for subspace in [SUBSPACE_COUNTER, SUBSPACE_QUOTA, SUBSPACE_IN_MEMORY_COUNTER] {
cfs.push(ColumnFamilyDescriptor::new(
std::str::from_utf8(&[subspace]).unwrap(),
cf_options(CfProfile::Counter, &cache, write_buffer_size),
));
}
// Blobs
let mut cf_opts = cf_options(CfProfile::Blob, &cache, write_buffer_size);
cf_opts.set_enable_blob_files(true);
cf_opts.set_min_blob_size(config.blob_size);
cf_opts.set_enable_blob_gc(true);
cf_opts.set_blob_gc_age_cutoff(1.0);
cf_opts.set_blob_gc_force_threshold(0.5);
cfs.push(ColumnFamilyDescriptor::new(CF_BLOBS, cf_opts));
// Other cfs
for (subspace, profile) in [
(SUBSPACE_INDEXES, CfProfile::Scan),
(SUBSPACE_ACL, CfProfile::Scan),
(SUBSPACE_TASK_QUEUE, CfProfile::Churn),
(SUBSPACE_DELETED_ITEMS, CfProfile::Churn),
(SUBSPACE_BLOB_LINK, CfProfile::Churn),
(SUBSPACE_IN_MEMORY_VALUE, CfProfile::Churn),
(SUBSPACE_PROPERTY, CfProfile::PointLookup),
(SUBSPACE_REGISTRY, CfProfile::PointLookup),
(SUBSPACE_QUEUE_MESSAGE, CfProfile::Churn),
(SUBSPACE_QUEUE_EVENT, CfProfile::Queue),
(SUBSPACE_REPORT_OUT, CfProfile::Churn),
(SUBSPACE_REPORT_IN, CfProfile::Churn),
(SUBSPACE_LOGS, CfProfile::Scan),
(SUBSPACE_TELEMETRY_SPAN, CfProfile::PointLookup),
(SUBSPACE_TELEMETRY_METRIC, CfProfile::Scan),
(SUBSPACE_SEARCH_INDEX, CfProfile::Scan),
(SUBSPACE_SPAM_SAMPLES, CfProfile::Churn),
(SUBSPACE_REGISTRY_IDX, CfProfile::Scan),
(SUBSPACE_REGISTRY_PK, CfProfile::PointLookup),
(SUBSPACE_DIRECTORY, CfProfile::PointLookup),
(LEGACY_SUBSPACE_BITMAP_TEXT, CfProfile::Scan),
(LEGACY_SUBSPACE_BITMAP_TAG, CfProfile::Scan),
] {
cfs.push(ColumnFamilyDescriptor::new(
std::str::from_utf8(&[subspace]).unwrap(),
cf_options(profile, &cache, write_buffer_size),
));
}
let mut db_opts = Options::default();
db_opts.create_missing_column_families(true);
db_opts.create_if_missing(true);
db_opts.set_max_background_jobs(std::cmp::max(num_cpus::get() as i32, 3));
db_opts.increase_parallelism(std::cmp::max(num_cpus::get() as i32, 3));
db_opts
.set_db_write_buffer_size((config.buffer_size as usize).max(MIN_DB_WRITE_BUFFER_SIZE));
db_opts.set_bytes_per_sync(BYTES_PER_SYNC);
db_opts.set_wal_bytes_per_sync(BYTES_PER_SYNC);
Ok(Store::RocksDb(Arc::new(RocksDbStore {
db: OptimisticTransactionDB::open_cf_descriptors(&db_opts, idx_path, cfs)
.map_err(|err| format!("Failed to open database: {:?}", err))?
.into(),
worker_pool: rayon::ThreadPoolBuilder::new()
.num_threads(std::cmp::max(
config
.pool_workers
.filter(|v| *v > 0)
.map(|v| v as usize)
.unwrap_or_else(num_cpus::get),
4,
))
.build()
.map_err(|err| format!("Failed to build worker pool: {:?}", err))?,
})))
}
pub async fn spawn_worker<U, V>(&self, mut f: U) -> trc::Result<V>
where
U: FnMut() -> trc::Result<V> + Send,
V: Sync + Send + 'static,
{
let (tx, rx) = oneshot::channel();
self.worker_pool.scope(|s| {
s.spawn(|_| {
tx.send(f()).ok();
});
});
match rx.await {
Ok(result) => result,
Err(err) => Err(trc::EventType::Server(trc::ServerEvent::ThreadError).reason(err)),
}
}
}
pub fn numeric_value_merge(
_key: &[u8],
value: Option<&[u8]>,
operands: &MergeOperands,
) -> Option<Vec<u8>> {
let mut value = if let Some(value) = value {
i64::from_le_bytes(value.try_into().ok()?)
} else {
0
};
for op in operands.iter() {
value += i64::from_le_bytes(op.try_into().ok()?);
}
let mut bytes = Vec::with_capacity(std::mem::size_of::<i64>());
bytes.extend_from_slice(&value.to_le_bytes());
Some(bytes)
}
fn cf_options(profile: CfProfile, cache: &Cache, write_buffer_size: usize) -> Options {
let mut block_opts = BlockBasedOptions::default();
block_opts.set_block_cache(cache);
block_opts.set_cache_index_and_filter_blocks(true);
block_opts.set_pin_l0_filter_and_index_blocks_in_cache(true);
let mut opts = Options::default();
opts.set_write_buffer_size(write_buffer_size);
opts.set_max_write_buffer_number(4);
match profile {
CfProfile::PointLookup => {
block_opts.set_bloom_filter(BLOOM_BITS_PER_KEY, false);
opts.set_compression_type(DBCompressionType::Lz4);
}
CfProfile::Scan => {
block_opts.set_block_size(SCAN_BLOCK_SIZE);
opts.set_compression_type(DBCompressionType::Lz4);
}
CfProfile::Churn => {
block_opts.set_bloom_filter(BLOOM_BITS_PER_KEY, false);
opts.set_compression_type(DBCompressionType::Lz4);
opts.set_target_file_size_base(CHURN_TARGET_FILE_SIZE);
opts.add_compact_on_deletion_collector_factory(
CHURN_DELETION_WINDOW,
CHURN_DELETION_TRIGGER,
CHURN_DELETION_RATIO,
);
}
CfProfile::Queue => {
block_opts.set_block_size(SCAN_BLOCK_SIZE);
opts.set_compression_type(DBCompressionType::None);
opts.set_target_file_size_base(CHURN_TARGET_FILE_SIZE);
opts.add_compact_on_deletion_collector_factory(
CHURN_DELETION_WINDOW,
CHURN_DELETION_TRIGGER,
CHURN_DELETION_RATIO,
);
}
CfProfile::Counter => {
block_opts.set_bloom_filter(BLOOM_BITS_PER_KEY, false);
opts.set_compression_type(DBCompressionType::None);
opts.set_merge_operator_associative("merge", numeric_value_merge);
}
CfProfile::Blob => {
block_opts.set_bloom_filter(BLOOM_BITS_PER_KEY, false);
opts.set_compression_type(DBCompressionType::None);
}
}
opts.set_block_based_table_factory(&block_opts);
opts
}
+43
View File
@@ -0,0 +1,43 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use std::sync::Arc;
use rocksdb::{BoundColumnFamily, MultiThreaded, OptimisticTransactionDB};
use crate::{SUBSPACE_BLOBS, SUBSPACE_INDEXES, SUBSPACE_LOGS};
pub mod blob;
pub mod main;
pub mod read;
pub mod write;
static CF_LOGS: &str = unsafe { std::str::from_utf8_unchecked(&[SUBSPACE_LOGS]) };
static CF_INDEXES: &str = unsafe { std::str::from_utf8_unchecked(&[SUBSPACE_INDEXES]) };
static CF_BLOBS: &str = unsafe { std::str::from_utf8_unchecked(&[SUBSPACE_BLOBS]) };
pub(crate) trait CfHandle {
fn subspace_handle(&self, subspace: u8) -> Arc<BoundColumnFamily<'_>>;
}
impl CfHandle for OptimisticTransactionDB<MultiThreaded> {
#[inline(always)]
fn subspace_handle(&self, subspace: u8) -> Arc<BoundColumnFamily<'_>> {
let subspace = &[subspace];
self.cf_handle(unsafe { std::str::from_utf8_unchecked(subspace) })
.unwrap()
}
}
pub struct RocksDbStore {
db: Arc<OptimisticTransactionDB<MultiThreaded>>,
worker_pool: rayon::ThreadPool,
}
#[inline(always)]
fn into_error(err: rocksdb::Error) -> trc::Error {
trc::StoreEvent::RocksdbError.reason(err)
}
+132
View File
@@ -0,0 +1,132 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{RocksDbStore, into_error};
use crate::{
Deserialize, IterateParams, Key, ValueKey, backend::rocksdb::CfHandle, write::ValueClass,
};
use rocksdb::ReadOptions;
impl RocksDbStore {
pub(crate) async fn get_value<U>(&self, key: impl Key) -> trc::Result<Option<U>>
where
U: Deserialize + 'static,
{
let db = self.db.clone();
self.spawn_worker(move || {
let subspace = &[key.subspace()];
let key = key.serialize(0);
db.get_pinned_cf(
&db.cf_handle(unsafe { std::str::from_utf8_unchecked(subspace.as_slice()) })
.unwrap(),
&key,
)
.map_err(into_error)
.and_then(|value| {
if let Some(value) = value {
U::deserialize_with_key(&key, &value).map(Some)
} else {
Ok(None)
}
})
})
.await
}
pub(crate) async fn key_exists(&self, key: impl Key) -> trc::Result<bool> {
let db = self.db.clone();
self.spawn_worker(move || {
let subspace = &[key.subspace()];
let key = key.serialize(0);
db.get_pinned_cf(
&db.cf_handle(unsafe { std::str::from_utf8_unchecked(subspace.as_slice()) })
.unwrap(),
&key,
)
.map_err(into_error)
.map(|value| value.is_some())
})
.await
}
pub(crate) async fn iterate<T: Key>(
&self,
params: IterateParams<T>,
mut cb: impl for<'x> FnMut(&'x [u8], &'x [u8]) -> trc::Result<bool> + Sync + Send,
) -> trc::Result<()> {
let db = self.db.clone();
self.spawn_worker(move || {
let cf = db.subspace_handle(params.begin.subspace());
let begin = params.begin.serialize(0);
let end = params.end.serialize(0);
let mut upper_bound = Vec::with_capacity(end.len() + 1);
upper_bound.extend_from_slice(&end);
upper_bound.push(0u8);
let mut read_opts = ReadOptions::default();
read_opts.set_iterate_lower_bound(begin.as_slice());
read_opts.set_iterate_upper_bound(upper_bound);
let mut it = db.raw_iterator_cf_opt(&cf, read_opts);
if params.ascending {
it.seek(&begin);
} else {
it.seek_for_prev(&end);
}
while it.valid() {
let Some(key) = it.key() else {
break;
};
let value = if params.values {
it.value().unwrap_or_default()
} else {
&[][..]
};
if !cb(key, value)? || params.first {
return Ok(());
}
if params.ascending {
it.next();
} else {
it.prev();
}
}
it.status().map_err(into_error)
})
.await
}
pub(crate) async fn get_counter(
&self,
key: impl Into<ValueKey<ValueClass>> + Sync + Send,
) -> trc::Result<i64> {
let key = key.into();
let db = self.db.clone();
self.spawn_worker(move || {
let cf = self.db.subspace_handle(key.subspace());
let key = key.serialize(0);
db.get_pinned_cf(&cf, &key)
.map_err(into_error)
.and_then(|bytes| {
Ok(if let Some(bytes) = bytes {
i64::from_le_bytes(bytes[..].try_into().map_err(|_| {
trc::Error::corrupted_key(&key, (&bytes[..]).into(), trc::location!())
})?)
} else {
0
})
})
})
.await
}
}
+306
View File
@@ -0,0 +1,306 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{CF_INDEXES, CF_LOGS, CfHandle, RocksDbStore, into_error};
use crate::{
Deserialize, IndexKey, Key, LogKey, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER,
SUBSPACE_QUOTA,
backend::deserialize_i64_le,
write::{
AssignedIds, Batch, MAX_COMMIT_ATTEMPTS, MAX_COMMIT_TIME, MergeResult, Operation,
ValueClass, ValueOp,
},
};
use rand::RngExt;
use rocksdb::{
BoundColumnFamily, ErrorKind, IteratorMode, OptimisticTransactionDB,
OptimisticTransactionOptions, WriteOptions,
};
use std::{
sync::Arc,
thread::sleep,
time::{Duration, Instant},
};
impl RocksDbStore {
pub(crate) async fn write(&self, mut batch: Batch<'_>) -> trc::Result<AssignedIds> {
let db = self.db.clone();
self.spawn_worker(move || {
let mut txn = RocksDBTransaction {
db: &db,
cf_indexes: db.cf_handle(CF_INDEXES).unwrap(),
cf_logs: db.cf_handle(CF_LOGS).unwrap(),
txn_opts: OptimisticTransactionOptions::default(),
batch: &mut batch,
};
txn.txn_opts.set_snapshot(true);
// Begin write
let mut retry_count = 0;
let start = Instant::now();
loop {
match txn.commit() {
Ok(result) => {
return Ok(result);
}
Err(CommitError::Internal(err)) => return Err(err),
Err(CommitError::RocksDB(err)) => match err.kind() {
ErrorKind::Busy | ErrorKind::MergeInProgress | ErrorKind::TryAgain
if retry_count < MAX_COMMIT_ATTEMPTS
&& start.elapsed() < MAX_COMMIT_TIME =>
{
let backoff = rand::rng().random_range(50..=300);
sleep(Duration::from_millis(backoff));
retry_count += 1;
}
_ => return Err(into_error(err)),
},
}
}
})
.await
}
pub(crate) async fn delete_range(&self, from: impl Key, to: impl Key) -> trc::Result<()> {
let db = self.db.clone();
self.spawn_worker(move || {
db.delete_range_cf(
&db.cf_handle(std::str::from_utf8(&[from.subspace()]).unwrap())
.unwrap(),
from.serialize(0),
to.serialize(0),
)
.map_err(into_error)
})
.await
}
pub(crate) async fn purge_store(&self) -> trc::Result<()> {
let db = self.db.clone();
self.spawn_worker(move || {
for subspace in [SUBSPACE_QUOTA, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER] {
let cf = db
.cf_handle(std::str::from_utf8(&[subspace]).unwrap())
.unwrap();
let mut delete_keys = Vec::new();
for row in db.iterator_cf(&cf, IteratorMode::Start) {
let (key, value) = row.map_err(into_error)?;
if i64::deserialize(&value)? == 0 {
delete_keys.push(key);
}
}
let txn_opts = OptimisticTransactionOptions::default();
for key in delete_keys {
let txn = db.transaction_opt(&WriteOptions::default(), &txn_opts);
if txn
.get_pinned_for_update_cf(&cf, &key, true)
.map_err(into_error)?
.map(|value| i64::deserialize(&value).map(|v| v == 0).unwrap_or(false))
.unwrap_or(false)
{
txn.delete_cf(&cf, key).map_err(into_error)?;
txn.commit().map_err(into_error)?;
} else {
txn.rollback().map_err(into_error)?;
}
}
}
Ok(())
})
.await
}
}
struct RocksDBTransaction<'x, 'y> {
db: &'x OptimisticTransactionDB,
cf_indexes: Arc<BoundColumnFamily<'x>>,
cf_logs: Arc<BoundColumnFamily<'x>>,
txn_opts: OptimisticTransactionOptions,
batch: &'x mut Batch<'y>,
}
enum CommitError {
Internal(trc::Error),
RocksDB(rocksdb::Error),
}
impl RocksDBTransaction<'_, '_> {
fn commit(&mut self) -> Result<AssignedIds, CommitError> {
let mut account_id = u32::MAX;
let mut collection = u8::MAX;
let mut document_id = u32::MAX;
let mut change_id = 0u64;
let mut result = AssignedIds::default();
let has_changes = !self.batch.changes.is_empty();
let txn = self
.db
.transaction_opt(&WriteOptions::default(), &self.txn_opts);
if has_changes {
let cf = self.db.cf_handle("n").unwrap();
for &account_id in self.batch.changes.keys() {
let key = ValueClass::ChangeId.serialize(account_id, 0, 0, 0);
let change_id = txn
.get_pinned_for_update_cf(&cf, &key, true)
.map_err(CommitError::from)
.and_then(|bytes| {
if let Some(bytes) = bytes {
deserialize_i64_le(&key, &bytes)
.map(|v| v + 1)
.map_err(CommitError::from)
} else {
Ok(1)
}
})?;
txn.put_cf(&cf, &key, &change_id.to_le_bytes()[..])?;
result.push_change_id(account_id, change_id as u64);
}
}
for op in self.batch.ops.iter_mut() {
match op {
Operation::AccountId {
account_id: account_id_,
} => {
account_id = *account_id_;
if has_changes {
change_id = result.set_current_change_id(account_id)?;
}
}
Operation::Collection {
collection: collection_,
} => {
collection = u8::from(*collection_);
}
Operation::DocumentId {
document_id: document_id_,
} => {
document_id = *document_id_;
}
Operation::Value { class, op } => {
let key = class.serialize(account_id, collection, document_id, 0);
let cf = self.db.subspace_handle(class.subspace(collection));
match op {
ValueOp::Set(value) => {
txn.put_cf(&cf, &key, value)?;
}
ValueOp::SetFnc(set_op) => {
let value = (set_op.fnc)(&set_op.params, &result)?;
txn.put_cf(&cf, &key, value)?;
}
ValueOp::MergeFnc(merge_op) => {
let merge_result = (merge_op.fnc)(
&merge_op.params,
&result,
txn.get_pinned_for_update_cf(&cf, &key, true)?.as_deref(),
)?;
match merge_result {
MergeResult::Update(value) => {
txn.put_cf(&cf, &key, value)?;
}
MergeResult::Delete => {
txn.delete_cf(&cf, &key)?;
}
MergeResult::Skip => (),
}
}
ValueOp::AtomicAdd(by) => {
txn.merge_cf(&cf, &key, &by.to_le_bytes()[..])?;
}
ValueOp::AddAndGet(by) => {
let num = txn
.get_pinned_for_update_cf(&cf, &key, true)
.map_err(CommitError::from)
.and_then(|bytes| {
if let Some(bytes) = bytes {
deserialize_i64_le(&key, &bytes)
.map(|v| v + *by)
.map_err(CommitError::from)
} else {
Ok(*by)
}
})?;
txn.put_cf(&cf, &key, &num.to_le_bytes()[..])?;
result.push_counter_id(num);
}
ValueOp::Clear => {
txn.delete_cf(&cf, &key)?;
}
}
}
Operation::Index { field, key, set } => {
let key = IndexKey {
account_id,
collection,
document_id,
field: *field,
key: &*key,
}
.serialize(0);
if *set {
txn.put_cf(&self.cf_indexes, &key, [])?;
} else {
txn.delete_cf(&self.cf_indexes, &key)?;
}
}
Operation::Log { collection, set } => {
let key = LogKey {
account_id,
collection: u8::from(*collection),
change_id,
}
.serialize(0);
txn.put_cf(&self.cf_logs, &key, set)?;
}
Operation::AssertValue {
class,
assert_value,
} => {
let key = class.serialize(account_id, collection, document_id, 0);
let cf = self.db.subspace_handle(class.subspace(collection));
let matches = txn
.get_pinned_for_update_cf(&cf, &key, true)?
.map(|value| assert_value.matches(&value))
.unwrap_or_else(|| assert_value.is_none());
if !matches {
txn.rollback()?;
return Err(CommitError::Internal(
trc::StoreEvent::AssertValueFailed.into(),
));
}
}
}
}
txn.commit().map(|_| result).map_err(Into::into)
}
}
impl From<rocksdb::Error> for CommitError {
fn from(err: rocksdb::Error) -> Self {
CommitError::RocksDB(err)
}
}
impl From<trc::Error> for CommitError {
fn from(err: trc::Error) -> Self {
CommitError::Internal(err)
}
}
+259
View File
@@ -0,0 +1,259 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use crate::BlobStore;
use registry::schema::structs;
use s3::{Bucket, Region, creds::Credentials};
use std::{io::Write, ops::Range, sync::Arc, time::Duration};
use utils::codec::base32_custom::Base32Writer;
pub struct S3Store {
bucket: Box<Bucket>,
prefix: Option<String>,
max_retries: u32,
verify_after_write: bool,
}
impl S3Store {
pub async fn open(config: structs::S3Store) -> Result<BlobStore, String> {
// Obtain region and endpoint from config
let region = match config.region {
structs::S3StoreRegion::UsEast1 => Region::UsEast1,
structs::S3StoreRegion::UsEast2 => Region::UsEast2,
structs::S3StoreRegion::UsWest1 => Region::UsWest1,
structs::S3StoreRegion::UsWest2 => Region::UsWest2,
structs::S3StoreRegion::CaCentral1 => Region::CaCentral1,
structs::S3StoreRegion::AfSouth1 => Region::Custom {
region: "af-south-1".into(),
endpoint: "s3.af-south-1.amazonaws.com".into(),
},
structs::S3StoreRegion::ApEast1 => Region::ApEast1,
structs::S3StoreRegion::ApSouth1 => Region::ApSouth1,
structs::S3StoreRegion::ApNortheast1 => Region::ApNortheast1,
structs::S3StoreRegion::ApNortheast2 => Region::ApNortheast2,
structs::S3StoreRegion::ApNortheast3 => Region::ApNortheast3,
structs::S3StoreRegion::ApSoutheast1 => Region::ApSoutheast1,
structs::S3StoreRegion::ApSoutheast2 => Region::ApSoutheast2,
structs::S3StoreRegion::CnNorth1 => Region::CnNorth1,
structs::S3StoreRegion::CnNorthwest1 => Region::CnNorthwest1,
structs::S3StoreRegion::EuNorth1 => Region::EuNorth1,
structs::S3StoreRegion::EuCentral1 => Region::EuCentral1,
structs::S3StoreRegion::EuCentral2 => Region::EuCentral2,
structs::S3StoreRegion::EuWest1 => Region::EuWest1,
structs::S3StoreRegion::EuWest2 => Region::EuWest2,
structs::S3StoreRegion::EuWest3 => Region::EuWest3,
structs::S3StoreRegion::IlCentral1 => Region::IlCentral1,
structs::S3StoreRegion::MeSouth1 => Region::MeSouth1,
structs::S3StoreRegion::SaEast1 => Region::SaEast1,
structs::S3StoreRegion::DoNyc3 => Region::DoNyc3,
structs::S3StoreRegion::DoAms3 => Region::DoAms3,
structs::S3StoreRegion::DoSgp1 => Region::DoSgp1,
structs::S3StoreRegion::DoFra1 => Region::DoFra1,
structs::S3StoreRegion::Yandex => Region::Yandex,
structs::S3StoreRegion::WaUsEast1 => Region::WaUsEast1,
structs::S3StoreRegion::WaUsEast2 => Region::WaUsEast2,
structs::S3StoreRegion::WaUsCentral1 => Region::WaUsCentral1,
structs::S3StoreRegion::WaUsWest1 => Region::WaUsWest1,
structs::S3StoreRegion::WaCaCentral1 => Region::WaCaCentral1,
structs::S3StoreRegion::WaEuCentral1 => Region::WaEuCentral1,
structs::S3StoreRegion::WaEuCentral2 => Region::WaEuCentral2,
structs::S3StoreRegion::WaEuWest1 => Region::WaEuWest1,
structs::S3StoreRegion::WaEuWest2 => Region::WaEuWest2,
structs::S3StoreRegion::WaApNortheast1 => Region::WaApNortheast1,
structs::S3StoreRegion::WaApNortheast2 => Region::WaApNortheast2,
structs::S3StoreRegion::WaApSoutheast1 => Region::WaApSoutheast1,
structs::S3StoreRegion::WaApSoutheast2 => Region::WaApSoutheast2,
structs::S3StoreRegion::Custom(custom) => Region::Custom {
region: custom.custom_region,
endpoint: custom.custom_endpoint,
},
};
let credentials = Credentials::new(
config.access_key.value().await?.as_deref(),
config.secret_key.secret().await?.as_deref(),
config.security_token.secret().await?.as_deref(),
config.session_token.secret().await?.as_deref(),
config.profile.as_deref(),
)
.map_err(|err| format!("Failed to create credentials: {err:?}"))?;
Ok(BlobStore::S3(Arc::new(S3Store {
bucket: Bucket::new(&config.bucket, region, credentials)
.map_err(|err| format!("Failed to create bucket: {err:?}"))?
.with_path_style()
.set_dangerous_config(config.allow_invalid_certs, config.allow_invalid_certs)
.map_err(|err| format!("Failed to create bucket: {err:?}"))?
.with_request_timeout(config.timeout.into_inner())
.map_err(|err| format!("Failed to create bucket: {err:?}"))?,
max_retries: config.max_retries as u32,
prefix: config.key_prefix,
verify_after_write: config.verify_after_write,
})))
}
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let path = self.build_key(key);
let mut retries_left = self.max_retries;
loop {
let response = if range.start != 0 || range.end != usize::MAX {
self.bucket
.get_object_range(
&path,
range.start as u64,
Some(range.end.saturating_sub(1) as u64),
)
.await
} else {
self.bucket.get_object(&path).await
}
.map_err(into_error)?;
match response.status_code() {
200..=299 => return Ok(Some(response.to_vec())),
404 => return Ok(None),
500..=599 if retries_left > 0 => {
// wait backoff
tokio::time::sleep(Duration::from_secs(
1 << (self.max_retries - retries_left).min(6),
))
.await;
retries_left -= 1;
}
code => {
return Err(trc::StoreEvent::S3Error
.reason(String::from_utf8_lossy(response.as_slice()))
.ctx(trc::Key::Code, code));
}
}
}
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
let path = self.build_key(key);
let mut retries_left = self.max_retries;
loop {
let response = self
.bucket
.put_object(&path, data)
.await
.map_err(into_error)?;
match response.status_code() {
200..=299 => {
if !self.verify_after_write {
return Ok(());
}
// Some S3-compatible backends acknowledge a PUT before the
// write is durable. HEAD the object to confirm it is visible
// to the read path before reporting success.
let (_, head_status) =
self.bucket.head_object(&path).await.map_err(into_error)?;
match head_status {
200..=299 => return Ok(()),
404 | 500..=599 if retries_left > 0 => {
tokio::time::sleep(Duration::from_secs(
1 << (self.max_retries - retries_left).min(6),
))
.await;
retries_left -= 1;
}
404 => {
return Err(trc::StoreEvent::S3Error
.reason(concat!(
"PUT acknowledged with 2xx but object not visible",
"to read path; backend may be silently losing writes"
))
.ctx(trc::Key::Code, head_status));
}
code => {
return Err(trc::StoreEvent::S3Error
.reason("HEAD verification failed after PUT")
.ctx(trc::Key::Code, code));
}
}
}
500..=599 if retries_left > 0 => {
// wait backoff
tokio::time::sleep(Duration::from_secs(
1 << (self.max_retries - retries_left).min(6),
))
.await;
retries_left -= 1;
}
code => {
return Err(trc::StoreEvent::S3Error
.reason(String::from_utf8_lossy(response.as_slice()))
.ctx(trc::Key::Code, code));
}
}
}
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
let mut retries_left = self.max_retries;
loop {
let response = self
.bucket
.delete_object(self.build_key(key))
.await
.map_err(into_error)?;
match response.status_code() {
200..=299 => return Ok(true),
404 => return Ok(false),
500..=599 if retries_left > 0 => {
// wait backoff
tokio::time::sleep(Duration::from_secs(
1 << (self.max_retries - retries_left).min(6),
))
.await;
retries_left -= 1;
}
code => {
return Err(trc::StoreEvent::S3Error
.reason(String::from_utf8_lossy(response.as_slice()))
.ctx(trc::Key::Code, code));
}
}
}
}
fn build_key(&self, key: &[u8]) -> String {
if let Some(prefix) = &self.prefix {
let mut writer =
Base32Writer::with_raw_capacity(prefix.len() + (key.len().div_ceil(4) * 5));
writer.push_string(prefix);
writer.write_all(key).unwrap();
writer.finalize()
} else {
Base32Writer::from_bytes(key).finalize()
}
}
}
fn into_error(err: impl std::error::Error) -> trc::Error {
let mut reason = err.to_string();
let mut source = err.source();
while let Some(err) = source {
reason.push_str(": ");
reason.push_str(&err.to_string());
source = err.source();
}
trc::StoreEvent::S3Error.reason(reason)
}
+70
View File
@@ -0,0 +1,70 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use std::ops::Range;
use rusqlite::OptionalExtension;
use super::{SqliteStore, into_error};
impl SqliteStore {
pub(crate) async fn get_blob(
&self,
key: &[u8],
range: Range<usize>,
) -> trc::Result<Option<Vec<u8>>> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
let mut result = conn
.prepare_cached("SELECT v FROM t WHERE k = ?")
.map_err(into_error)?;
result
.query_row([&key], |row| {
Ok({
let bytes = row.get_ref(0)?.as_bytes()?;
if range.start == 0 && range.end == usize::MAX {
bytes.to_vec()
} else {
bytes
.get(range.start..std::cmp::min(bytes.len(), range.end))
.unwrap_or_default()
.to_vec()
}
})
})
.optional()
.map_err(into_error)
})
.await
}
pub(crate) async fn put_blob(&self, key: &[u8], data: &[u8]) -> trc::Result<()> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
conn.prepare_cached("INSERT OR REPLACE INTO t (k, v) VALUES (?, ?)")
.map_err(into_error)?
.execute([key, data])
.map_err(into_error)
.map(|_| ())
})
.await
}
pub(crate) async fn delete_blob(&self, key: &[u8]) -> trc::Result<bool> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
conn.prepare_cached("DELETE FROM t WHERE k = ?")
.map_err(into_error)?
.execute([key])
.map_err(into_error)
.map(|_| true)
})
.await
}
}
+145
View File
@@ -0,0 +1,145 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use rusqlite::{Row, Rows, ToSql, types::FromSql};
use crate::{IntoRows, QueryResult, QueryType, Value};
use super::{SqliteStore, into_error};
impl SqliteStore {
pub(crate) async fn sql_query<T: QueryResult>(
&self,
query: &str,
params_: &[Value<'_>],
) -> trc::Result<T> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
let mut s = conn.prepare_cached(query).map_err(into_error)?;
let params = params_
.iter()
.map(|v| v as &dyn rusqlite::types::ToSql)
.collect::<Vec<_>>();
match T::query_type() {
QueryType::Execute => s
.execute(params.as_slice())
.map_or_else(|e| Err(into_error(e)), |r| Ok(T::from_exec(r))),
QueryType::Exists => s
.exists(params.as_slice())
.map(T::from_exists)
.map_err(into_error),
QueryType::QueryOne => s
.query(params.as_slice())
.and_then(|mut rows| Ok(T::from_query_one(rows.next()?)))
.map_err(into_error),
QueryType::QueryAll => Ok(T::from_query_all(
s.query(params.as_slice()).map_err(into_error)?,
)),
}
})
.await
}
}
impl ToSql for Value<'_> {
fn to_sql(&self) -> rusqlite::Result<rusqlite::types::ToSqlOutput<'_>> {
match self {
Value::Integer(value) => value.to_sql(),
Value::Bool(value) => value.to_sql(),
Value::Float(value) => value.to_sql(),
Value::Text(value) => value.to_sql(),
Value::Blob(value) => value.to_sql(),
Value::Null => Ok(rusqlite::types::ToSqlOutput::Owned(
rusqlite::types::Value::Null,
)),
}
}
}
impl FromSql for Value<'static> {
fn column_result(value: rusqlite::types::ValueRef<'_>) -> rusqlite::types::FromSqlResult<Self> {
Ok(match value {
rusqlite::types::ValueRef::Null => Value::Null,
rusqlite::types::ValueRef::Integer(v) => Value::Integer(v),
rusqlite::types::ValueRef::Real(v) => Value::Float(v),
rusqlite::types::ValueRef::Text(v) => {
Value::Text(String::from_utf8_lossy(v).into_owned().into())
}
rusqlite::types::ValueRef::Blob(v) => Value::Blob(v.to_vec().into()),
})
}
}
impl IntoRows for Rows<'_> {
fn into_rows(mut self) -> crate::Rows {
let column_count = self.as_ref().map(|s| s.column_count()).unwrap_or_default();
let mut rows = crate::Rows { rows: Vec::new() };
while let Ok(Some(row)) = self.next() {
rows.rows.push(crate::Row {
values: (0..column_count)
.map(|idx| row.get::<_, Value>(idx).unwrap_or(Value::Null))
.collect(),
});
}
rows
}
fn into_named_rows(mut self) -> crate::NamedRows {
let (column_count, names) = self
.as_ref()
.map(|s| {
(
s.column_count(),
s.column_names()
.into_iter()
.map(String::from)
.collect::<Vec<_>>(),
)
})
.unwrap_or((0, Vec::new()));
let mut rows = crate::NamedRows {
names,
rows: Vec::new(),
};
while let Ok(Some(row)) = self.next() {
rows.rows.push(crate::Row {
values: (0..column_count)
.map(|idx| row.get::<_, Value>(idx).unwrap_or(Value::Null))
.collect(),
});
}
rows
}
fn into_row(self) -> Option<crate::Row> {
unreachable!()
}
}
impl IntoRows for Option<&Row<'_>> {
fn into_row(self) -> Option<crate::Row> {
self.map(|row| crate::Row {
values: (0..row.as_ref().column_count())
.map(|idx| row.get::<_, Value>(idx).unwrap_or(Value::Null))
.collect(),
})
}
fn into_rows(self) -> crate::Rows {
unreachable!()
}
fn into_named_rows(self) -> crate::NamedRows {
unreachable!()
}
}
+146
View File
@@ -0,0 +1,146 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{SqliteStore, into_error, pool::SqliteConnectionManager};
use crate::*;
use ::registry::schema::structs;
use r2d2::Pool;
use tokio::sync::oneshot;
impl SqliteStore {
pub fn open(config: structs::SqliteStore) -> Result<Store, String> {
Ok(Store::SQLite(Arc::new(SqliteStore {
conn_pool: Pool::builder()
.max_size(config.pool_max_connections as u32)
.build(SqliteConnectionManager::file(&config.path).with_init(|c| {
c.execute_batch(concat!(
"PRAGMA journal_mode = WAL; ",
"PRAGMA synchronous = NORMAL; ",
"PRAGMA temp_store = memory;",
"PRAGMA busy_timeout = 30000;"
))
}))
.map_err(|err| format!("Failed to build connection pool: {err}"))?,
worker_pool: rayon::ThreadPoolBuilder::new()
.num_threads(std::cmp::max(
config
.pool_workers
.filter(|v| *v > 0)
.map(|v| v as usize)
.unwrap_or_else(num_cpus::get),
4,
))
.build()
.map_err(|err| format!("Failed to build worker pool: {err}"))?,
})))
}
#[cfg(feature = "test_mode")]
pub fn open_memory() -> trc::Result<Self> {
use super::into_error;
let db = Self {
conn_pool: Pool::builder()
.max_size(1)
.build(SqliteConnectionManager::memory())
.map_err(into_error)?,
worker_pool: rayon::ThreadPoolBuilder::new()
.num_threads(num_cpus::get())
.build()
.map_err(|err| {
into_error(err).ctx(trc::Key::Reason, "Failed to build worker pool")
})?,
};
db.create_tables()?;
Ok(db)
}
pub(crate) fn create_tables(&self) -> trc::Result<()> {
let conn = self.conn_pool.get().map_err(into_error)?;
for table in [
SUBSPACE_ACL,
SUBSPACE_TASK_QUEUE,
SUBSPACE_DELETED_ITEMS,
SUBSPACE_SPAM_SAMPLES,
SUBSPACE_BLOB_LINK,
SUBSPACE_IN_MEMORY_VALUE,
SUBSPACE_PROPERTY,
SUBSPACE_REGISTRY,
SUBSPACE_REGISTRY_PK,
SUBSPACE_QUEUE_MESSAGE,
SUBSPACE_QUEUE_EVENT,
SUBSPACE_REPORT_OUT,
SUBSPACE_REPORT_IN,
SUBSPACE_LOGS,
SUBSPACE_BLOBS,
SUBSPACE_TELEMETRY_SPAN,
SUBSPACE_TELEMETRY_METRIC,
SUBSPACE_SEARCH_INDEX,
SUBSPACE_DIRECTORY,
] {
let table = char::from(table);
conn.execute(
&format!(
"CREATE TABLE IF NOT EXISTS {table} (
k BLOB PRIMARY KEY,
v BLOB NOT NULL
)"
),
[],
)
.map_err(into_error)?;
}
for table in [SUBSPACE_INDEXES, SUBSPACE_REGISTRY_IDX] {
let table = char::from(table);
conn.execute(
&format!(
"CREATE TABLE IF NOT EXISTS {table} (
k BLOB PRIMARY KEY
)"
),
[],
)
.map_err(into_error)?;
}
for table in [SUBSPACE_COUNTER, SUBSPACE_QUOTA, SUBSPACE_IN_MEMORY_COUNTER] {
conn.execute(
&format!(
"CREATE TABLE IF NOT EXISTS {} (
k BLOB PRIMARY KEY,
v INTEGER NOT NULL DEFAULT 0
)",
char::from(table)
),
[],
)
.map_err(into_error)?;
}
Ok(())
}
pub async fn spawn_worker<U, V>(&self, mut f: U) -> trc::Result<V>
where
U: FnMut() -> trc::Result<V> + Send,
V: Sync + Send + 'static,
{
let (tx, rx) = oneshot::channel();
self.worker_pool.scope(|s| {
s.spawn(|_| {
tx.send(f()).ok();
});
});
match rx.await {
Ok(result) => result,
Err(err) => Err(trc::EventType::Server(trc::ServerEvent::ThreadError).reason(err)),
}
}
}
+26
View File
@@ -0,0 +1,26 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use self::pool::SqliteConnectionManager;
use r2d2::Pool;
use std::fmt::Display;
pub mod blob;
pub mod lookup;
pub mod main;
pub mod pool;
pub mod read;
pub mod write;
pub struct SqliteStore {
pub(crate) conn_pool: Pool<SqliteConnectionManager>,
pub(crate) worker_pool: rayon::ThreadPool,
}
#[inline(always)]
fn into_error(err: impl Display) -> trc::Error {
trc::StoreEvent::SqliteError.reason(err)
}
+118
View File
@@ -0,0 +1,118 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use rusqlite::{Connection, Error, OpenFlags};
use std::fmt;
use std::path::{Path, PathBuf};
#[derive(Debug)]
enum Source {
File(PathBuf),
Memory,
}
type InitFn = dyn Fn(&mut Connection) -> Result<(), rusqlite::Error> + Send + Sync + 'static;
/// An `r2d2::ManageConnection` for `rusqlite::Connection`s.
pub struct SqliteConnectionManager {
source: Source,
flags: OpenFlags,
init: Option<Box<InitFn>>,
}
impl fmt::Debug for SqliteConnectionManager {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let mut builder = f.debug_struct("SqliteConnectionManager");
let _ = builder.field("source", &self.source);
let _ = builder.field("flags", &self.source);
let _ = builder.field("init", &self.init.as_ref().map(|_| "InitFn"));
builder.finish()
}
}
impl SqliteConnectionManager {
/// Creates a new `SqliteConnectionManager` from file.
///
/// See `rusqlite::Connection::open`
pub fn file<P: AsRef<Path>>(path: P) -> Self {
Self {
source: Source::File(path.as_ref().to_path_buf()),
flags: OpenFlags::default(),
init: None,
}
}
/// Creates a new `SqliteConnectionManager` from memory.
pub fn memory() -> Self {
Self {
source: Source::Memory,
flags: OpenFlags::default(),
init: None,
}
}
/// Converts `SqliteConnectionManager` into one that sets OpenFlags upon
/// connection creation.
///
/// See `rustqlite::OpenFlags` for a list of available flags.
pub fn with_flags(self, flags: OpenFlags) -> Self {
Self { flags, ..self }
}
/// Converts `SqliteConnectionManager` into one that calls an initialization
/// function upon connection creation. Could be used to set PRAGMAs, for
/// example.
///
/// ### Example
///
/// Make a `SqliteConnectionManager` that sets the `foreign_keys` pragma to
/// true for every connection.
///
/// ```rust,no_run
/// # use r2d2_sqlite::{SqliteConnectionManager};
/// let manager = SqliteConnectionManager::file("app.db")
/// .with_init(|c| c.execute_batch("PRAGMA foreign_keys=1;"));
/// ```
pub fn with_init<F>(self, init: F) -> Self
where
F: Fn(&mut Connection) -> Result<(), rusqlite::Error> + Send + Sync + 'static,
{
let init: Option<Box<InitFn>> = Some(Box::new(init));
Self { init, ..self }
}
}
fn sleeper(_: i32) -> bool {
std::thread::sleep(std::time::Duration::from_millis(200));
true
}
impl r2d2::ManageConnection for SqliteConnectionManager {
type Connection = Connection;
type Error = rusqlite::Error;
fn connect(&self) -> Result<Connection, Error> {
match self.source {
Source::File(ref path) => Connection::open_with_flags(path, self.flags),
Source::Memory => Connection::open_in_memory_with_flags(self.flags),
}
.and_then(|mut c| {
c.busy_handler(Some(sleeper))?;
match self.init {
None => Ok(c),
Some(ref init) => init(&mut c).map(|_| c),
}
})
}
fn is_valid(&self, conn: &mut Connection) -> Result<(), Error> {
conn.execute_batch("")
}
fn has_broken(&self, _: &mut Connection) -> bool {
false
}
}
+152
View File
@@ -0,0 +1,152 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{SqliteStore, into_error};
use crate::{Deserialize, IterateParams, Key, ValueKey, write::ValueClass};
use rusqlite::OptionalExtension;
impl SqliteStore {
pub(crate) async fn get_value<U>(&self, key: impl Key) -> trc::Result<Option<U>>
where
U: Deserialize + 'static,
{
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
let mut result = conn
.prepare_cached(&format!(
"SELECT v FROM {} WHERE k = ?",
char::from(key.subspace())
))
.map_err(into_error)?;
let key = key.serialize(0);
result
.query_row([&key], |row| {
U::deserialize_with_key(&key, row.get_ref(0)?.as_bytes()?)
.map_err(|err| rusqlite::Error::ToSqlConversionFailure(err.into()))
})
.optional()
.map_err(into_error)
})
.await
}
pub(crate) async fn key_exists(&self, key: impl Key) -> trc::Result<bool> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
let mut result = conn
.prepare_cached(&format!(
"SELECT 1 FROM {} WHERE k = ?",
char::from(key.subspace())
))
.map_err(into_error)?;
let key = key.serialize(0);
result
.query_row([&key], |_| Ok(()))
.optional()
.map(|opt| opt.is_some())
.map_err(into_error)
})
.await
}
pub(crate) async fn iterate<T: Key>(
&self,
params: IterateParams<T>,
mut cb: impl for<'x> FnMut(&'x [u8], &'x [u8]) -> trc::Result<bool> + Sync + Send,
) -> trc::Result<()> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
let table = char::from(params.begin.subspace());
let begin = params.begin.serialize(0);
let end = params.end.serialize(0);
let keys = if params.values { "k, v" } else { "k" };
let mut query = conn
.prepare_cached(&match (params.first, params.ascending) {
(true, true) => {
format!(
"SELECT {keys} FROM {table} WHERE k >= ? AND k <= ? ORDER BY k ASC LIMIT 1"
)
}
(true, false) => {
format!(
"SELECT {keys} FROM {table} WHERE k >= ? AND k <= ? ORDER BY k DESC LIMIT 1"
)
}
(false, true) => {
format!("SELECT {keys} FROM {table} WHERE k >= ? AND k <= ? ORDER BY k ASC")
}
(false, false) => {
format!(
"SELECT {keys} FROM {table} WHERE k >= ? AND k <= ? ORDER BY k DESC"
)
}
})
.map_err(into_error)?;
let mut rows = query.query([&begin, &end]).map_err(into_error)?;
if params.values {
while let Some(row) = rows.next().map_err(into_error)? {
let key = row
.get_ref(0)
.map_err(into_error)?
.as_bytes()
.map_err(into_error)?;
let value = row
.get_ref(1)
.map_err(into_error)?
.as_bytes()
.map_err(into_error)?;
if !cb(key, value)? {
break;
}
}
} else {
while let Some(row) = rows.next().map_err(into_error)? {
if !cb(
row.get_ref(0)
.map_err(into_error)?
.as_bytes()
.map_err(into_error)?,
b"",
)? {
break;
}
}
}
Ok(())
})
.await
}
pub(crate) async fn get_counter(
&self,
key: impl Into<ValueKey<ValueClass>> + Sync + Send,
) -> trc::Result<i64> {
let key = key.into();
let table = char::from(key.subspace());
let key = key.serialize(0);
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
match conn
.prepare_cached(&format!("SELECT v FROM {table} WHERE k = ?"))
.map_err(into_error)?
.query_row([&key], |row| row.get::<_, i64>(0))
{
Ok(value) => Ok(value),
Err(rusqlite::Error::QueryReturnedNoRows) => Ok(0),
Err(e) => Err(into_error(e)),
}
})
.await
}
}
+319
View File
@@ -0,0 +1,319 @@
/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL
*/
use super::{SqliteStore, into_error};
use crate::{
IndexKey, Key, LogKey, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER, SUBSPACE_QUOTA,
SUBSPACE_REGISTRY_IDX,
write::{AssignedIds, Batch, MergeResult, Operation, ValueClass, ValueOp},
};
use rusqlite::{OptionalExtension, TransactionBehavior, params};
use trc::AddContext;
impl SqliteStore {
pub(crate) async fn write(&self, batch: Batch<'_>) -> trc::Result<AssignedIds> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let mut conn = manager.get().map_err(into_error)?;
let mut account_id = u32::MAX;
let mut collection = u8::MAX;
let mut document_id = u32::MAX;
let mut change_id = 0u64;
let trx = conn
.transaction_with_behavior(TransactionBehavior::Immediate)
.map_err(into_error)
.caused_by(trc::location!())?;
let mut result = AssignedIds::default();
let has_changes = !batch.changes.is_empty();
if has_changes {
for &account_id in batch.changes.keys() {
let key = ValueClass::ChangeId.serialize(account_id, 0, 0, 0);
let change_id = trx
.prepare_cached(concat!(
"INSERT INTO n (k, v) VALUES (?, ?) ",
"ON CONFLICT(k) DO UPDATE SET v = v + ",
"excluded.v RETURNING v"
))
.map_err(into_error)
.caused_by(trc::location!())?
.query_row(params![&key, &1i64], |row| row.get::<_, i64>(0))
.map_err(into_error)
.caused_by(trc::location!())?;
result.push_change_id(account_id, change_id as u64);
}
}
for op in batch.ops.iter_mut() {
match op {
Operation::AccountId {
account_id: account_id_,
} => {
account_id = *account_id_;
if has_changes {
change_id = result.set_current_change_id(account_id)?;
}
}
Operation::Collection {
collection: collection_,
} => {
collection = u8::from(*collection_);
}
Operation::DocumentId {
document_id: document_id_,
} => {
document_id = *document_id_;
}
Operation::Value { class, op } => {
let key = class.serialize(account_id, collection, document_id, 0);
let subspace = class.subspace(collection);
let table = char::from(subspace);
match op {
ValueOp::Set(value) => {
if subspace != SUBSPACE_REGISTRY_IDX {
trx.prepare_cached(&format!(
"INSERT OR REPLACE INTO {} (k, v) VALUES (?, ?)",
table
))
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key, value])
.map_err(into_error)
.caused_by(trc::location!())?;
} else {
trx.prepare_cached("INSERT OR IGNORE INTO b (k) VALUES (?)")
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key])
.map_err(into_error)
.caused_by(trc::location!())?;
}
}
ValueOp::SetFnc(set_op) => {
let value = (set_op.fnc)(&set_op.params, &result)?;
trx.prepare_cached(&format!(
"INSERT OR REPLACE INTO {} (k, v) VALUES (?, ?)",
table
))
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key, &value])
.map_err(into_error)
.caused_by(trc::location!())?;
}
ValueOp::MergeFnc(merge_op) => {
let merge_result = trx
.prepare_cached(&format!("SELECT v FROM {} WHERE k = ?", table))
.map_err(into_error)
.caused_by(trc::location!())?
.query_row([&key], |row| {
Ok((merge_op.fnc)(
&merge_op.params,
&result,
Some(row.get_ref(0)?.as_bytes()?),
))
})
.optional()
.map_err(into_error)
.caused_by(trc::location!())?
.unwrap_or_else(|| {
(merge_op.fnc)(&merge_op.params, &result, None)
})?;
match merge_result {
MergeResult::Update(value) => {
trx.prepare_cached(&format!(
"INSERT OR REPLACE INTO {} (k, v) VALUES (?, ?)",
table
))
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key, &value])
.map_err(into_error)
.caused_by(trc::location!())?;
}
MergeResult::Delete => {
trx.prepare_cached(&format!(
"DELETE FROM {} WHERE k = ?",
table
))
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key])
.map_err(into_error)
.caused_by(trc::location!())?;
}
MergeResult::Skip => (),
}
}
ValueOp::AtomicAdd(by) => {
if *by >= 0 {
trx.prepare_cached(&format!(
concat!(
"INSERT INTO {} (k, v) VALUES (?, ?) ",
"ON CONFLICT(k) DO UPDATE SET v = v + excluded.v"
),
table
))
.map_err(into_error)
.caused_by(trc::location!())?
.execute(params![&key, *by])
.map_err(into_error)
.caused_by(trc::location!())?;
} else {
trx.prepare_cached(&format!(
"UPDATE {table} SET v = v + ? WHERE k = ?"
))
.map_err(into_error)
.caused_by(trc::location!())?
.execute(params![*by, &key])
.map_err(into_error)
.caused_by(trc::location!())?;
}
}
ValueOp::AddAndGet(by) => {
result.push_counter_id(
trx.prepare_cached(&format!(
concat!(
"INSERT INTO {} (k, v) VALUES (?, ?) ",
"ON CONFLICT(k) DO UPDATE SET v = v + ",
"excluded.v RETURNING v"
),
table
))
.map_err(into_error)
.caused_by(trc::location!())?
.query_row(params![&key, &*by], |row| row.get::<_, i64>(0))
.map_err(into_error)
.caused_by(trc::location!())?,
);
}
ValueOp::Clear => {
trx.prepare_cached(&format!("DELETE FROM {} WHERE k = ?", table))
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key])
.map_err(into_error)
.caused_by(trc::location!())?;
}
}
}
Operation::Index { field, key, set } => {
let key = IndexKey {
account_id,
collection,
document_id,
field: *field,
key: &*key,
}
.serialize(0);
if *set {
trx.prepare_cached("INSERT OR IGNORE INTO i (k) VALUES (?)")
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key])
.map_err(into_error)
.caused_by(trc::location!())?;
} else {
trx.prepare_cached("DELETE FROM i WHERE k = ?")
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key])
.map_err(into_error)
.caused_by(trc::location!())?;
}
}
Operation::Log { collection, set } => {
let key = LogKey {
account_id,
collection: u8::from(*collection),
change_id,
}
.serialize(0);
trx.prepare_cached("INSERT OR REPLACE INTO l (k, v) VALUES (?, ?)")
.map_err(into_error)
.caused_by(trc::location!())?
.execute([&key, set])
.map_err(into_error)
.caused_by(trc::location!())?;
}
Operation::AssertValue {
class,
assert_value,
} => {
let key = class.serialize(account_id, collection, document_id, 0);
let table = char::from(class.subspace(collection));
let matches = trx
.prepare_cached(&format!("SELECT v FROM {} WHERE k = ?", table))
.map_err(into_error)
.caused_by(trc::location!())?
.query_row([&key], |row| {
Ok(assert_value.matches(row.get_ref(0)?.as_bytes()?))
})
.optional()
.map_err(into_error)
.caused_by(trc::location!())?
.unwrap_or_else(|| assert_value.is_none());
if !matches {
trx.rollback()
.map_err(into_error)
.caused_by(trc::location!())?;
return Err(trc::StoreEvent::AssertValueFailed
.into_err()
.caused_by(trc::location!()));
}
}
}
}
trx.commit().map(|_| result).map_err(into_error)
})
.await
}
pub(crate) async fn purge_store(&self) -> trc::Result<()> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
for subspace in [SUBSPACE_QUOTA, SUBSPACE_COUNTER, SUBSPACE_IN_MEMORY_COUNTER] {
conn.prepare_cached(&format!("DELETE FROM {} WHERE v = 0", char::from(subspace),))
.map_err(into_error)
.caused_by(trc::location!())?
.execute([])
.map_err(into_error)
.caused_by(trc::location!())?;
}
Ok(())
})
.await
}
pub(crate) async fn delete_range(&self, from: impl Key, to: impl Key) -> trc::Result<()> {
let manager = self.conn_pool.clone();
self.spawn_worker(move || {
let conn = manager.get().map_err(into_error)?;
conn.prepare_cached(&format!(
"DELETE FROM {} WHERE k >= ? AND k < ?",
char::from(from.subspace()),
))
.map_err(into_error)
.caused_by(trc::location!())?
.execute([from.serialize(0), to.serialize(0)])
.map_err(into_error)
.caused_by(trc::location!())?;
Ok(())
})
.await
}
}