re-organized backend

This commit is contained in:
2026-01-19 21:42:34 +00:00
parent 6b08f55983
commit 8ecadcf2e5
7 changed files with 581 additions and 221 deletions

View File

@@ -0,0 +1,63 @@
from fastapi import Request, Response, HTTPException
from fastapi.responses import JSONResponse, PlainTextResponse
import aiosqlite
from user_mgmt import *
import db_setup as db
async def register_endpoint(req: Request):
body = await req.json()
try:
create_user(body["first_name"], body["last_name"], body["email"], body["password"])
except KeyError:
raise HTTPException(400, "Missing args")
async def login_endpoint(req: Request):
"""
The login endpoint
Takes email and password
Returns 401 or 400 in error.
Returns the user's full info if correct
"""
body = await req.json()
#print(body)
try:
user_data = await fetch_user_info(body["email"], "email")
if (user_data is None):
print("User doesn't exist")
raise HTTPException(status_code=401, detail="User doesn't exist")
else:
pwd = user_data.get("password")
if (bcrypt.checkpw(body["password"].encode(), user_data.get("password").encode())):
return user_data
else:
print("Wrong password")
raise HTTPException(status_code=401, detail="Wrong password")
except KeyError:
raise HTTPException(status_code=400)
async def me_get_endpoint(req: Request):
"""
Retrieves user profile data.
Trusts Nuxt's backend to allow it or not.
"""
args = dict(req.query_params)
me = None
try:
# Prefer id over e-mail
if "id" in args.keys(): me = fetch_user_info(args["id"], "id")
elif "email" in args.keys(): me = fetch_user_info(args["email"], "email")
else: raise HTTPException(400, "Provide either email or user id")
#print(args)
if (me is None):
raise HTTPException(404, "User doesn't exist")
else: return me
except KeyError:
raise HTTPException(400, "Please provide email")

View File

@@ -0,0 +1,31 @@
import aiomysql, aiosqlite, os
from fastapi import FastAPI
from contextlib import asynccontextmanager
# Either "mysql" or "sqlite"
mode = os.environ.get("DB_MODE", "sqlite")
# For SQLite usage
sqlite_file = os.environ.get("SQLITE_FILE","./database.db")
# For MySQL usage
mysql_pool = None
mysql_host = os.environ.get("MYSQL_HOST","127.0.0.1")
mysql_user = os.environ.get("MYSQL_USER","yt2mp3")
mysql_db_name = os.environ.get("MYSQL_DB","yt2mp3")
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:
yield
finally:
pool.close()
await pool.wait_closed()
else:
try:
yield
finally: pass # no action required now

View File

@@ -1,222 +1,124 @@
from fastapi import FastAPI, HTTPException, Request, Response # src/main.py
#from fastapi import Body, Cookie, File, Form, Header, Path, Query from fastapi import FastAPI, HTTPException, status, Request
from contextlib import asynccontextmanager from fastapi.responses import JSONResponse, PlainTextResponse
from pydantic import BaseModel, ValidationError
from typing import Optional
import uvicorn
import aiomysql
import os import os
import bcrypt import bcrypt
import asyncio import asyncio
import aiosqlite import logging
import subprocess
from yt_dlp import YoutubeDL # Configure logging
from yt_dlp.utils import DownloadError logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
#from pytube import YouTube from . import db_setup as db # Import the entire db_setup module
#from moviepy.audio.io import AudioFileClip from .user_mgmt import *
from .media_mgmt import *
#import asyncio app = FastAPI(
title="YT2MP3 API",
description="Simple API to download YouTube videos as MP3.",
version="1.0.0",
license_info="MIT",
lifespan=db.lifespan
)
# In this FastAPI example, we skip BaseModel verification cause NuxtJS does that for us # Global error handler
# In other words, this backend only recieves requests from nuxt js backend @app.exception_handler(Exception)
# So we can trust it async def global_exception_handler(request, exc):
logger.exception(f"An unexpected error occurred: {exc}")
return JSONResponse(content={"error": "Internal Server Error"}, status_code=500)
class db: # Health check endpoint
@app.get("/health", status_code=status.HTTP_200_OK)
async def health_check():
return {"status": "ok"}
mode = os.environ.get("DB_MODE", "sqlite") #
# User-related endpoints
# For SQLite usage #
@app.get("/me")
file = os.environ.get("SQLITE_FILE","./database.db") async def me_endpoint(req: Request):
query = dict(req.query_params)
# For MySQL usage if "id" not in query.keys(): raise HTTPException(400, "Give ID")
pool = None
user = await fetch_user_info(query["id"],"id")
host = os.environ.get("MYSQL_HOST","127.0.0.1") if user is None: raise HTTPException(404, "User not found")
user = os.environ.get("MYSQL_USER","yt2mp3")
db_name = os.environ.get("MYSQL_DB","yt2mp3")
pwd = os.environ.get("MYSQL_PASSWORD","")
@asynccontextmanager
async def lifespan(app: FastAPI):
if db.mode == "mysql":
db.pool = await aiomysql.create_pool(host=db.host,user=db.user,password=db.pwd,db=db.db_name, minsize=1, maxsize=10)
try:
yield
finally:
db.pool.close()
await db.pool.wait_closed()
else: else:
try: return user
yield
finally: pass # no action required now
app = FastAPI(lifespan=lifespan)
async def fetch_user_info(arg: str, type: str):
if db.mode == "mysql":
async with db.pool.acquire() as conn:
async with conn.cursor(aiomysql.DictCursor) as cur:
await cur.execute("SELECT * FROM users WHERE %s = %s", (type, arg))
return await cur.fetchone()
elif db.mode == "sqlite":
async with aiosqlite.connect(db.file) as conn:
async with conn.execute("SELECT * FROM users WHERE ? = ?", (type, arg)) as cursor:
return await cursor.fetchone()
def download_audio(video_url: str, output_folder: str) -> str:
ydl_opts = {
'format': 'bestaudio/best',
'outtmpl': os.path.join(output_folder, '%(title)s.%(ext)s'),
'postprocessors': [{
'key': 'FFmpegExtractAudio',
'preferredcodec': 'mp3',
'preferredquality': '192',
}, {
'key': 'FFmpegMetadata',
'add_metadata': True,
}],
'quiet': True,
'noplaylist': True
}
with YoutubeDL(ydl_opts) as ydl:
info = ydl.extract_info(video_url, download=True)
filename = ydl.prepare_filename(info)
return os.path.splitext(filename)[0] + ".mp3" # path to final MP3
def set_metadata_and_rename(mp3_file: str, name: str, artist: str = None, album: str = None) -> str:
new_file = os.path.abspath(os.path.join(os.path.dirname(mp3_file), f"{name}.mp3"))
cmd = ["ffmpeg", "-y", "-i", mp3_file]
# Only add metadata if provided
cmd += ["-metadata", f"title={name}"]
if artist:
cmd += ["-metadata", f"artist={artist}"]
if album:
cmd += ["-metadata", f"album={album}"]
cmd += [new_file]
subprocess.run(cmd, check=True)
os.remove(mp3_file) # clean original
return new_file
# Configuration (customize these!)
DOWNLOAD_FOLDER = os.environ.get("DOWNLOAD_FOLDER","downloads") # Where the converted files will be saved
if not os.path.exists(DOWNLOAD_FOLDER):
os.makedirs(DOWNLOAD_FOLDER)
# Throw error if the env is incorrect
try:
AUTODELETE_DELAY_SECS = int(os.environ.get("FILE_AUTODELETE_DELAY_SECS", 0))
except ValueError:
raise EnvironmentError("You must set AUTODELETE_DELAY_SECS to a valid number (int) of seconds! To disable auto delete (not recommended), set it to 0")
async def delay_delete_file(file:str):
if (AUTODELETE_DELAY_SECS > 0):
await asyncio.sleep(AUTODELETE_DELAY_SECS)
if os.path.exists(file): os.unlink(file)
@app.post("/convert")
async def create_conversion(req: Request):
"""Converts a YouTube video to MP3."""
body = await req.json()
print(body)
video_url = body.get("video_url")
user_id = body.get("user_id")
if (video_url is None): raise HTTPException(400, "Please provide an URL on body")
if (body.get("user_id") is None): raise HTTPException(401, "No User ID")
USER_FOLDER = os.path.join(DOWNLOAD_FOLDER, str(user_id))
if not os.path.exists(USER_FOLDER): os.mkdir(USER_FOLDER)
try:
file = await asyncio.to_thread(download_audio, video_url, USER_FOLDER)
except DownloadError as e:
raise HTTPException(500, "Error while downloading " + e.msg)
final_file = await asyncio.to_thread(set_metadata_and_rename, file, body.get("song"), body.get("artist"), body.get("album"))
file_name = final_file.split("/")[-1]
print("File served: " + final_file)
asyncio.create_task(delay_delete_file(final_file))
return { "file" : file_name }
@app.post("/register")
async def create_user(req: Request):
body = await req.json()
# Auto exceptions give more verbose to Nuxt.
body["password"] = bcrypt.hashpw(body["password"].encode(), bcrypt.gensalt()).decode()
if db.mode == "mysql":
async with db.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)", (body.first_name, body.last_name, body.email, body.password))
await conn.commit()
elif db.mode == "sqlite":
async with aiosqlite.connect(db.file) as conn:
await conn.execute("INSERT INTO users (`first_name`,`last_name`,`email`,`password`) VALUES (?, ?, ?, ?)", (body.first_name, body.last_name, body.email, body.password))
await conn.commit()
@app.post("/login") @app.post("/login")
async def login_endpoint(req: Request): async def login_endpoint(req: Request):
body = await req.json() body = await req.json()
#print(body) needed_keys = ("email","password")
for key in needed_keys:
if body.get(key) is None:
raise HTTPException(400, "Missing args")
login_obj = login_user(body["email"], body["password"])
if login_obj is None:
raise HTTPException(401, "Incorrect login")
else:
return login_obj
@app.post("/register")
async def register_endpoint(req: Request):
body = await req.json()
for key in ("email","password", "first_name", "last_name"):
if body.get(key) is None:
raise HTTPException(400, "Missing args")
# check if user exists.
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"])
return PlainTextResponse("User registered")
@app.put("/edit_user")
async def edit_endpoint(req: Request):
body:dict = await req.json()
if "user_id" not in body.keys(): raise HTTPException(400, "Provide user ID")
params = body
params.pop("user_id")
edit_user(body["user_id"], params)
return PlainTextResponse("User edited")
#
# Track-specific endpoints
#
@app.post("/convert")
async def begin_conversion(req: Request):
body:dict = await req.json()
video_url = body.get("video_url")
user_id = body.get("user_id")
quality = body.get("kbps")
if (video_url is None): raise HTTPException(400, "Please provide an URL on body")
if (user_id is None): raise HTTPException(401, "No User ID")
# User folder is DOWNLOAD_FOLDER/<userid>
USER_FOLDER = os.path.join(DOWNLOAD_FOLDER, str(user_id))
if not os.path.exists(USER_FOLDER): os.mkdir(USER_FOLDER)
try: try:
user_data = await fetch_user_info(body["email"], "email") file = await asyncio.to_thread(download_audio, video_url, USER_FOLDER, quality)
if (user_data is None): final_file = await asyncio.to_thread(set_metadata_and_rename, file, body.get("song"), body.get("artist"), body.get("album"))
print("User doesn't exist") file_name = final_file.split("/")[-1]
raise HTTPException(status_code=401, detail="User doesn't exist")
else:
pwd = user_data.get("password")
if (bcrypt.checkpw(body["password"].encode(), user_data.get("password").encode())):
return user_data
else:
print("Wrong password")
raise HTTPException(status_code=401, detail="Wrong password")
except KeyError:
raise HTTPException(status_code=400)
#except ValidationError:
# raise HTTPException(status_code=400)
print("File served: " + final_file)
asyncio.create_task(delay_delete_file(final_file))
@app.get("/me") return { "file" : file_name }
async def me_get_endpoint(req: Request): except KeyError as e:
""" raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=f"Missing parameter: {e}")
Retrieves user profile data. except DownloadError as e:
logger.exception(f"Download error: {e}")
Trusts Nuxt's backend to allow it or not. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to download video")
""" except Exception as e:
args = dict(req.query_params) logger.exception(f"Error during media conversion: {e}")
me = None raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to convert media")
try:
# Prefer id over e-mail
if "id" in args.keys(): me = fetch_user_info(args["id"], "id")
elif "email" in args.keys(): me = fetch_user_info(args["email"], "email")
else: raise HTTPException(400, "Provide either email or user id")
#print(args)
if (me is None):
raise HTTPException(404, "User doesn't exist")
else: return me
except KeyError:
raise HTTPException(400, "Please provide email")
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

@@ -0,0 +1,234 @@
from fastapi import FastAPI, HTTPException, Request, Response
#from fastapi import Body, Cookie, File, Form, Header, Path, Query
from contextlib import asynccontextmanager
from pydantic import BaseModel, ValidationError
from typing import Optional
import uvicorn
import aiomysql
import os
import bcrypt
import asyncio
import aiosqlite
import subprocess
from yt_dlp import YoutubeDL
from yt_dlp.utils import DownloadError
class db:
mode = os.environ.get("DB_MODE", "sqlite")
# For SQLite usage
file = os.environ.get("SQLITE_FILE","./database.db")
# For MySQL usage
pool = None
host = os.environ.get("MYSQL_HOST","127.0.0.1")
user = os.environ.get("MYSQL_USER","yt2mp3")
db_name = os.environ.get("MYSQL_DB","yt2mp3")
pwd = os.environ.get("MYSQL_PASSWORD","")
@asynccontextmanager
async def lifespan(app: FastAPI):
if db.mode == "mysql":
db.pool = await aiomysql.create_pool(host=db.host,user=db.user,password=db.pwd,db=db.db_name, minsize=1, maxsize=10)
try:
yield
finally:
db.pool.close()
await db.pool.wait_closed()
else:
try:
yield
finally: pass # no action required now
app = FastAPI(lifespan=lifespan)
#
# Fetch user info.
# Returns a JSON of the user or None if nothing exists (like an incorrect email)
#
async def fetch_user_info(arg: str, type: str):
if db.mode == "mysql":
async with db.pool.acquire() as conn:
async with conn.cursor(aiomysql.DictCursor) as cur:
await cur.execute("SELECT * FROM users WHERE %s = %s", (type, arg))
return await cur.fetchone()
elif db.mode == "sqlite":
async with aiosqlite.connect(db.file) as conn:
async with conn.execute("SELECT * FROM users WHERE ? = ?", (type, arg)) as cursor:
return await cursor.fetchone()
#
# Downloads the video and extracts the audio
#
def download_audio(video_url: str, output_folder: str, kbps:int) -> str:
ydl_opts = {
'format': 'bestaudio/best',
'outtmpl': os.path.join(output_folder, '%(title)s.%(ext)s'),
'postprocessors': [{
'key': 'FFmpegExtractAudio',
'preferredcodec': 'mp3',
'preferredquality': str(kbps),
}, {
'key': 'FFmpegMetadata',
'add_metadata': True,
}],
'quiet': True,
'noplaylist': True
}
with YoutubeDL(ydl_opts) as ydl:
info = ydl.extract_info(video_url, download=True)
filename = ydl.prepare_filename(info)
return os.path.splitext(filename)[0] + ".mp3" # path to final MP3
#
# Uses FFmpeg to edit the artist and album if specified
#
def set_metadata_and_rename(mp3_file: str, name: str, artist: str = None, album: str = None) -> str:
new_file = os.path.abspath(os.path.join(os.path.dirname(mp3_file), f"{name}.mp3"))
cmd = ["ffmpeg", "-y", "-i", mp3_file]
# Only add metadata if provided
cmd += ["-metadata", f"title={name}"]
if artist:
cmd += ["-metadata", f"artist={artist}"]
if album:
cmd += ["-metadata", f"album={album}"]
cmd += [new_file]
subprocess.run(cmd, check=True)
os.remove(mp3_file) # clean original
return new_file
#
# The download folder.
# It goes:
# FOLDER/<user_id>/<filename>
#
DOWNLOAD_FOLDER = os.environ.get("DOWNLOAD_FOLDER","/tmp/downloads") # Where the converted files will be saved
if not os.path.exists(DOWNLOAD_FOLDER):
os.makedirs(DOWNLOAD_FOLDER)
# Throw error if the env is incorrect
try:
AUTODELETE_DELAY_SECS = int(os.environ.get("FILE_AUTODELETE_DELAY_SECS", 0))
except ValueError:
raise EnvironmentError("You must set AUTODELETE_DELAY_SECS to a valid number (int) of seconds! To disable auto delete (not recommended), set it to 0")
#
# Asyncioway of scheduling a file delete with asyncio.sleep() and os.unlink()
#
async def delay_delete_file(file:str):
if (AUTODELETE_DELAY_SECS > 0):
await asyncio.sleep(AUTODELETE_DELAY_SECS)
if os.path.exists(file): os.unlink(file)
@app.post("/convert")
async def create_conversion(req: Request):
"""Converts a YouTube video to MP3."""
body = await req.json()
print(body)
video_url = body.get("video_url")
user_id = body.get("user_id")
quality = body.get("kbps")
if (video_url is None): raise HTTPException(400, "Please provide an URL on body")
if (body.get("user_id") is None): raise HTTPException(401, "No User ID")
# User folder is DOWNLOAD_FOLDER/<userid>
USER_FOLDER = os.path.join(DOWNLOAD_FOLDER, str(user_id))
if not os.path.exists(USER_FOLDER): os.mkdir(USER_FOLDER)
try:
file = await asyncio.to_thread(download_audio, video_url, USER_FOLDER, quality)
except DownloadError as e:
raise HTTPException(500, "Error while downloading " + e.msg)
final_file = await asyncio.to_thread(set_metadata_and_rename, file, body.get("song"), body.get("artist"), body.get("album"))
file_name = final_file.split("/")[-1]
print("File served: " + final_file)
asyncio.create_task(delay_delete_file(final_file))
return { "file" : file_name }
@app.post("/register")
async def create_user(req: Request):
body = await req.json()
body["password"] = bcrypt.hashpw(body["password"].encode(), bcrypt.gensalt()).decode()
if db.mode == "mysql":
async with db.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)", (body.first_name, body.last_name, body.email, body.password))
await conn.commit()
elif db.mode == "sqlite":
async with aiosqlite.connect(db.file) as conn:
await conn.execute("INSERT INTO users (`first_name`,`last_name`,`email`,`password`) VALUES (?, ?, ?, ?)", (body.first_name, body.last_name, body.email, body.password))
await conn.commit()
@app.post("/login")
async def login_endpoint(req: Request):
body = await req.json()
#print(body)
try:
user_data = await fetch_user_info(body["email"], "email")
if (user_data is None):
print("User doesn't exist")
raise HTTPException(status_code=401, detail="User doesn't exist")
else:
pwd = user_data.get("password")
if (bcrypt.checkpw(body["password"].encode(), user_data.get("password").encode())):
return user_data
else:
print("Wrong password")
raise HTTPException(status_code=401, detail="Wrong password")
except KeyError:
raise HTTPException(status_code=400)
#except ValidationError:
# raise HTTPException(status_code=400)
@app.get("/me")
async def me_get_endpoint(req: Request):
"""
Retrieves user profile data.
Trusts Nuxt's backend to allow it or not.
"""
args = dict(req.query_params)
me = None
try:
# Prefer id over e-mail
if "id" in args.keys(): me = fetch_user_info(args["id"], "id")
elif "email" in args.keys(): me = fetch_user_info(args["email"], "email")
else: raise HTTPException(400, "Provide either email or user id")
#print(args)
if (me is None):
raise HTTPException(404, "User doesn't exist")
else: return me
except KeyError:
raise HTTPException(400, "Please provide email")
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

@@ -0,0 +1,73 @@
import subprocess, os, asyncio
from yt_dlp import YoutubeDL
from yt_dlp.utils import DownloadError
#
# The download folder.
# It goes:
# FOLDER/<user_id>/<filename>
#
DOWNLOAD_FOLDER = os.environ.get("DOWNLOAD_FOLDER","/tmp/downloads") # Where the converted files will be saved
if not os.path.exists(DOWNLOAD_FOLDER):
os.makedirs(DOWNLOAD_FOLDER)
# Throw error if the env is incorrect
try:
AUTODELETE_DELAY_SECS = int(os.environ.get("FILE_AUTODELETE_DELAY_SECS", 0))
except ValueError:
raise EnvironmentError("You must set AUTODELETE_DELAY_SECS to a valid number (int) of seconds! To disable auto delete (not recommended), set it to 0")
#
# Asyncioway of scheduling a file delete with asyncio.sleep() and os.unlink()
#
async def delay_delete_file(file:str):
if (AUTODELETE_DELAY_SECS > 0):
await asyncio.sleep(AUTODELETE_DELAY_SECS)
if os.path.exists(file): os.unlink(file)
#
# Downloads the video and extracts the audio
#
def download_audio(video_url: str, output_folder: str, kbps:int) -> str:
ydl_opts = {
'format': 'bestaudio/best',
'outtmpl': os.path.join(output_folder, '%(title)s.%(ext)s'),
'postprocessors': [{
'key': 'FFmpegExtractAudio',
'preferredcodec': 'mp3',
'preferredquality': str(kbps),
}, {
'key': 'FFmpegMetadata',
'add_metadata': True,
}],
'quiet': True,
'noplaylist': True
}
with YoutubeDL(ydl_opts) as ydl:
info = ydl.extract_info(video_url, download=True)
filename = ydl.prepare_filename(info)
return os.path.splitext(filename)[0] + ".mp3" # path to final MP3
#
# Uses FFmpeg to edit the artist and album if specified
#
def set_metadata_and_rename(mp3_file: str, name: str, artist: str = None, album: str = None) -> str:
new_file = os.path.abspath(os.path.join(os.path.dirname(mp3_file), f"{name}.mp3"))
cmd = ["ffmpeg", "-y", "-i", mp3_file]
# Only add metadata if provided
cmd += ["-metadata", f"title={name}"]
if artist:
cmd += ["-metadata", f"artist={artist}"]
if album:
cmd += ["-metadata", f"album={album}"]
cmd += [new_file]
subprocess.run(cmd, check=True)
os.remove(mp3_file) # clean original
return new_file

View File

@@ -1,16 +0,0 @@
aiomysql==0.3.2
annotated-doc==0.0.4
annotated-types==0.7.0
anyio==4.12.1
bcrypt==5.0.0
click==8.3.1
fastapi==0.128.0
h11==0.16.0
idna==3.11
pydantic==2.12.5
pydantic_core==2.41.5
starlette==0.50.0
typing-inspection==0.4.2
typing_extensions==4.15.0
uvicorn==0.40.0

View File

@@ -0,0 +1,73 @@
# fastapi_server/user_mgmt.py
from . import db_setup as db
import bcrypt
import asyncio
from aiomysql import DictCursor
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))
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:
return await cursor.fetchone()
return None
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:
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:
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
return False
async def login_user(email: str, password: str):
"""Logs in a user."""
user_data = await fetch_user_info(email, "email")
if user_data is None:
return None # User doesn't exist
else:
if bcrypt.checkpw(password.encode(), user_data.get("password").encode()):
return user_data # Login successful
else:
return None # Incorrect password
async def edit_user(user_id: int, updates: dict):
"""Edits an existing user."""
if db.mode == "mysql":
placeholders = ", ".join(["%s"] * len(updates))
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 conn.cursor() as cur:
await cur.execute(query, values)
await conn.commit()
return True
elif db.mode == "sqlite":
placeholders = ", ".join(["?"] * len(updates))
values = list(updates.values())
values.append(user_id)
query = f"UPDATE users SET {', '.join([f'{key} = ?' for key in updates.keys()])} WHERE id = ?"
async with db.file.connect() as conn:
await conn.execute(query, values)
await conn.commit()
return True
return False