diff --git a/server/py/fastapi_server/db_setup.py b/server/py/fastapi_server/db_setup.py index baef295..1bee33e 100644 --- a/server/py/fastapi_server/db_setup.py +++ b/server/py/fastapi_server/db_setup.py @@ -19,13 +19,18 @@ mysql_pwd = os.environ.get("MYSQL_PASSWORD","") @asynccontextmanager async def lifespan(app: FastAPI): if mode == "mysql": - pool = await aiomysql.create_pool(host=mysql_host,user=mysql_user,password=mysql_pwd,db=mysql_db_name, minsize=1, maxsize=10) try: + print("Going with mysql") + global mysql_pool + mysql_pool = await aiomysql.create_pool(host=mysql_host,user=mysql_user,password=mysql_pwd,db=mysql_db_name, minsize=1, maxsize=10) yield finally: - pool.close() - await pool.wait_closed() - else: + mysql_pool.close() + await mysql_pool.wait_closed() + elif mode == "sqlite": try: + print("Going with sqlite") yield - finally: pass # no action required now \ No newline at end of file + finally: pass # no action required now + else: + raise EnvironmentError("DB_MODE must either be mysql or sqlite") \ No newline at end of file diff --git a/server/py/fastapi_server/main.py b/server/py/fastapi_server/main.py index 9cbc6f2..476751e 100644 --- a/server/py/fastapi_server/main.py +++ b/server/py/fastapi_server/main.py @@ -5,14 +5,16 @@ import os import bcrypt import asyncio import logging +import uvicorn +import copy # Configure logging logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') logger = logging.getLogger(__name__) -from . import db_setup as db # Import the entire db_setup module -from .user_mgmt import * -from .media_mgmt import * +import db_setup as db # Import the entire db_setup module +from user_mgmt import * +from media_mgmt import * app = FastAPI( title="YT2MP3 API", @@ -39,6 +41,7 @@ async def health_check(): @app.get("/me") async def me_endpoint(req: Request): query = dict(req.query_params) + print(query) if "id" not in query.keys(): raise HTTPException(400, "Give ID") user = await fetch_user_info(query["id"],"id") @@ -54,7 +57,7 @@ async def login_endpoint(req: Request): if body.get(key) is None: raise HTTPException(400, "Missing args") - login_obj = login_user(body["email"], body["password"]) + login_obj = await login_user(body["email"], body["password"]) if login_obj is None: raise HTTPException(401, "Incorrect login") else: @@ -72,19 +75,20 @@ async def register_endpoint(req: Request): if await fetch_user_info(body["email"],"email") is not None: raise HTTPException(403, "User already exists") - create_user(body["fist_name"], body["last_name"], body["email"], body["password"]) + await create_user(body["first_name"], body["last_name"], body["email"], body["password"]) return PlainTextResponse("User registered") @app.put("/edit_user") async def edit_endpoint(req: Request): body:dict = await req.json() + print(body.keys()) if "user_id" not in body.keys(): raise HTTPException(400, "Provide user ID") - params = body + params = copy.deepcopy(body) # to create a completely new version (ridiculous but okay) params.pop("user_id") - edit_user(body["user_id"], params) + await edit_user(body["user_id"], params) return PlainTextResponse("User edited") # @@ -121,4 +125,10 @@ async def begin_conversion(req: Request): raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to download video") except Exception as e: logger.exception(f"Error during media conversion: {e}") - raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to convert media") \ No newline at end of file + raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to convert media") + + +if __name__ == "__main__": + listen_host = str(os.environ.get("HOST", "127.0.0.1")) + listen_port = int(os.environ.get("PORT", 3000)) + uvicorn.run(app,host=listen_host,port=listen_port) \ No newline at end of file diff --git a/server/py/fastapi_server/user_mgmt.py b/server/py/fastapi_server/user_mgmt.py index 48db11d..263d61e 100644 --- a/server/py/fastapi_server/user_mgmt.py +++ b/server/py/fastapi_server/user_mgmt.py @@ -1,21 +1,24 @@ # fastapi_server/user_mgmt.py -from . import db_setup as db +import db_setup as db import bcrypt import asyncio -from aiomysql import DictCursor +import aiomysql + +import db_setup as db async def fetch_user_info(arg: str, type: str): """Fetches user info. This kind of query is supposed to either return ONE user or nothing. """ if db.mode == "mysql": - async with db.pool.acquire() as conn: - async with conn.cursor(DictCursor) as cur: - await cur.execute("SELECT * FROM users WHERE %s = %s", (type, arg)) + async with db.mysql_pool.acquire() as conn: + async with conn.cursor(aiomysql.DictCursor) as cur: + await cur.execute(f"SELECT * FROM users WHERE {type} = %s", (arg,)) + #print(cur.description) return await cur.fetchone() elif db.mode == "sqlite": async with db.file.connect() as conn: - async with conn.execute("SELECT * FROM users WHERE ? = ?", (type, arg)) as cursor: + async with conn.execute(f"SELECT * FROM users WHERE {type} = ?", (arg,)) as cursor: return await cursor.fetchone() return None @@ -23,14 +26,14 @@ async def fetch_user_info(arg: str, type: str): async def create_user(first_name: str, last_name: str, email: str, password: str): """Creates a new user.""" hashed_password = bcrypt.hashpw(password.encode(), bcrypt.gensalt()) - if db.db.mode == "mysql": - async with db.db.pool.acquire() as conn: + if db.mode == "mysql": + async with db.mysql_pool.acquire() as conn: async with conn.cursor() as cur: await cur.execute("INSERT INTO users (`first_name`,`last_name`,`email`,`password`) VALUES (%s,%s,%s,%s)", (first_name, last_name, email, hashed_password.decode())) await conn.commit() return True - elif db.db.mode == "sqlite": - async with db.db.file.connect() as conn: + elif db.mode == "sqlite": + async with db.file.connect() as conn: await conn.execute("INSERT INTO users (`first_name`,`last_name`,`email`,`password`) VALUES (?, ?, ?, ?)", (first_name, last_name, email, hashed_password.decode())) await conn.commit() return True @@ -40,6 +43,7 @@ async def create_user(first_name: str, last_name: str, email: str, password: str async def login_user(email: str, password: str): """Logs in a user.""" user_data = await fetch_user_info(email, "email") + print(user_data) if user_data is None: return None # User doesn't exist else: @@ -56,7 +60,7 @@ async def edit_user(user_id: int, updates: dict): values = list(updates.values()) query = f"UPDATE users SET {', '.join([f'{key} = %s' for key in updates.keys()])} WHERE id = %s" values.append(user_id) - async with db.pool.acquire() as conn: + async with db.mysql_pool.acquire() as conn: async with conn.cursor() as cur: await cur.execute(query, values) await conn.commit()