use std::collections::HashMap;
use async_trait::async_trait;
use atuin_domain::record::{
EncryptedData, HostId, Record, RecordIdx, RecordSeriesKey, RecordStatus, RecordTag,
};
use atuin_server_database::models::{NewSession, NewUser, Session, User};
use atuin_server_database::{Database, DbError, DbResult, DbSettings};
use rand::Rng;
use sqlx::postgres::PgPoolOptions;
use tracing::instrument;
use uuid::Uuid;
use wrappers::DbRecord;
mod wrappers;
const MIN_PG_VERSION: u32 = 14;
#[derive(Clone)]
pub struct Postgres {
pool: sqlx::Pool<sqlx::postgres::Postgres>,
read_pool: Option<sqlx::Pool<sqlx::postgres::Postgres>>,
}
impl Postgres {
fn read_pool(&self) -> &sqlx::Pool<sqlx::postgres::Postgres> {
self.read_pool.as_ref().unwrap_or(&self.pool)
}
}
#[async_trait]
impl Database for Postgres {
async fn new(settings: &DbSettings) -> DbResult<Self> {
let pool =
PgPoolOptions::new().max_connections(100).connect(settings.db_uri.as_str()).await?;
let pg_major_version: u32 = pool
.acquire()
.await?
.server_version_num()
.ok_or(DbError::Other(eyre::Report::msg("could not get PostgreSQL version")))?
/ 10000;
if pg_major_version < MIN_PG_VERSION {
return Err(DbError::Other(eyre::Report::msg(format!(
"unsupported PostgreSQL version {pg_major_version}, minimum required is \
{MIN_PG_VERSION}"
))));
}
sqlx::migrate!("./migrations")
.run(&pool)
.await
.map_err(|error| DbError::Other(error.into()))?;
let read_pool = if let Some(read_db_uri) = &settings.read_db_uri {
tracing::info!("Connecting to read replica database");
let read_pool =
PgPoolOptions::new().max_connections(100).connect(read_db_uri.as_str()).await?;
let read_pg_major_version: u32 =
read_pool.acquire().await?.server_version_num().ok_or(DbError::Other(
eyre::Report::msg("could not get PostgreSQL version from read replica"),
))? / 10000;
if read_pg_major_version < MIN_PG_VERSION {
return Err(DbError::Other(eyre::Report::msg(format!(
"unsupported PostgreSQL version {read_pg_major_version} on read replica, \
minimum required is {MIN_PG_VERSION}"
))));
}
Some(read_pool)
} else {
None
};
Ok(Self { pool, read_pool })
}
#[instrument(skip_all)]
async fn get_session(&self, token: &str) -> DbResult<Session> {
sqlx::query_as("select id, user_id, token from sessions where token = $1")
.bind(token)
.fetch_one(self.read_pool())
.await
.map_err(Into::into)
}
#[instrument(skip_all)]
async fn get_user(&self, username: &str) -> DbResult<User> {
sqlx::query_as("select id, username, email, password from users where username = $1")
.bind(username)
.fetch_one(self.read_pool())
.await
.map_err(Into::into)
}
#[instrument(skip_all)]
async fn get_session_user(&self, token: &str) -> DbResult<User> {
sqlx::query_as(
"select users.id, users.username, users.email, users.password from users
inner join sessions
on users.id = sessions.user_id
and sessions.token = $1",
)
.bind(token)
.fetch_one(self.read_pool())
.await
.map_err(Into::into)
}
async fn delete_store(&self, user: &User) -> DbResult<()> {
let mut tx = self.pool.begin().await?;
sqlx::query(
"delete from store
where user_id = $1",
)
.bind(user.id)
.execute(&mut *tx)
.await?;
sqlx::query(
"delete from store_idx_cache
where user_id = $1",
)
.bind(user.id)
.execute(&mut *tx)
.await?;
tx.commit().await?;
Ok(())
}
#[instrument(skip_all)]
async fn delete_user(&self, u: &User) -> DbResult<()> {
sqlx::query("delete from sessions where user_id = $1")
.bind(u.id)
.execute(&self.pool)
.await?;
sqlx::query("delete from history where user_id = $1")
.bind(u.id)
.execute(&self.pool)
.await?;
sqlx::query("delete from store where user_id = $1").bind(u.id).execute(&self.pool).await?;
sqlx::query("delete from total_history_count_user where user_id = $1")
.bind(u.id)
.execute(&self.pool)
.await?;
sqlx::query("delete from users where id = $1").bind(u.id).execute(&self.pool).await?;
Ok(())
}
#[instrument(skip_all)]
async fn update_user_password(&self, user: &User) -> DbResult<()> {
sqlx::query(
"update users
set password = $1
where id = $2",
)
.bind(&user.password)
.bind(user.id)
.execute(&self.pool)
.await?;
Ok(())
}
#[instrument(skip_all)]
async fn add_user(&self, user: &NewUser) -> DbResult<i64> {
let email: &str = &user.email;
let username: &str = &user.username;
let password: &str = &user.password;
let res: (i64,) = sqlx::query_as(
"insert into users
(username, email, password)
values($1, $2, $3)
returning id",
)
.bind(username)
.bind(email)
.bind(password)
.fetch_one(&self.pool)
.await?;
Ok(res.0)
}
#[instrument(skip_all)]
async fn add_session(&self, session: &NewSession) -> DbResult<()> {
let token: &str = &session.token;
sqlx::query(
"insert into sessions
(user_id, token)
values($1, $2)",
)
.bind(session.user_id)
.bind(token)
.execute(&self.pool)
.await?;
Ok(())
}
#[instrument(skip_all)]
async fn get_user_session(&self, u: &User) -> DbResult<Session> {
sqlx::query_as("select id, user_id, token from sessions where user_id = $1")
.bind(u.id)
.fetch_one(self.read_pool())
.await
.map_err(Into::into)
}
#[instrument(skip_all)]
async fn add_records(&self, user: &User, records: &[Record<EncryptedData>]) -> DbResult<()> {
let mut tx = self.pool.begin().await?;
let mut heads = HashMap::<(HostId, &str), u64>::new();
for i in records {
let id = atuin_common::utils::uuid_v7();
let result = sqlx::query(
"insert into store
(id, client_id, host, idx, timestamp, version, tag, data, cek, user_id)
values ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
on conflict do nothing
",
)
.bind(id)
.bind(i.id)
.bind(i.host.id)
.bind(i.idx as i64)
.bind(i.timestamp as i64) .bind(i.version.as_str())
.bind(i.tag.as_str())
.bind(&i.data.raw)
.bind(&i.data.cek)
.bind(user.id)
.execute(&mut *tx)
.await?;
if result.rows_affected() > 0 {
heads
.entry((i.host.id, i.tag.as_str()))
.and_modify(|e| {
if i.idx > *e {
*e = i.idx
}
})
.or_insert(i.idx);
}
}
for ((host, tag), idx) in heads {
sqlx::query(
"insert into store_idx_cache
(user_id, host, tag, idx)
values ($1, $2, $3, $4)
on conflict(user_id, host, tag) do update set idx = greatest(store_idx_cache.idx, \
$4)
",
)
.bind(user.id)
.bind(host)
.bind(tag)
.bind(idx as i64)
.execute(&mut *tx)
.await?;
}
tx.commit().await?;
Ok(())
}
#[instrument(skip_all)]
async fn next_records(
&self,
user: &User,
series: &RecordSeriesKey,
start: Option<RecordIdx>,
count: u64,
) -> DbResult<Vec<Record<EncryptedData>>> {
tracing::debug!("{:?} - {:?} - {:?}", series.host_id, series.tag, start);
let start = start.unwrap_or(0);
let records: Result<Vec<DbRecord>, DbError> = sqlx::query_as(
"select client_id, host, idx, timestamp, version, tag, data, cek from store
where user_id = $1
and tag = $2
and host = $3
and idx >= $4
order by idx asc
limit $5",
)
.bind(user.id)
.bind(series.tag.as_str())
.bind(series.host_id)
.bind(start as i64)
.bind(count as i64)
.fetch_all(self.read_pool())
.await
.map_err(Into::into);
let ret = match records {
Ok(records) => {
let records: Vec<Record<EncryptedData>> = records
.into_iter()
.map(|f| {
let record: Record<EncryptedData> = f.into();
record
})
.collect();
records
}
Err(DbError::NotFound) => {
tracing::debug!("no records found in store: {:?}/{}", series.host_id, series.tag);
return Ok(vec![]);
}
Err(e) => return Err(e),
};
Ok(ret)
}
async fn status(&self, user: &User) -> DbResult<RecordStatus> {
const STATUS_SQL: &str =
"select host, tag, max(idx) from store where user_id = $1 group by host, tag";
let idx_cache_rollout = std::env::var("IDX_CACHE_ROLLOUT").unwrap_or("0".to_string());
let idx_cache_rollout = idx_cache_rollout.parse::<f64>().unwrap_or(0.0);
let use_idx_cache = rand::thread_rng().gen_bool(idx_cache_rollout / 100.0);
let mut res: Vec<(Uuid, String, i64)> = if use_idx_cache {
tracing::debug!("using idx cache for user {}", user.id);
sqlx::query_as("select host, tag, idx from store_idx_cache where user_id = $1")
.bind(user.id)
.fetch_all(self.read_pool())
.await?
} else {
tracing::debug!("using aggregate query for user {}", user.id);
sqlx::query_as(STATUS_SQL).bind(user.id).fetch_all(self.read_pool()).await?
};
res.sort();
let mut status = RecordStatus::new();
for i in &res {
status.set_raw(
RecordSeriesKey::new(HostId(i.0), RecordTag::from(i.1.clone())),
i.2 as u64,
);
}
Ok(status)
}
}