157 lines
5.1 KiB
Python
157 lines
5.1 KiB
Python
# src/main.py
|
|
from fastapi import FastAPI, HTTPException, status, Request
|
|
from fastapi.responses import JSONResponse, PlainTextResponse
|
|
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__)
|
|
|
|
import db_setup as db # Import the entire db_setup module
|
|
from user_mgmt import *
|
|
from media_mgmt import *
|
|
|
|
pathRoot = os.environ.get("ROOT_PATH","/")
|
|
|
|
app = FastAPI(
|
|
title="YT2MP3 API",
|
|
description="Simple API to download YouTube videos as MP3.",
|
|
version="1.0.0",
|
|
license_info="MIT",
|
|
root_path=pathRoot,
|
|
lifespan=db.lifespan
|
|
)
|
|
|
|
# Global error handler
|
|
@app.exception_handler(Exception)
|
|
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)
|
|
|
|
# Health check endpoint
|
|
@app.get("/health", status_code=status.HTTP_200_OK)
|
|
async def health_check():
|
|
return {"status": "ok"}
|
|
|
|
#
|
|
# User-related endpoints
|
|
#
|
|
@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")
|
|
if user is None: raise HTTPException(404, "User not found")
|
|
else:
|
|
return user
|
|
|
|
@app.post("/login")
|
|
async def login_endpoint(req: Request):
|
|
body = await req.json()
|
|
needed_keys = ("email","password")
|
|
for key in needed_keys:
|
|
if body.get(key) is None:
|
|
raise HTTPException(400, "Missing args")
|
|
|
|
login_obj = await 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(400, "User already exists")
|
|
|
|
await create_user(body["first_name"], body["last_name"], body["email"], body["password"])
|
|
return PlainTextResponse("User registered")
|
|
|
|
@app.post("/edit_user")
|
|
async def edit_endpoint(req: Request):
|
|
body:dict = await req.json()
|
|
print(body)
|
|
|
|
if "user_id" not in body.keys(): raise HTTPException(400, "Provide user ID")
|
|
|
|
params = copy.deepcopy(body) # to create a completely new version (ridiculous but okay)
|
|
params.pop("user_id")
|
|
|
|
await edit_user(body["user_id"], params)
|
|
return PlainTextResponse("User edited")
|
|
|
|
@app.post("/changepw")
|
|
async def changepw_endpoint(req: Request):
|
|
body:dict = await req.json()
|
|
print(body)
|
|
|
|
if "user_id" not in body.keys(): raise HTTPException(400, "Provide user ID")
|
|
|
|
params = copy.deepcopy(body) # to create a completely new version (ridiculous but okay)
|
|
params.pop("user_id")
|
|
params.pop("old_password")
|
|
|
|
user = await fetch_user_info(body["user_id"],"id")
|
|
|
|
if bcrypt.checkpw(body["old_password"].encode(), user["password"].encode()) == False:
|
|
raise HTTPException(403, "Incorrect password")
|
|
|
|
|
|
|
|
await edit_user(body["user_id"], params)
|
|
return PlainTextResponse("Password changed")
|
|
#
|
|
# 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", 256)
|
|
|
|
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:
|
|
file = await asyncio.to_thread(download_audio, video_url, USER_FOLDER, quality)
|
|
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 }
|
|
except KeyError as e:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=f"Missing parameter: {e}")
|
|
except DownloadError as e:
|
|
logger.exception(f"Download error: {e}")
|
|
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")
|
|
|
|
|
|
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) |