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
+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)
}
}