diff --git a/src/app.py b/src/app.py index fdb392e..532196a 100644 --- a/src/app.py +++ b/src/app.py @@ -2,8 +2,8 @@ from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware from routes import admin_router, score_router, user_router, card_router, movie_router -from models import * -from lib.database import init_db + +from lib.database.database import init_db origins = [ diff --git a/src/lib/database/database.py b/src/lib/database/database.py new file mode 100644 index 0000000..00b04e3 --- /dev/null +++ b/src/lib/database/database.py @@ -0,0 +1,46 @@ +from contextlib import contextmanager +import duckdb +import bcrypt +from rich import print + + +from lib import settings +from models.user import UserIn, UserOut + + +@contextmanager +def get_db(): + try: + conn = duckdb.connect(database=settings.db_url) + yield conn + finally: + conn.close() + + +def db_run(sql): + with get_db() as conn: + conn.execute(sql) + + +def init_db(): + if settings.db_url.exists(): + # TODO: Development feature. Remove this in production + from os import unlink + + unlink(settings.db_url) + # return + sql = "" + with open("lib/database/database.sql") as f: + sql = f.read() + + db_run(sql) + + # Initialize default user + + hashed_password = bcrypt.hashpw("password".encode("utf-8"), bcrypt.gensalt()) + + sql = f"INSERT INTO users (username, password, is_admin) VALUES ('admin', '{hashed_password.decode('utf-8')}', true)" + db_run(sql) + + sql = f"INSERT INTO users (username, password, is_admin) VALUES ('test', '{hashed_password.decode('utf-8')}', false)" + db_run(sql) diff --git a/src/lib/database.sql b/src/lib/database/database.sql similarity index 100% rename from src/lib/database.sql rename to src/lib/database/database.sql diff --git a/src/lib/database.py b/src/lib/database/user.py similarity index 55% rename from src/lib/database.py rename to src/lib/database/user.py index d300b2e..efa1753 100644 --- a/src/lib/database.py +++ b/src/lib/database/user.py @@ -1,49 +1,5 @@ -from contextlib import contextmanager -import duckdb -import bcrypt -from rich import print - - -from lib import settings from models.user import UserIn, UserOut - - -@contextmanager -def get_db(): - try: - conn = duckdb.connect(database=settings.db_url) - yield conn - finally: - conn.close() - - -def db_run(sql): - with get_db() as conn: - conn.execute(sql) - - -def init_db(): - if settings.db_url.exists(): - # TODO: Development feature. Remove this in production - from os import unlink - - unlink(settings.db_url) - # return - sql = "" - with open("lib/database.sql") as f: - sql = f.read() - - db_run(sql) - - # Initialize default user - - hashed_password = bcrypt.hashpw("password".encode("utf-8"), bcrypt.gensalt()) - - sql = f"INSERT INTO users (username, password, is_admin) VALUES ('admin', '{hashed_password.decode('utf-8')}', true)" - db_run(sql) - - sql = f"INSERT INTO users (username, password, is_admin) VALUES ('test', '{hashed_password.decode('utf-8')}', false)" - db_run(sql) +from lib.database.database import get_db def get_user_by_username(username) -> UserIn | None: diff --git a/src/routes/user.py b/src/routes/user.py index d48fc50..da572ea 100644 --- a/src/routes/user.py +++ b/src/routes/user.py @@ -3,7 +3,7 @@ import bcrypt from models.user import UserIn, UserOut, UserCredentials from models.general import Message -from lib.database import get_all_users, get_user_by_username +from lib.database.user import get_all_users, get_user_by_username router = APIRouter(