diff --git a/src/db/mod.rs b/src/db/mod.rs index 6286d35..13f9ea5 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -40,3 +40,11 @@ pub fn open_database() -> Result, Pixiv } Err(Box::new(String::from(gettext("Unknown database type.")))) } + +/// Open the database and initialize it +pub async fn open_and_init_database( +) -> Result, PixivDownloaderDbError> { + let db = open_database()?; + db.init().await?; + Ok(db) +} diff --git a/src/db/sqlite/db.rs b/src/db/sqlite/db.rs index 09c6e51..ba0f1d6 100644 --- a/src/db/sqlite/db.rs +++ b/src/db/sqlite/db.rs @@ -3,14 +3,55 @@ use super::super::{ }; use super::SqliteError; use futures_util::lock::Mutex; -use rusqlite::{Connection, OpenFlags, OptionalExtension}; +use rusqlite::{Connection, OpenFlags}; use std::collections::HashMap; +const AUTHORS_TABLE: &'static str = "CREATE TABLE authors ( +id INT, +name TEXT, +creator_id TEXT, +icon INT, +big_icon INT, +background INT, +comment TEXT, +webpage TEXT, +PRIMARY KEY (id) +);"; +const FILES_TABLE: &'static str = "CREATE TABLE files ( +id INT, +path TEXT, +last_modified DATETIME, +etag TEXT, +url TEXT, +PRIMARY KEY (id) +);"; +const PIXIV_ARTWORK_TAGS_TABLE: &'static str = "CREATE TABLE pixiv_artwork_tags ( +id INT, +tag_id INT, +);"; +const PIXIV_ARTWORKS_TABLE: &'static str = "CREATE TABLE pixiv_artworks ( +id INT, +title TEXT, +author TEXT, +uid INT, +description TEXT, +count INT, +);"; +const PIXIV_FILES_TABLE: &'static str = "CREATE TABLE pixiv_files ( +id INT, +file_id INT, +page INT, +);"; const TAGS_TABLE: &'static str = "CREATE TABLE tags ( id INT, name TEXT, PRIMARY KEY (id) );"; +const TAGS_I18N_TABLE: &'static str = "CREATE TABLE tags_i18n ( +id INT, +lang TEXT, +translated TEXT, +);"; const VERSION_TABLE: &'static str = "CREATE TABLE version ( id TEXT, v1 INT, @@ -26,6 +67,57 @@ pub struct PixivDownloaderSqlite { } impl PixivDownloaderSqlite { + /// Check if the database needed create all tables. + async fn _check_database(&self) -> Result { + let tables = self._get_exists_table().await?; + let db_version = if tables.contains_key("version") { + self._read_version().await? + } else { + None + }; + let db_version = match db_version { + Some(v) => v, + None => { + return Ok(false); + } + }; + if db_version > VERSION { + return Err(SqliteError::DatabaseVersionTooNew); + } + Ok(true) + } + + /// Create tables + async fn _create_table(&self) -> Result<(), SqliteError> { + let tables = self._get_exists_table().await?; + if !tables.contains_key("version") { + self.db.lock().await.execute(VERSION_TABLE, [])?; + self._write_version().await?; + } + if !tables.contains_key("authors") { + self.db.lock().await.execute(AUTHORS_TABLE, [])?; + } + if !tables.contains_key("files") { + self.db.lock().await.execute(FILES_TABLE, [])?; + } + if !tables.contains_key("pixiv_artwork_tags") { + self.db.lock().await.execute(PIXIV_ARTWORK_TAGS_TABLE, [])?; + } + if !tables.contains_key("pixiv_artworks") { + self.db.lock().await.execute(PIXIV_ARTWORKS_TABLE, [])?; + } + if !tables.contains_key("pixiv_files") { + self.db.lock().await.execute(PIXIV_FILES_TABLE, [])?; + } + if !tables.contains_key("tags") { + self.db.lock().await.execute(TAGS_TABLE, [])?; + } + if !tables.contains_key("tags_i18n") { + self.db.lock().await.execute(TAGS_I18N_TABLE, [])?; + } + Ok(()) + } + /// Get all exists tables async fn _get_exists_table(&self) -> Result, SqliteError> { let con = self.db.lock().await; @@ -38,6 +130,26 @@ impl PixivDownloaderSqlite { Ok(tables) } + async fn _read_version(&self) -> Result, SqliteError> { + let con = self.db.lock().await; + let mut stmt = con.prepare("SELECT v1, v2, v3, v4 FROM version WHERE id='main';")?; + let mut rows = stmt.query([])?; + if let Some(row) = rows.next()? { + Ok(Some([row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?])) + } else { + Ok(None) + } + } + + async fn _write_version(&self) -> Result<(), SqliteError> { + let con = self.db.lock().await; + let mut stmt = con.prepare( + "INSERT OR REPLACE INTO INTO version (id, v1, v2, v3, v4) VALUES ('main', ?, ?, ?, ?);", + )?; + stmt.execute([VERSION[0], VERSION[1], VERSION[2], VERSION[3]])?; + Ok(()) + } + fn _new(cfg: &PixivDownloaderSqliteConfig) -> Result { let db = Connection::open_with_flags( &cfg.path, @@ -48,6 +160,12 @@ impl PixivDownloaderSqlite { )?; Ok(Self { db: Mutex::new(db) }) } + + /// Optimize the database + pub async fn vacuum(&self) -> Result<(), SqliteError> { + self.db.lock().await.execute("VACUUM;", [])?; + Ok(()) + } } #[async_trait] @@ -65,6 +183,9 @@ impl PixivDownloaderDb for PixivDownloaderSqlite { } async fn init(&self) -> Result<(), PixivDownloaderDbError> { + if !self._check_database().await? { + self._create_table().await?; + } Ok(()) } } diff --git a/src/db/sqlite/error.rs b/src/db/sqlite/error.rs index 2cb144d..dc70347 100644 --- a/src/db/sqlite/error.rs +++ b/src/db/sqlite/error.rs @@ -1,4 +1,5 @@ #[derive(derive_more::Display, derive_more::From)] pub enum SqliteError { DbError(rusqlite::Error), + DatabaseVersionTooNew, } diff --git a/src/main.rs b/src/main.rs index 0743444..4ec57a7 100644 --- a/src/main.rs +++ b/src/main.rs @@ -137,7 +137,7 @@ impl Main { #[cfg(feature = "server")] Command::Server => { let addr = get_helper().server(); - match server::service::start_server(&addr) { + match server::service::start_server(&addr).await { Ok(server) => { println!("Listening on http://{}", addr); match server.await { diff --git a/src/server/context.rs b/src/server/context.rs index a3ab330..b735f98 100644 --- a/src/server/context.rs +++ b/src/server/context.rs @@ -1,18 +1,17 @@ use super::cors::CorsContext; -use crate::db::{open_database, PixivDownloaderDb}; +use crate::db::{open_and_init_database, PixivDownloaderDb}; use crate::gettext; -use std::default::Default; pub struct ServerContext { pub cors: CorsContext, pub db: Box, } -impl Default for ServerContext { - fn default() -> Self { +impl ServerContext { + pub async fn default() -> Self { Self { cors: CorsContext::default(), - db: match open_database() { + db: match open_and_init_database().await { Ok(db) => db, Err(e) => panic!("{} {}", gettext("Failed to open database:"), e), }, diff --git a/src/server/service.rs b/src/server/service.rs index d3abfb1..64dfeba 100644 --- a/src/server/service.rs +++ b/src/server/service.rs @@ -64,9 +64,9 @@ pub struct PixivDownloaderMakeSvc { } impl PixivDownloaderMakeSvc { - pub fn new() -> Self { + pub async fn new() -> Self { Self { - context: Arc::new(ServerContext::default()), + context: Arc::new(ServerContext::default().await), routes: Arc::new(ServerRoutes::new()), } } @@ -90,8 +90,8 @@ impl Service for PixivDownloaderMakeSvc { } /// Start the server -pub fn start_server( +pub async fn start_server( addr: &SocketAddr, ) -> Result, hyper::Error> { - Ok(Server::try_bind(addr)?.serve(PixivDownloaderMakeSvc::new())) + Ok(Server::try_bind(addr)?.serve(PixivDownloaderMakeSvc::new().await)) }