login works backend

This commit is contained in:
2026-01-20 01:33:55 +00:00
parent 797a696aff
commit accf35e364
3 changed files with 43 additions and 24 deletions

View File

@@ -19,13 +19,18 @@ mysql_pwd = os.environ.get("MYSQL_PASSWORD","")
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):
if mode == "mysql": 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: 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 yield
finally: finally:
pool.close() mysql_pool.close()
await pool.wait_closed() await mysql_pool.wait_closed()
else: elif mode == "sqlite":
try: try:
print("Going with sqlite")
yield yield
finally: pass # no action required now finally: pass # no action required now
else:
raise EnvironmentError("DB_MODE must either be mysql or sqlite")

View File

@@ -5,14 +5,16 @@ import os
import bcrypt import bcrypt
import asyncio import asyncio
import logging import logging
import uvicorn
import copy
# Configure logging # Configure logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from . import db_setup as db # Import the entire db_setup module import db_setup as db # Import the entire db_setup module
from .user_mgmt import * from user_mgmt import *
from .media_mgmt import * from media_mgmt import *
app = FastAPI( app = FastAPI(
title="YT2MP3 API", title="YT2MP3 API",
@@ -39,6 +41,7 @@ async def health_check():
@app.get("/me") @app.get("/me")
async def me_endpoint(req: Request): async def me_endpoint(req: Request):
query = dict(req.query_params) query = dict(req.query_params)
print(query)
if "id" not in query.keys(): raise HTTPException(400, "Give ID") if "id" not in query.keys(): raise HTTPException(400, "Give ID")
user = await fetch_user_info(query["id"],"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: if body.get(key) is None:
raise HTTPException(400, "Missing args") 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: if login_obj is None:
raise HTTPException(401, "Incorrect login") raise HTTPException(401, "Incorrect login")
else: else:
@@ -72,19 +75,20 @@ async def register_endpoint(req: Request):
if await fetch_user_info(body["email"],"email") is not None: if await fetch_user_info(body["email"],"email") is not None:
raise HTTPException(403, "User already exists") 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") return PlainTextResponse("User registered")
@app.put("/edit_user") @app.put("/edit_user")
async def edit_endpoint(req: Request): async def edit_endpoint(req: Request):
body:dict = await req.json() body:dict = await req.json()
print(body.keys())
if "user_id" not in body.keys(): raise HTTPException(400, "Provide user ID") 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") params.pop("user_id")
edit_user(body["user_id"], params) await edit_user(body["user_id"], params)
return PlainTextResponse("User edited") return PlainTextResponse("User edited")
# #
@@ -122,3 +126,9 @@ async def begin_conversion(req: Request):
except Exception as e: except Exception as e:
logger.exception(f"Error during media conversion: {e}") logger.exception(f"Error during media conversion: {e}")
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to convert media") 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)

View File

@@ -1,21 +1,24 @@
# fastapi_server/user_mgmt.py # fastapi_server/user_mgmt.py
from . import db_setup as db import db_setup as db
import bcrypt import bcrypt
import asyncio import asyncio
from aiomysql import DictCursor import aiomysql
import db_setup as db
async def fetch_user_info(arg: str, type: str): async def fetch_user_info(arg: str, type: str):
"""Fetches user info. """Fetches user info.
This kind of query is supposed to either return ONE user or nothing. This kind of query is supposed to either return ONE user or nothing.
""" """
if db.mode == "mysql": if db.mode == "mysql":
async with db.pool.acquire() as conn: async with db.mysql_pool.acquire() as conn:
async with conn.cursor(DictCursor) as cur: async with conn.cursor(aiomysql.DictCursor) as cur:
await cur.execute("SELECT * FROM users WHERE %s = %s", (type, arg)) await cur.execute(f"SELECT * FROM users WHERE {type} = %s", (arg,))
#print(cur.description)
return await cur.fetchone() return await cur.fetchone()
elif db.mode == "sqlite": elif db.mode == "sqlite":
async with db.file.connect() as conn: 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 await cursor.fetchone()
return None 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): async def create_user(first_name: str, last_name: str, email: str, password: str):
"""Creates a new user.""" """Creates a new user."""
hashed_password = bcrypt.hashpw(password.encode(), bcrypt.gensalt()) hashed_password = bcrypt.hashpw(password.encode(), bcrypt.gensalt())
if db.db.mode == "mysql": if db.mode == "mysql":
async with db.db.pool.acquire() as conn: async with db.mysql_pool.acquire() as conn:
async with conn.cursor() as cur: 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 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() await conn.commit()
return True return True
elif db.db.mode == "sqlite": elif db.mode == "sqlite":
async with db.db.file.connect() as conn: 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.execute("INSERT INTO users (`first_name`,`last_name`,`email`,`password`) VALUES (?, ?, ?, ?)", (first_name, last_name, email, hashed_password.decode()))
await conn.commit() await conn.commit()
return True 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): async def login_user(email: str, password: str):
"""Logs in a user.""" """Logs in a user."""
user_data = await fetch_user_info(email, "email") user_data = await fetch_user_info(email, "email")
print(user_data)
if user_data is None: if user_data is None:
return None # User doesn't exist return None # User doesn't exist
else: else:
@@ -56,7 +60,7 @@ async def edit_user(user_id: int, updates: dict):
values = list(updates.values()) values = list(updates.values())
query = f"UPDATE users SET {', '.join([f'{key} = %s' for key in updates.keys()])} WHERE id = %s" query = f"UPDATE users SET {', '.join([f'{key} = %s' for key in updates.keys()])} WHERE id = %s"
values.append(user_id) 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: async with conn.cursor() as cur:
await cur.execute(query, values) await cur.execute(query, values)
await conn.commit() await conn.commit()