diff --git a/app.py b/app.py index dc66acc..527df79 100644 --- a/app.py +++ b/app.py @@ -1,5 +1,3 @@ -# app.py - import os import base64 import datetime @@ -12,7 +10,7 @@ import hashlib import smtplib import mimetypes import json -import logging # Import the logging module +import logging from email.mime.text import MIMEText from pyramid.config import Configurator @@ -28,19 +26,21 @@ from sqlalchemy import ( Boolean, Integer, Index, - func, ) from sqlalchemy.orm import declarative_base from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import scoped_session +from sqlalchemy.pool import StaticPool from waitress import serve -# Import renderers for templates from pyramid.renderers import render_to_response -# Import for events from pyramid.events import subscriber from pyramid_jinja2 import IJinja2Environment +import transaction # Import transaction management +from zope.sqlalchemy import register # Import zope.sqlalchemy + ################################################################################ # Set up logging ################################################################################ @@ -51,17 +51,20 @@ log = logging.getLogger(__name__) # Helper Functions ################################################################################ + def slugify(text): text = text.lower() text = re.sub(r"\s+", "-", text) text = re.sub(r"[^\w\-]", "", text) return text + def get_gravatar_url(email, size=100): email = email.strip().lower() hash_code = hashlib.md5(email.encode("utf-8")).hexdigest() return f"https://www.gravatar.com/avatar/{hash_code}?s={size}&d=identicon" + def send_email(to_email, subject, body): # For testing purposes, print the email content to the console log.debug("======= Email Sent =======") @@ -82,14 +85,17 @@ def send_email(to_email, subject, body): except Exception as e: log.error(f"Error sending email: {e}") + def admin_required(view_func): def wrapper(request): user = request.user if not user or not user.is_admin: return HTTPForbidden("You must be an admin to access this page.") return view_func(request) + return wrapper + def get_current_user(request): """Return the current user (authenticated or guest) from session.""" user_id = request.session.get("user_id") @@ -118,7 +124,7 @@ def get_current_user(request): user_id = str(user_uuid) short_id = uuid_to_short_id(user_uuid) - # Create a new guest user + # Create a new guest user (do not create user database) guest_user = User( id=user_id, short_id=short_id, @@ -127,13 +133,14 @@ def get_current_user(request): is_verified=False, ) s.add(guest_user) - s.commit() + s.flush() # Use flush instead of commit in pyramid_tm # Store the user ID in the session request.session["user_id"] = guest_user.id return guest_user + def get_mime_type(filename): # Guess the MIME type based on the file extension mime_type, _ = mimetypes.guess_type(filename) @@ -146,6 +153,7 @@ def uuid_to_short_id(u): """Encode UUID to a URL-safe base64 string without padding.""" return base64.urlsafe_b64encode(u.bytes).decode("ascii").rstrip("=") + def short_id_to_uuid(sid): """Decode the short ID back to UUID, trying different padding lengths.""" # Try with 0 to 3 padding characters @@ -157,17 +165,21 @@ def short_id_to_uuid(sid): return uuid.UUID(bytes=bytes_data) except Exception as e: continue - + # If we get here, none of the padding attempts worked - log.error(f"Failed to convert short_id {sid} to UUID after trying all padding lengths") + log.error( + f"Failed to convert short_id {sid} to UUID after trying all padding lengths" + ) return None + def get_user_db_url(user_id): """Return the database URL for the user's SQLite database.""" APP_DIR = os.path.dirname(os.path.abspath(__file__)) db_file = os.path.join(APP_DIR, f"user_{user_id}.db") return f"sqlite:///{db_file}" + def filesizeformat(value): """Returns the human-readable file size.""" for unit in ["bytes", "KB", "MB", "GB", "TB"]: @@ -176,6 +188,7 @@ def filesizeformat(value): value /= 1024.0 return f"{value:.2f} PB" + ################################################################################ # Database Setup ################################################################################ @@ -186,6 +199,7 @@ log.debug(f"Using database URL: {DB_URL}") # For debugging Base = declarative_base() + class User(Base): __tablename__ = "users" id = Column(String, primary_key=True) # UUID @@ -207,6 +221,7 @@ class User(Base): def __repr__(self): return f"" + class Media(Base): __tablename__ = "media" id = Column(String, primary_key=True) # UUID @@ -226,41 +241,76 @@ class Media(Base): Index("ix_media_short_id", "short_id"), ) -################################################################################ -# Main Database (for Users) -################################################################################ - -engine = create_engine(DB_URL, echo=False) -Session = sessionmaker(bind=engine) -Base.metadata.create_all(engine) ################################################################################ # Jinja2 Environment and Custom Filters ################################################################################ + @subscriber(IJinja2Environment) def add_jinja2_filters(event): env = event.environment env.filters["filesizeformat"] = filesizeformat + +################################################################################ +# Request Methods +################################################################################ + + +def add_user_dbsession(request): + """Adds user_dbsession to request for verified users.""" + if request.user and request.user.is_verified: + user_dbsession = get_user_dbsession_by_user_id(request.user.id, request) + return user_dbsession + else: + return None # Guests do not have user_dbsession + + +def get_user_dbsession_by_user_id(user_id, request): + """Helper function to get a user_dbsession for a given user_id.""" + user_db_url = get_user_db_url(user_id) + db_file = os.path.join(APP_DIR, f"user_{user_id}.db") + if not os.path.exists(db_file): + return None # User database does not exist + user_engine = create_engine( + user_db_url, connect_args={"check_same_thread": False}, poolclass=StaticPool + ) + UserSessionFactory = sessionmaker(bind=user_engine) + user_dbsession = scoped_session(UserSessionFactory) + register(user_dbsession) # Register with zope.sqlalchemy + + # Attach cleanup callbacks + def cleanup(request): + user_dbsession.remove() + user_engine.dispose() + + request.add_finished_callback(cleanup) + return user_dbsession + + ################################################################################ # Routes ################################################################################ + @view_config(route_name="home", renderer="home.html.j2") def home_view(request): return { "request": request, } + ################################################################################ # Authentication Views ################################################################################ + @view_config(route_name="login", request_method="GET", renderer="login.html.j2") def login_get_view(request): return {"request": request} + @view_config(route_name="login", request_method="POST") def login_post_view(request): email = request.POST.get("email", "").strip().lower() @@ -284,7 +334,7 @@ def login_post_view(request): is_verified=False, ) session.add(user) - session.commit() + session.flush() # Use flush instead of commit # Generate 6-digit code code_str = f"{random.randint(0,999999):06d}" @@ -294,7 +344,7 @@ def login_post_view(request): user.code_hash = code_hash user.code_expires = datetime.datetime.now() + datetime.timedelta(minutes=15) user.is_verified = False - session.commit() + session.flush() # Send code via email email_body = f"Your verification code is: {code_str}" @@ -305,10 +355,12 @@ def login_post_view(request): return HTTPFound(location=request.route_url("verify")) + @view_config(route_name="verify", request_method="GET", renderer="verify.html.j2") def verify_get_view(request): return {"request": request} + @view_config(route_name="verify", request_method="POST") def verify_post_view(request): code_entered = request.POST.get("code", "").strip() @@ -342,23 +394,35 @@ def verify_post_view(request): user.is_verified = True user.code_hash = None user.code_expires = None - s.commit() + s.flush() # Remove the email from the session del request.session["login_email"] request.session["user_id"] = user.id + + # Create user database upon verification + user_db_url = get_user_db_url(user.id) + user_engine = create_engine( + user_db_url, connect_args={"check_same_thread": False}, poolclass=StaticPool + ) + Base.metadata.create_all(user_engine) # Create tables in user's database + user_engine.dispose() # Dispose the engine + return HTTPFound(location=request.route_url("home")) + @view_config(route_name="logout") def logout_view(request): request.session.invalidate() return HTTPFound(location=request.route_url("home")) + ################################################################################ # Profile and Media Upload ################################################################################ + @view_config(route_name="profile", request_method="GET", renderer="profile.html.j2") def profile_get_view(request): user = request.user @@ -366,18 +430,12 @@ def profile_get_view(request): return Response("You must be logged in to access your profile.", status=403) gravatar_url = get_gravatar_url(user.email) if user.enable_gravatar else "" - # Calculate stats + # Initialize stats total_uploads = 0 total_size = 0 # in bytes - user_db_url = get_user_db_url(user.id) - db_file = os.path.join(APP_DIR, f"user_{user.id}.db") - if os.path.exists(db_file): - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - Base.metadata.create_all(user_engine) - user_dbsession = UserSession() - + user_dbsession = request.user_dbsession + if user_dbsession: media_items = user_dbsession.query(Media).all() total_uploads = len(media_items) total_size = sum(media.size for media in media_items) @@ -390,6 +448,7 @@ def profile_get_view(request): "total_size": total_size, } + @view_config(route_name="profile", request_method="POST") def profile_post_view(request): if not request.user: @@ -411,9 +470,10 @@ def profile_post_view(request): return Response("Username is already in use.", status=400) request.user.username = new_username - s.commit() + s.flush() return HTTPFound(location=request.route_url("profile")) + @view_config(route_name="download_database") def download_database_view(request): user = request.user @@ -425,13 +485,17 @@ def download_database_view(request): with open(db_file, "rb") as f: data = f.read() response = Response(body=data, content_type="application/octet-stream") - response.headers["Content-Disposition"] = f'attachment; filename="user_{user.id}.db"' + response.headers["Content-Disposition"] = ( + f'attachment; filename="user_{user.id}.db"' + ) return response + ################################################################################ # User Record Export (Users) and Import (Admins) ################################################################################ + @view_config(route_name="export_user_record") def export_user_record_view(request): user = request.user @@ -453,16 +517,20 @@ def export_user_record_view(request): # Create response response = Response(body=user_json, content_type="application/json") - response.headers[ - "Content-Disposition" - ] = 'attachment; filename="user_record.json"' + response.headers["Content-Disposition"] = 'attachment; filename="user_record.json"' return response -@view_config(route_name="import_user_record", request_method="GET", renderer="import_user_record.html.j2") + +@view_config( + route_name="import_user_record", + request_method="GET", + renderer="import_user_record.html.j2", +) @admin_required def import_user_record_get_view(request): return {"request": request} + @view_config(route_name="import_user_record", request_method="POST") @admin_required def import_user_record_post_view(request): @@ -470,7 +538,10 @@ def import_user_record_post_view(request): # Get the uploaded user record file user_record_file = request.POST.get("user_record_file") - if user_record_file is None or not getattr(user_record_file, "filename", "").strip(): + if ( + user_record_file is None + or not getattr(user_record_file, "filename", "").strip() + ): return Response("No user record file uploaded.", status=400) # Read and parse the JSON data @@ -498,15 +569,19 @@ def import_user_record_post_view(request): is_verified=True, # Assume verified ) s.add(user) - s.commit() + s.flush() + + # Do not create user database upon import request.session.flash(f"User {user.username} imported successfully.") return HTTPFound(location=request.route_url("home")) + ################################################################################ # Media Upload, Listing, and Management ################################################################################ + @view_config( route_name="upload_media", request_method="GET", renderer="upload_media.html.j2" ) @@ -516,12 +591,15 @@ def upload_media_get_view(request): return Response("You must be logged in to upload media.", status=403) return {"request": request} + @view_config(route_name="upload_media", request_method="POST") def upload_media_post_view(request): user = request.user if not user or not user.is_verified: return Response("You must be logged in to upload media.", status=403) + user_dbsession = request.user_dbsession # Assume this exists for verified users + media_file = request.POST.get("media_file") if media_file is None or not getattr(media_file, "filename", "").strip(): return Response("No file uploaded.", status=400) @@ -547,13 +625,6 @@ def upload_media_post_view(request): # Encode content to base64 encoded_str = base64.b64encode(raw_bytes).decode("utf-8") - # Handle user's database - user_db_url = get_user_db_url(user.id) - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - Base.metadata.create_all(user_engine) - user_dbsession = UserSession() - # Generate UUID and short ID for the media media_uuid = uuid.uuid4() media_id = str(media_uuid) @@ -573,7 +644,8 @@ def upload_media_post_view(request): size=file_size, ) user_dbsession.add(media) - user_dbsession.commit() + user_dbsession.flush() + # No need to commit; transaction manager will handle it return HTTPFound( location=request.route_url( @@ -583,6 +655,7 @@ def upload_media_post_view(request): ) ) + @view_config(route_name="list_media", renderer="list_media.html.j2") def list_media_view(request): # Aggregate public media from all users @@ -590,14 +663,10 @@ def list_media_view(request): users = s.query(User).filter(User.is_verified == True).all() media_list = [] for user in users: - user_db_url = get_user_db_url(user.id) - db_file = os.path.join(APP_DIR, f"user_{user.id}.db") - if not os.path.exists(db_file): - continue - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - Base.metadata.create_all(user_engine) - user_dbsession = UserSession() + # Get user_dbsession + user_dbsession = get_user_dbsession_by_user_id(user.id, request) + if not user_dbsession: + continue # Skip users without database media_items = user_dbsession.query(Media).filter(Media.is_public == True).all() for media in media_items: @@ -617,6 +686,7 @@ def list_media_view(request): "media_list": media_list, } + @view_config(route_name="user_media", renderer="user_media.html.j2") def user_media_view(request): # View uploads by a particular user @@ -640,21 +710,19 @@ def user_media_view(request): viewer = request.user is_owner = viewer and viewer.id == user.id - user_db_url = get_user_db_url(user.id) - db_file = os.path.join(APP_DIR, f"user_{user.id}.db") - if not os.path.exists(db_file): + # Get user_dbsession + user_dbsession = get_user_dbsession_by_user_id(user.id, request) + if not user_dbsession: return Response("User has no uploads.", status=404) - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - Base.metadata.create_all(user_engine) - user_dbsession = UserSession() if is_owner: # Show all media (public and private) media_items = user_dbsession.query(Media).all() else: # Show only public media - media_items = user_dbsession.query(Media).filter(Media.is_public == True).all() + media_items = ( + user_dbsession.query(Media).filter(Media.is_public == True).all() + ) # Sort media by upload date (recent first) media_items.sort(key=lambda media: media.upload_date, reverse=True) @@ -670,6 +738,7 @@ def user_media_view(request): log.exception(f"Error processing user_short_id {user_short_id}") return Response("Error processing request.", status=500) + @view_config(route_name="view_media_details", renderer="view_media_details.html.j2") def view_media_details_view(request): media_short_id = request.matchdict.get("media_short_id") @@ -683,13 +752,10 @@ def view_media_details_view(request): return Response("User not found.", status=404) user_id = user.id - user_db_url = get_user_db_url(user_id) - db_file = os.path.join(APP_DIR, f"user_{user_id}.db") - if not os.path.exists(db_file): + # Get user_dbsession + user_dbsession = get_user_dbsession_by_user_id(user_id, request) + if not user_dbsession: return Response("User database not found.", status=404) - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - user_dbsession = UserSession() # Lookup media by short_id media = user_dbsession.query(Media).filter_by(short_id=media_short_id).first() @@ -710,6 +776,7 @@ def view_media_details_view(request): "is_owner": is_owner, } + @view_config(route_name="delete_media", request_method="POST") def delete_media_view(request): media_short_id = request.matchdict.get("media_short_id") @@ -722,10 +789,7 @@ def delete_media_view(request): if viewer.short_id != user_short_id: return Response("You are not authorized to delete this media.", status=403) - user_db_url = get_user_db_url(viewer.id) - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - user_dbsession = UserSession() + user_dbsession = request.user_dbsession # Assume this exists for verified users media = user_dbsession.query(Media).filter_by(short_id=media_short_id).first() if not media: @@ -733,13 +797,17 @@ def delete_media_view(request): # Delete the media user_dbsession.delete(media) - user_dbsession.commit() + user_dbsession.flush() + # No need to commit; transaction manager will handle it return HTTPFound( location=request.route_url("user_media", user_short_id=viewer.short_id) ) -@view_config(route_name="edit_media", request_method="GET", renderer="edit_media.html.j2") + +@view_config( + route_name="edit_media", request_method="GET", renderer="edit_media.html.j2" +) def edit_media_get_view(request): media_short_id = request.matchdict.get("media_short_id") user_short_id = request.matchdict.get("user_short_id") @@ -751,10 +819,7 @@ def edit_media_get_view(request): if viewer.short_id != user_short_id: return Response("You are not authorized to edit this media.", status=403) - user_db_url = get_user_db_url(viewer.id) - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - user_dbsession = UserSession() + user_dbsession = request.user_dbsession # Assume this exists for verified users media = user_dbsession.query(Media).filter_by(short_id=media_short_id).first() if not media: @@ -765,6 +830,7 @@ def edit_media_get_view(request): "media": media, } + @view_config(route_name="edit_media", request_method="POST") def edit_media_post_view(request): media_short_id = request.matchdict.get("media_short_id") @@ -777,10 +843,7 @@ def edit_media_post_view(request): if viewer.short_id != user_short_id: return Response("You are not authorized to edit this media.", status=403) - user_db_url = get_user_db_url(viewer.id) - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - user_dbsession = UserSession() + user_dbsession = request.user_dbsession # Assume this exists for verified users media = user_dbsession.query(Media).filter_by(short_id=media_short_id).first() if not media: @@ -815,7 +878,8 @@ def edit_media_post_view(request): is_public = request.POST.get("is_public") == "on" media.is_public = is_public - user_dbsession.commit() + user_dbsession.flush() + # No need to commit; transaction manager will handle it return HTTPFound( location=request.route_url( @@ -825,10 +889,12 @@ def edit_media_post_view(request): ) ) + ################################################################################ # Media Viewing and Downloading ################################################################################ + @view_config(route_name="view_media") def view_media_view(request): media_short_id = request.matchdict.get("media_short_id") @@ -842,13 +908,10 @@ def view_media_view(request): return Response("User not found.", status=404) user_id = user.id - user_db_url = get_user_db_url(user_id) - db_file = os.path.join(APP_DIR, f"user_{user_id}.db") - if not os.path.exists(db_file): + # Get user_dbsession + user_dbsession = get_user_dbsession_by_user_id(user_id, request) + if not user_dbsession: return Response("User database not found.", status=404) - user_engine = create_engine(user_db_url) - UserSession = sessionmaker(bind=user_engine) - user_dbsession = UserSession() # Lookup media by short_id media = user_dbsession.query(Media).filter_by(short_id=media_short_id).first() @@ -874,23 +937,28 @@ def view_media_view(request): # Use original filename download_filename = media.filename + # Check if the user wants to download the file + download = request.GET.get("download", "false").lower() == "true" + content_disposition = "attachment" if download else "inline" + # Serve the media content with appropriate headers response = Response(body=media_data, content_type=mime_type) response.headers.update( { "Access-Control-Allow-Origin": "*", - "Content-Disposition": f'inline; filename="{download_filename}"', + "Content-Disposition": f'{content_disposition}; filename="{download_filename}"', } ) + return response + ################################################################################ # Main ################################################################################ -def main(global_config=None, **settings): - from pyramid.decorator import reify +def main(global_config=None, **settings): # Configure logging logging.basicConfig(level=logging.DEBUG) @@ -912,6 +980,7 @@ def main(global_config=None, **settings): config = Configurator(settings=settings, session_factory=session_factory) config.include("pyramid_jinja2") + config.include("pyramid_tm") # Include pyramid_tm for transaction management # Add .html.j2 extension for Jinja2 templates # Set up Jinja2 template search path @@ -920,14 +989,29 @@ def main(global_config=None, **settings): # The Jinja2 filters are added via the event subscriber above + # Set up SQLAlchemy + engine = create_engine( + DB_URL, connect_args={"check_same_thread": False}, poolclass=StaticPool + ) + session_factory = sessionmaker(bind=engine) + Base.metadata.bind = engine + + # Use a scoped_session + DBSession = scoped_session(session_factory) + register(DBSession) # Register with zope.sqlalchemy + # Add user to all requests. config.add_request_method(callable=get_current_user, name="user", reify=True) + # Provide dbsession to requests def dbsession(request): - return Session() + return DBSession config.add_request_method(dbsession, "dbsession", reify=True) + # Add user_dbsession to requests for verified users + config.add_request_method(add_user_dbsession, "user_dbsession", reify=True) + # Routes config.add_route("home", "/") @@ -945,7 +1029,9 @@ def main(global_config=None, **settings): config.add_route("upload_media", "/media/upload") config.add_route("list_media", "/media/list") config.add_route("user_media", "/media/user/{user_short_id}") # Moved earlier - config.add_route("view_media_details", "/media/{user_short_id}/{media_short_id}/details") + config.add_route( + "view_media_details", "/media/{user_short_id}/{media_short_id}/details" + ) config.add_route("edit_media", "/media/{user_short_id}/{media_short_id}/edit") config.add_route("delete_media", "/media/{user_short_id}/{media_short_id}/delete") config.add_route("view_media", "/media/{user_short_id}/{media_short_id}") @@ -956,6 +1042,7 @@ def main(global_config=None, **settings): config.scan() return config.make_wsgi_app() + if __name__ == "__main__": app = main() log.info("Serving on http://localhost:6544") diff --git a/requirements.txt b/requirements.txt index 43e487a..bd0af6c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,6 +5,9 @@ pyramid_debugtoolbar waitress pyramid_retry +pyramid_tm +zope.sqlalchemy + bcrypt sqlalchemy werkzeug diff --git a/templates/list_media.html.j2 b/templates/list_media.html.j2 index 1c63f01..cc5d5b9 100644 --- a/templates/list_media.html.j2 +++ b/templates/list_media.html.j2 @@ -16,7 +16,7 @@ View | Download | - View Details + Details {% if request.user and request.user.id == media.user_id %} | Edit {% endif %} diff --git a/templates/user_media.html.j2 b/templates/user_media.html.j2 index 6e51146..41cd4f4 100644 --- a/templates/user_media.html.j2 +++ b/templates/user_media.html.j2 @@ -19,7 +19,7 @@ View | Download | - View Details + Details {% if request.user and request.user.id == media.user_id %} | Edit {% endif %} diff --git a/templates/view_media_details.html.j2 b/templates/view_media_details.html.j2 index 224653e..b28cffb 100644 --- a/templates/view_media_details.html.j2 +++ b/templates/view_media_details.html.j2 @@ -13,16 +13,12 @@

Size: {{ media.size | filesizeformat }}

-

- View | - Download -

+View | +Download {% if is_owner %} - -

- Edit -

+ +| Edit {% endif %} {% endblock %}