login works backend
This commit is contained in:
@@ -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")
|
||||||
@@ -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)
|
||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user