Initial commit
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
/target
|
||||
.direnv
|
||||
.envrc
|
||||
.env
|
||||
Generated
+1183
File diff suppressed because it is too large
Load Diff
+23
@@ -0,0 +1,23 @@
|
||||
[package]
|
||||
name = "rukosync"
|
||||
version = "0.1.0"
|
||||
edition = "2024"
|
||||
|
||||
[dependencies]
|
||||
anyhow = { version = "1.0.104", features = ["backtrace"] }
|
||||
axum = { version = "0.8.9", features = ["macros"] }
|
||||
dotenvy = "0.15.7"
|
||||
hmac = "0.13.0"
|
||||
secrecy = { version = "0.10.3", features = ["serde"] }
|
||||
serde = { version = "1.0.229", features = ["derive"] }
|
||||
serde_json = "1.0.151"
|
||||
sha2 = "0.11.0"
|
||||
tempfile = "3.27.0"
|
||||
thiserror = "2.0.21"
|
||||
tokio = { version = "1.53.2", features = ["full"] }
|
||||
tokio-rusqlite = { version = "0.8.0", features = ["bundled", "uuid"] }
|
||||
tower = { version = "0.5.3", features = ["tracing"] }
|
||||
tower-http = { version = "0.7.1", features = ["trace", "request-id"] }
|
||||
tracing = "0.1.44"
|
||||
tracing-subscriber = { version = "0.3.23", features = ["env-filter"] }
|
||||
uuid = { version = "1.27.0", features = ["v4", "serde"] }
|
||||
Generated
+61
@@ -0,0 +1,61 @@
|
||||
{
|
||||
"nodes": {
|
||||
"flake-utils": {
|
||||
"inputs": {
|
||||
"systems": "systems"
|
||||
},
|
||||
"locked": {
|
||||
"lastModified": 1731533236,
|
||||
"narHash": "sha256-l0KFg5HjrsfsO/JpG+r7fRrqm12kzFHyUHqHCVpMMbI=",
|
||||
"owner": "numtide",
|
||||
"repo": "flake-utils",
|
||||
"rev": "11707dc2f618dd54ca8739b309ec4fc024de578b",
|
||||
"type": "github"
|
||||
},
|
||||
"original": {
|
||||
"owner": "numtide",
|
||||
"repo": "flake-utils",
|
||||
"type": "github"
|
||||
}
|
||||
},
|
||||
"nixpkgs": {
|
||||
"locked": {
|
||||
"lastModified": 1791048980,
|
||||
"narHash": "sha256-KgItSKML8Xvte0B7/uGnBDsYzSnnKHcOaiUbgWBXLXw=",
|
||||
"owner": "NixOS",
|
||||
"repo": "nixpkgs",
|
||||
"rev": "a7868a727837f3c09cee2ce0ca671c76b1589fed",
|
||||
"type": "github"
|
||||
},
|
||||
"original": {
|
||||
"owner": "NixOS",
|
||||
"ref": "nixos-unstable",
|
||||
"repo": "nixpkgs",
|
||||
"type": "github"
|
||||
}
|
||||
},
|
||||
"root": {
|
||||
"inputs": {
|
||||
"flake-utils": "flake-utils",
|
||||
"nixpkgs": "nixpkgs"
|
||||
}
|
||||
},
|
||||
"systems": {
|
||||
"locked": {
|
||||
"lastModified": 1681028828,
|
||||
"narHash": "sha256-Vy1rq5AaRuLzOxct8nz4T6wlgyUR7zLU309k9mBC768=",
|
||||
"owner": "nix-systems",
|
||||
"repo": "default",
|
||||
"rev": "da67096a3b9bf56a91d16901293e51ba5b49a27e",
|
||||
"type": "github"
|
||||
},
|
||||
"original": {
|
||||
"owner": "nix-systems",
|
||||
"repo": "default",
|
||||
"type": "github"
|
||||
}
|
||||
}
|
||||
},
|
||||
"root": "root",
|
||||
"version": 7
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
{
|
||||
description = "Rust development environment";
|
||||
|
||||
inputs = {
|
||||
nixpkgs.url = "github:NixOS/nixpkgs/nixos-unstable";
|
||||
flake-utils.url = "github:numtide/flake-utils";
|
||||
};
|
||||
|
||||
outputs = { self, nixpkgs, flake-utils }:
|
||||
flake-utils.lib.eachDefaultSystem (system:
|
||||
let
|
||||
pkgs = nixpkgs.legacyPackages.${system};
|
||||
# Read the file relative to the flake's root
|
||||
overrides = (builtins.fromTOML (builtins.readFile (self + "/rust-toolchain.toml")));
|
||||
libPath = with pkgs; lib.makeLibraryPath [
|
||||
# load external libraries that you need in your rust project here
|
||||
];
|
||||
|
||||
mealie-cli = pkgs.rustPlatform.buildRustPackage {
|
||||
pname = "mealie-cli";
|
||||
version = "0.5.0";
|
||||
src = self;
|
||||
cargoLock = {
|
||||
lockFile = ./Cargo.lock;
|
||||
};
|
||||
nativeBuildInputs = [ pkgs.pkg-config ];
|
||||
buildInputs = [ ];
|
||||
};
|
||||
in
|
||||
{
|
||||
packages.default = mealie-cli;
|
||||
packages.mealie-cli = mealie-cli;
|
||||
|
||||
apps.default = flake-utils.lib.mkApp {
|
||||
drv = mealie-cli;
|
||||
};
|
||||
|
||||
devShells.default = pkgs.mkShell rec {
|
||||
nativeBuildInputs = [ pkgs.pkg-config ];
|
||||
buildInputs = with pkgs; [
|
||||
clang
|
||||
llvmPackages.bintools
|
||||
rustup
|
||||
];
|
||||
|
||||
RUSTC_VERSION = overrides.toolchain.channel;
|
||||
|
||||
# https://github.com/rust-lang/rust-bindgen#environment-variables
|
||||
LIBCLANG_PATH = pkgs.lib.makeLibraryPath [ pkgs.llvmPackages_latest.libclang.lib ];
|
||||
|
||||
shellHook = ''
|
||||
export PATH=$PATH:''${CARGO_HOME:-~/.cargo}/bin
|
||||
export PATH=$PATH:''${RUSTUP_HOME:-~/.rustup}/toolchains/$RUSTC_VERSION-x86_64-unknown-linux-gnu/bin/
|
||||
'';
|
||||
|
||||
# Add precompiled library to rustc search path
|
||||
RUSTFLAGS = (builtins.map (a: ''-L ${a}/lib'') [
|
||||
# add libraries here (e.g. pkgs.libvmi)
|
||||
]);
|
||||
|
||||
LD_LIBRARY_PATH = pkgs.lib.makeLibraryPath (buildInputs ++ nativeBuildInputs);
|
||||
|
||||
|
||||
# Add glibc, clang, glib, and other headers to bindgen search path
|
||||
BINDGEN_EXTRA_CLANG_ARGS =
|
||||
# Includes normal include path
|
||||
(builtins.map (a: ''-I"${a}/include"'') [
|
||||
# add dev libraries here (e.g. pkgs.libvmi.dev)
|
||||
pkgs.glibc.dev
|
||||
])
|
||||
# Includes with special directory paths
|
||||
++ [
|
||||
''-I"${pkgs.llvmPackages_latest.libclang.lib}/lib/clang/${pkgs.llvmPackages_latest.libclang.version}/include"''
|
||||
''-I"${pkgs.glib.dev}/include/glib-2.0"''
|
||||
''-I${pkgs.glib.out}/lib/glib-2.0/include/''
|
||||
];
|
||||
};
|
||||
}
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
[toolchain]
|
||||
channel = "stable"
|
||||
@@ -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