Initial commit
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
pub mod users;
|
||||
pub mod syncs;
|
||||
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use anyhow::Context as _;
|
||||
use axum::{Json, Router, extract::{Path, State}, response::{IntoResponse, Response}, routing};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::instrument;
|
||||
|
||||
use crate::{app::AppState, app_error::AppError, auth_user::AuthUser, sync::Document};
|
||||
|
||||
|
||||
pub fn syncs_routes() -> Router<AppState> {
|
||||
Router::new()
|
||||
.route("/progress", routing::put(put_progress))
|
||||
.route("/progress/{document}", routing::get(get_progress))
|
||||
}
|
||||
|
||||
#[derive(Debug,Deserialize)]
|
||||
pub struct PutProgress {
|
||||
pub document: String,
|
||||
pub progress: String,
|
||||
pub percentage: f32,
|
||||
pub device: String,
|
||||
pub device_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug,Serialize)]
|
||||
pub struct ProgressCreated {
|
||||
document: String,
|
||||
timestamp: i64,
|
||||
}
|
||||
|
||||
#[instrument(skip(st))]
|
||||
#[axum::debug_handler]
|
||||
pub async fn put_progress(State(st) : State<AppState>, AuthUser{username} : AuthUser, Json(progress): Json<PutProgress>) -> Result<axum::Json<ProgressCreated>, AppError> {
|
||||
tracing::info!("Storing a new progress");
|
||||
let timestamp = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.context("System clock is before unix epoch")?
|
||||
.as_secs() as i64;
|
||||
let document = Document {
|
||||
timestamp,
|
||||
document_id: progress.document,
|
||||
progress: progress.progress,
|
||||
percentage: progress.percentage,
|
||||
device: progress.device,
|
||||
device_id: progress.device_id,
|
||||
};
|
||||
st.sync.set_progress(&username, &document).await?;
|
||||
Ok(Json(ProgressCreated{timestamp: document.timestamp, document: document.document_id}))
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct Progress {
|
||||
document: String,
|
||||
progress: String,
|
||||
percentage: f32,
|
||||
device: String,
|
||||
timestamp: i64,
|
||||
device_id: Option<String>,
|
||||
}
|
||||
|
||||
#[instrument(skip(st))]
|
||||
#[axum::debug_handler]
|
||||
pub async fn get_progress(State(st) : State<AppState>, AuthUser{username} : AuthUser, Path(document): Path<String>) -> Result<Response, AppError> {
|
||||
let d = st.sync.get_progress(&username, &document).await?;
|
||||
tracing::info!("{:?}", d);
|
||||
match d {
|
||||
None => Ok(Json(serde_json::json!({})).into_response()),
|
||||
Some(d) =>
|
||||
Ok(Json(Progress{
|
||||
document: d.document_id,
|
||||
progress: d.progress,
|
||||
percentage: d.percentage,
|
||||
device: d.device,
|
||||
timestamp: d.timestamp,
|
||||
device_id: d.device_id,
|
||||
}).into_response())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
use axum::{Json, Router, extract::State, routing};
|
||||
use secrecy::SecretString;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::instrument;
|
||||
|
||||
use crate::{app::AppState, app_error::AppError, auth_user::AuthUser};
|
||||
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct CreateUser {
|
||||
username: String,
|
||||
password: String,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Debug)]
|
||||
pub struct UserCreated {
|
||||
username: String
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
pub fn users_routes() -> Router<AppState> {
|
||||
Router::new()
|
||||
.route("/create", routing::post(create_user))
|
||||
.route("/auth", routing::get(check_auth))
|
||||
.route("/me", routing::delete(delete_user))
|
||||
}
|
||||
|
||||
#[instrument(skip_all)]
|
||||
#[axum::debug_handler]
|
||||
pub async fn create_user(State(st) : State<AppState>, Json(create_user): Json<CreateUser>) -> Result<axum::Json<UserCreated>, AppError> {
|
||||
let username = st.users
|
||||
.create_user(&create_user.username, SecretString::from(create_user.password))
|
||||
.await?;
|
||||
Ok(Json(UserCreated{username}))
|
||||
}
|
||||
|
||||
#[derive(Serialize, Debug)]
|
||||
pub struct AuthResult {
|
||||
authorized: String
|
||||
}
|
||||
|
||||
#[instrument(skip_all)]
|
||||
#[axum::debug_handler]
|
||||
pub async fn check_auth(_: State<AppState>, _ : AuthUser) -> Result<Json<AuthResult>, AppError> {
|
||||
Ok(Json(AuthResult{authorized: "OK".to_string()}))
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct DeleteResult { deleted: bool }
|
||||
|
||||
#[instrument(skip_all)]
|
||||
#[axum::debug_handler]
|
||||
pub async fn delete_user(State(app_state): State<AppState>, AuthUser{username} : AuthUser) -> Result<Json<DeleteResult>, AppError> {
|
||||
app_state.users.delete_user(&username).await?;
|
||||
Ok(Json(DeleteResult{deleted: true}))
|
||||
}
|
||||
+22
@@ -0,0 +1,22 @@
|
||||
use std::{path::Path, sync::Arc};
|
||||
|
||||
use axum::extract::FromRef;
|
||||
|
||||
use crate::{config::Config, reader::Reader, sync::ProgressSync, users::Users, writer::Writer};
|
||||
|
||||
#[derive(Clone, FromRef)]
|
||||
pub struct AppState {
|
||||
pub users: Users,
|
||||
pub sync: ProgressSync,
|
||||
}
|
||||
|
||||
impl AppState {
|
||||
pub async fn new(db_path: &Path, config: &Config) -> anyhow::Result<Self> {
|
||||
let writer = Arc::new(Writer::new(&db_path).await?);
|
||||
let reader = Arc::new(Reader::new(&db_path).await?);
|
||||
let users = Users::new(writer.clone(), reader.clone(), config.server_secret.clone());
|
||||
let sync = ProgressSync::new(writer.clone(), reader.clone());
|
||||
Ok(Self{users, sync})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
use axum::{Json, http::StatusCode, response::IntoResponse};
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::{sync::SyncError, users::UserError};
|
||||
|
||||
|
||||
// The AppError represents the different error types from upstream kosync server
|
||||
#[derive(Error, Debug)]
|
||||
pub enum AppError {
|
||||
#[error("Username already taken")]
|
||||
UsernameTaken,
|
||||
#[error("Unauthorized")]
|
||||
Unauthorized,
|
||||
#[error("Account not found")]
|
||||
AccountNotFound,
|
||||
#[error(transparent)]
|
||||
Internal(#[from] anyhow::Error),
|
||||
}
|
||||
|
||||
impl From<UserError> for AppError {
|
||||
fn from(value: UserError) -> Self {
|
||||
match value {
|
||||
UserError::UnknownError(error) => Self::Internal(error),
|
||||
UserError::UserAlreadyExists => Self::UsernameTaken,
|
||||
UserError::DatabaseFailure(error) => Self::Internal(error.into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<SyncError> for AppError {
|
||||
fn from(value: SyncError) -> Self {
|
||||
match value {
|
||||
SyncError::DatabaseFailure(error) => Self::Internal(error.into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoResponse for AppError {
|
||||
fn into_response(self) -> axum::response::Response {
|
||||
let (status, error_code, msg) = match self {
|
||||
AppError::UsernameTaken => (StatusCode::PAYMENT_REQUIRED, 2002, "Username is already registered"),
|
||||
AppError::Internal(error) => {
|
||||
tracing::error!(err=?error, "Internal server error");
|
||||
(StatusCode::INTERNAL_SERVER_ERROR, 2000, "Internal server error")
|
||||
},
|
||||
AppError::Unauthorized => {
|
||||
(StatusCode::UNAUTHORIZED, 2001, "Unauthorized")
|
||||
}
|
||||
AppError::AccountNotFound => {
|
||||
(StatusCode::UNAUTHORIZED, 2006, "Account not found")
|
||||
}
|
||||
};
|
||||
let body = serde_json::json!({
|
||||
"code": error_code,
|
||||
"message": msg
|
||||
});
|
||||
(status, Json(body)).into_response()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
use axum::extract::{FromRef as _, FromRequestParts};
|
||||
use secrecy::SecretString;
|
||||
|
||||
use crate::{app::AppState, app_error::AppError, users};
|
||||
|
||||
pub struct AuthUser { pub username: String }
|
||||
|
||||
impl FromRequestParts<AppState> for AuthUser
|
||||
{
|
||||
type Rejection = AppError;
|
||||
|
||||
async fn from_request_parts(
|
||||
parts: &mut axum::http::request::Parts,
|
||||
state: &AppState,
|
||||
) -> Result<Self, Self::Rejection> {
|
||||
let app_state = AppState::from_ref(state);
|
||||
let username = parts
|
||||
.headers
|
||||
.get("x-auth-user").and_then(|h| h.to_str().ok())
|
||||
.ok_or(AppError::Unauthorized)?;
|
||||
let key = parts
|
||||
.headers
|
||||
.get("x-auth-key").and_then(|h| h.to_str().ok().map(SecretString::from))
|
||||
.ok_or(AppError::Unauthorized)?;
|
||||
match app_state.users.authenticate_user(username, key).await? {
|
||||
users::Authenticated::Authenticated(username) => Ok(AuthUser{ username }),
|
||||
users::Authenticated::UserNotFound => Err(AppError::AccountNotFound),
|
||||
users::Authenticated::InvalidSecret => Err(AppError::Unauthorized),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
use std::env;
|
||||
|
||||
use secrecy::SecretString;
|
||||
use thiserror::Error;
|
||||
|
||||
pub struct Config {
|
||||
pub server_secret: SecretString,
|
||||
}
|
||||
|
||||
#[derive(Error, Debug)]
|
||||
pub enum ConfigError {
|
||||
#[error("Missing SERVER_SECRET")]
|
||||
MissingServerSecret,
|
||||
#[error("Unknown error: #{0}")]
|
||||
Unknown(anyhow::Error),
|
||||
}
|
||||
|
||||
impl Config {
|
||||
pub fn new() -> Result<Self, ConfigError> {
|
||||
dotenvy::dotenv().map_err(|err| ConfigError::Unknown(err.into()))?;
|
||||
let server_secret = env::var("SERVER_SECRET").map(SecretString::from).map_err(|_| ConfigError::MissingServerSecret)?;
|
||||
Ok(Self{server_secret})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
pub mod config;
|
||||
pub mod writer;
|
||||
pub mod users;
|
||||
pub mod sync;
|
||||
pub mod reader;
|
||||
pub mod api;
|
||||
pub mod app;
|
||||
pub mod auth_user;
|
||||
pub mod app_error;
|
||||
+39
@@ -0,0 +1,39 @@
|
||||
use std::path::Path;
|
||||
|
||||
use axum::{Json, Router, response::{IntoResponse as _, Response}, routing};
|
||||
use rukosync::{api::{syncs, users}, app::AppState, config::Config, };
|
||||
use tower_http::trace::TraceLayer;
|
||||
use tracing::level_filters::LevelFilter;
|
||||
use tracing_subscriber::{EnvFilter, fmt, layer::SubscriberExt as _, util::SubscriberInitExt as _};
|
||||
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> anyhow::Result<()> {
|
||||
let level_filter = EnvFilter::builder()
|
||||
.with_default_directive(LevelFilter::INFO.into())
|
||||
.from_env_lossy();
|
||||
tracing_subscriber::registry()
|
||||
.with(fmt::layer())
|
||||
.with(level_filter)
|
||||
.init();
|
||||
let config = Config::new()?;
|
||||
|
||||
let app_state = AppState::new(Path::new("/tmp/kosync.db"), &config).await?;
|
||||
|
||||
tracing::info!("Hello");
|
||||
|
||||
let app = Router::new()
|
||||
.route("/healthcheck", routing::get(healthcheck))
|
||||
.nest("/users", users::users_routes())
|
||||
.nest("/syncs", syncs::syncs_routes())
|
||||
.with_state(app_state)
|
||||
.layer(TraceLayer::new_for_http());
|
||||
let listener = tokio::net::TcpListener::bind("0.0.0.0:3000").await?;
|
||||
axum::serve(listener, app).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[axum::debug_handler]
|
||||
async fn healthcheck() -> Response {
|
||||
Json(serde_json::json!({"state": "OK"})).into_response()
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
use std::path::Path;
|
||||
|
||||
use tokio_rusqlite::{Connection, OpenFlags};
|
||||
|
||||
pub struct Reader {
|
||||
pub connection: Connection,
|
||||
}
|
||||
|
||||
impl Reader {
|
||||
pub async fn new(path: &Path) -> Result<Self, tokio_rusqlite::Error> {
|
||||
let connection = Connection::open_with_flags(path, OpenFlags::SQLITE_OPEN_READ_ONLY).await?;
|
||||
Ok(Self{connection})
|
||||
}
|
||||
}
|
||||
+205
@@ -0,0 +1,205 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use thiserror::Error;
|
||||
use tokio_rusqlite::{OptionalExtension as _, named_params, rusqlite};
|
||||
|
||||
use crate::{reader::Reader, writer::Writer};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ProgressSync {
|
||||
writer: Arc<Writer>,
|
||||
reader: Arc<Reader>,
|
||||
}
|
||||
|
||||
|
||||
#[derive(Error, Debug)]
|
||||
pub enum SyncError {
|
||||
#[error("Database call failed {0}")]
|
||||
DatabaseFailure(#[from] tokio_rusqlite::Error)
|
||||
}
|
||||
|
||||
#[derive(Debug, PartialEq, Clone)]
|
||||
pub struct Document {
|
||||
// document identifier, name of the file, md5 of the file or something
|
||||
pub document_id: String,
|
||||
// opaque progress
|
||||
pub progress: String,
|
||||
pub percentage: f32,
|
||||
pub device: String,
|
||||
pub device_id: Option<String>,
|
||||
pub timestamp: i64,
|
||||
}
|
||||
|
||||
impl ProgressSync {
|
||||
pub fn new(writer: Arc<Writer>, reader: Arc<Reader>) -> Self {
|
||||
Self{writer,reader}
|
||||
}
|
||||
|
||||
pub async fn set_progress(&self, username: &str, document: &Document) -> Result<(), SyncError> {
|
||||
let username = username.to_owned();
|
||||
let doc = document.clone();
|
||||
self.writer.connection.call(move |conn| -> rusqlite::Result<()> {
|
||||
conn.execute("insert into progress (user_id, document_id, progress, percentage, device, device_id, timestamp)
|
||||
select id, :document_id, :progress, :percentage, :device, :device_id, :timestamp from users
|
||||
where username = :username
|
||||
on conflict (user_id, document_id) do update
|
||||
set progress = excluded.progress,
|
||||
percentage = excluded.percentage,
|
||||
device = excluded.device,
|
||||
device_id = excluded.device_id,
|
||||
timestamp = excluded.timestamp", named_params!
|
||||
{ ":username": username,
|
||||
":document_id": doc.document_id,
|
||||
":progress": doc.progress,
|
||||
":percentage": doc.percentage,
|
||||
":device": doc.device,
|
||||
":device_id": doc.device_id,
|
||||
":timestamp": doc.timestamp,
|
||||
})?;
|
||||
Ok(())
|
||||
}).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn get_progress(&self, username: &str, document_id: &str) -> Result<Option<Document>, SyncError> {
|
||||
let username = username.to_owned();
|
||||
let document_id = document_id.to_owned();
|
||||
let doc = self.reader.connection.call(move |conn| -> rusqlite::Result<Option<Document>> {
|
||||
let x : Option<Document> = conn.query_row("select
|
||||
document_id, progress, percentage, device, device_id, timestamp
|
||||
from progress
|
||||
left join users on progress.user_id = users.id
|
||||
where users.username = :username and progress.document_id = :document_id", named_params!
|
||||
{ ":username": username,
|
||||
":document_id": document_id,
|
||||
}, |row| {
|
||||
let document_id : String = row.get(0)?;
|
||||
let progress : String = row.get(1)?;
|
||||
let percentage : f32 = row.get(2)?;
|
||||
let device : String = row.get(3)?;
|
||||
let device_id : Option<String> = row.get(4)?;
|
||||
let timestamp : i64 = row.get(5)?;
|
||||
Ok(Document{ document_id, progress, percentage, device, device_id, timestamp })
|
||||
}).optional()?;
|
||||
Ok(x)
|
||||
}).await?;
|
||||
Ok(doc)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use secrecy::SecretString;
|
||||
use tempfile::{TempDir, tempdir};
|
||||
|
||||
use crate::users::Users;
|
||||
|
||||
use super::*;
|
||||
|
||||
async fn create_data() -> anyhow::Result<(TempDir, Users, ProgressSync)> {
|
||||
let temp_dir = tempdir()?;
|
||||
let path = temp_dir.path().join("kosync.db");
|
||||
let writer = Arc::new(Writer::new(&path).await?);
|
||||
let reader = Arc::new(Reader::new(&path).await?);
|
||||
let users = Users::new(writer.clone(), reader.clone(), SecretString::from("secret"));
|
||||
let sync = ProgressSync::new(writer.clone(), reader.clone());
|
||||
Ok((temp_dir, users, sync))
|
||||
}
|
||||
|
||||
fn now() -> anyhow::Result<i64> {
|
||||
let n = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs() as i64;
|
||||
Ok(n)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn put_get_values() -> anyhow::Result<()> {
|
||||
let (_t, users, sync) = create_data().await?;
|
||||
|
||||
let username = "foo@example.com";
|
||||
let username = users.create_user(username, SecretString::from("password")).await?;
|
||||
|
||||
let doc = Document {
|
||||
document_id: "foo".to_string(),
|
||||
progress: "prog".to_string(),
|
||||
percentage: 0.5,
|
||||
device: "kobo".to_string(),
|
||||
device_id: None,
|
||||
timestamp: now()?,
|
||||
};
|
||||
|
||||
sync.set_progress(&username, &doc).await?;
|
||||
let got = sync.get_progress(&username, "foo").await?;
|
||||
|
||||
assert_eq!(got, Some(doc));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn put_gets_latest() -> anyhow::Result<()> {
|
||||
let (_t, users, sync) = create_data().await?;
|
||||
|
||||
let username = "foo@example.com";
|
||||
let username = users.create_user(username, SecretString::from("password")).await?;
|
||||
let document_id = "foo".to_string();
|
||||
|
||||
let doc1 = Document {
|
||||
document_id: document_id.clone(),
|
||||
progress: "prog".to_string(),
|
||||
percentage: 0.5,
|
||||
device: "kobo".to_string(),
|
||||
device_id: None,
|
||||
timestamp: now()?,
|
||||
};
|
||||
let doc2 = Document {
|
||||
document_id: document_id.clone(),
|
||||
progress: "prog2".to_string(),
|
||||
percentage: 0.5,
|
||||
device: "kobo".to_string(),
|
||||
device_id: None,
|
||||
timestamp: now()?,
|
||||
};
|
||||
|
||||
sync.set_progress(&username, &doc1).await?;
|
||||
sync.set_progress(&username, &doc2).await?;
|
||||
let got = sync.get_progress(&username, &document_id).await?;
|
||||
|
||||
assert_eq!(got, Some(doc2));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn put_gets_new_device() -> anyhow::Result<()> {
|
||||
let (_t, users, sync) = create_data().await?;
|
||||
|
||||
let username = "foo@example.com";
|
||||
let username = users.create_user(username, SecretString::from("password")).await?;
|
||||
|
||||
let document_id = "foo".to_string();
|
||||
|
||||
let doc1 = Document {
|
||||
document_id: document_id.clone(),
|
||||
progress: "prog".to_string(),
|
||||
percentage: 0.5,
|
||||
device: "kobo".to_string(),
|
||||
device_id: None,
|
||||
timestamp: now()?,
|
||||
};
|
||||
let doc2 = Document {
|
||||
document_id: document_id.clone(),
|
||||
progress: "prog2".to_string(),
|
||||
percentage: 0.5,
|
||||
device: "x4".to_string(),
|
||||
device_id: None,
|
||||
timestamp: now()? + 3,
|
||||
};
|
||||
|
||||
sync.set_progress(&username, &doc1).await?;
|
||||
sync.set_progress(&username, &doc2).await?;
|
||||
let got = sync.get_progress(&username, &document_id).await?;
|
||||
|
||||
assert_eq!(got, Some(doc2));
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
+190
@@ -0,0 +1,190 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use hmac::{Hmac, KeyInit as _, Mac as _, digest::InvalidLength};
|
||||
use secrecy::{ExposeSecret as _, SecretBox, SecretSlice, SecretString};
|
||||
use sha2::Sha256;
|
||||
use thiserror::Error;
|
||||
use tokio_rusqlite::{OptionalExtension as _, rusqlite};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::{reader::Reader, writer::Writer};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Users {
|
||||
writer : Arc<Writer>,
|
||||
reader : Arc<Reader>,
|
||||
secret: SecretString
|
||||
}
|
||||
|
||||
type Hmac256 = Hmac<Sha256>;
|
||||
|
||||
#[derive(Error, Debug)]
|
||||
pub enum UserError {
|
||||
#[error("Unknown error #{0}")]
|
||||
UnknownError(anyhow::Error),
|
||||
#[error("User already exists")]
|
||||
UserAlreadyExists,
|
||||
#[error("Database failure #{0}")]
|
||||
DatabaseFailure(tokio_rusqlite::Error),
|
||||
|
||||
}
|
||||
|
||||
impl From<InvalidLength> for UserError {
|
||||
fn from(value: InvalidLength) -> Self {
|
||||
Self::UnknownError(value.into())
|
||||
}
|
||||
}
|
||||
|
||||
impl From<tokio_rusqlite::Error> for UserError {
|
||||
fn from(value: tokio_rusqlite::Error) -> Self {
|
||||
Self::DatabaseFailure(value)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Eq, PartialEq)]
|
||||
pub enum Authenticated {
|
||||
Authenticated(String),
|
||||
UserNotFound,
|
||||
InvalidSecret
|
||||
}
|
||||
|
||||
impl Users {
|
||||
pub fn new(writer: Arc<Writer>, reader: Arc<Reader>, secret: SecretString) -> Self {
|
||||
Self{writer, reader, secret}
|
||||
}
|
||||
|
||||
fn create_secret(&self, password: &SecretString) -> Result<SecretBox<[u8]>, UserError> {
|
||||
let mut mac = Hmac256::new_from_slice(self.secret.expose_secret().as_bytes())?;
|
||||
mac.update(password.expose_secret().as_bytes());
|
||||
let secret = SecretSlice::from(mac.finalize().into_bytes().as_slice().to_vec());
|
||||
Ok(secret)
|
||||
}
|
||||
|
||||
pub async fn create_user(&self, username: &str, password: SecretString) -> Result<String, UserError> {
|
||||
let secret = self.create_secret(&password)?;
|
||||
let username = username.to_owned();
|
||||
let created_username = self.writer.connection.call(move |conn| -> rusqlite::Result<Option<String>> {
|
||||
let params = (Uuid::new_v4(), &username, secret.expose_secret());
|
||||
let inserted = conn.execute("insert into users (id, username, secret) values (?, ?, ?) on conflict(username) do nothing", params)?;
|
||||
if inserted == 1 { Ok(Some(username)) } else { Ok(None) }
|
||||
}).await?;
|
||||
created_username.ok_or(UserError::UserAlreadyExists)
|
||||
}
|
||||
|
||||
|
||||
pub async fn authenticate_user(&self, username: &str, password: SecretString) -> Result<Authenticated, UserError> {
|
||||
let secret = self.create_secret(&password)?;
|
||||
let username = username.to_owned();
|
||||
let authenticated = self.reader.connection.call(move |conn| -> rusqlite::Result<Authenticated> {
|
||||
let user = conn.query_row("select username, secret from users where username = ?", [&username], |row| {
|
||||
let u : String = row.get(0)?;
|
||||
let s : SecretBox<[u8]> = SecretSlice::from(row.get::<_, Box<[u8]>>(1)?);
|
||||
Ok((u,s))
|
||||
}).optional()?;
|
||||
match user {
|
||||
None => Ok(Authenticated::UserNotFound),
|
||||
Some((username, s)) => {
|
||||
if secret.expose_secret() == s.expose_secret()
|
||||
{ Ok(Authenticated::Authenticated(username)) }
|
||||
else { Ok(Authenticated::InvalidSecret) }
|
||||
}
|
||||
}
|
||||
}).await?;
|
||||
Ok(authenticated)
|
||||
}
|
||||
|
||||
pub async fn delete_user(&self, username: &str) -> Result<(), UserError> {
|
||||
let username = username.to_owned();
|
||||
self.writer.connection.call(move |conn| -> rusqlite::Result<()> {
|
||||
conn.execute("delete from users where username = ?", [username])?;
|
||||
Ok(())
|
||||
}).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use tempfile::{TempDir, tempdir};
|
||||
|
||||
use super::*;
|
||||
|
||||
async fn create_users() -> anyhow::Result<(TempDir, Users)> {
|
||||
let temp_dir = tempdir()?;
|
||||
let path = temp_dir.path().join("kosync.db");
|
||||
let writer = Arc::new(Writer::new(&path).await?);
|
||||
let reader = Arc::new(Reader::new(&path).await?);
|
||||
let users = Users::new(writer, reader, SecretString::from("secret"));
|
||||
Ok((temp_dir, users))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_user_returns_username() -> anyhow::Result<()> {
|
||||
let (_t, users) = create_users().await?;
|
||||
let username = "foo@example.com";
|
||||
let got = users.create_user(username, SecretString::from("password")).await?;
|
||||
assert_eq!(username, got);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn user_can_authenticate() -> anyhow::Result<()> {
|
||||
let (_t, users) = create_users().await?;
|
||||
let username = "foo@example.com";
|
||||
let password = SecretString::from("password");
|
||||
let _ = users.create_user(username, password.clone()).await?;
|
||||
let auth = users.authenticate_user(username, password.clone()).await?;
|
||||
let wanted = Authenticated::Authenticated(username.to_owned());
|
||||
std::assert_eq!(auth, wanted);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn missing_user_not_found() -> anyhow::Result<()> {
|
||||
let (_t, users) = create_users().await?;
|
||||
let username = "foo@example.com";
|
||||
let password = SecretString::from("password");
|
||||
let _ = users.create_user(username, password.clone()).await?;
|
||||
let auth = users.authenticate_user("bar@example.com", password.clone()).await?;
|
||||
let wanted = Authenticated::UserNotFound;
|
||||
std::assert_eq!(auth, wanted);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn wrong_password_invalid() -> anyhow::Result<()> {
|
||||
let (_t, users) = create_users().await?;
|
||||
let username = "foo@example.com";
|
||||
let password = SecretString::from("password");
|
||||
let wrong_password = SecretString::from("passwordxyz");
|
||||
let _ = users.create_user(username, password.clone()).await?;
|
||||
let auth = users.authenticate_user(username, wrong_password.clone()).await?;
|
||||
let wanted = Authenticated::InvalidSecret;
|
||||
std::assert_eq!(auth, wanted);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn creating_user_twice_fails() -> anyhow::Result<()> {
|
||||
let (_t, users) = create_users().await?;
|
||||
let username = "foo@example.com";
|
||||
let _ = users.create_user(username, SecretString::from("password")).await?;
|
||||
let got = users.create_user(username, SecretString::from("password")).await;
|
||||
std::assert_matches!(got, Err(UserError::UserAlreadyExists));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn user_not_found_after_delete() -> anyhow::Result<()> {
|
||||
let (_t, users) = create_users().await?;
|
||||
let username = "foo@example.com";
|
||||
let password = SecretString::from("password");
|
||||
let _ = users.create_user(username, password.clone()).await?;
|
||||
let _ = users.delete_user(username).await?;
|
||||
let auth = users.authenticate_user(username, password.clone()).await?;
|
||||
let wanted = Authenticated::UserNotFound;
|
||||
std::assert_eq!(auth, wanted);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
use std::path::Path;
|
||||
|
||||
use tokio_rusqlite::{Connection, rusqlite};
|
||||
|
||||
|
||||
pub struct Writer {
|
||||
// Note that I don't need to wrap this in a mutex because
|
||||
// tokio_rusqlite already serializes each request into a single executor
|
||||
pub connection: Connection,
|
||||
}
|
||||
|
||||
impl Writer {
|
||||
pub async fn new(path: &Path) -> Result<Self, tokio_rusqlite::Error> {
|
||||
let connection = Connection::open(path).await?;
|
||||
|
||||
connection.call(|conn| -> rusqlite::Result<()> {
|
||||
conn.execute("create table if not exists users (id primary key, username not null, secret not null, unique(username))", [])?;
|
||||
conn.execute("create table if not exists progress (user_id not null, document_id not null, progress not null, percentage not null, device not null, device_id, timestamp not null, unique(user_id, document_id), foreign key (user_id) references users(id))", [])?;
|
||||
Ok(())
|
||||
}).await?;
|
||||
|
||||
Ok(Self{connection})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user