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:
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user