Spaces:
Running
Running
| """User module""" | |
| import hashlib | |
| import secrets | |
| import time | |
| from datetime import datetime, timedelta, timezone | |
| from db import get_service_client | |
| from engineHelper import upload_media_to_cloudinary | |
| _dev_cache: set | None = None | |
| class PasswordResetUnavailableError(RuntimeError): | |
| """Supabase is missing the password_reset_tokens table (migration not applied).""" | |
| RESET_TABLE_INSTRUCTIONS = ( | |
| "Password reset is not set up: the database table is missing. " | |
| "In Supabase → SQL → run the script in docs/Design/password_reset_tokens.sql, then try again." | |
| ) | |
| def _missing_password_reset_table(exc: BaseException) -> bool: | |
| s = str(exc) | |
| if "PGRST205" in s: | |
| return True | |
| return "password_reset_tokens" in s and "Could not find" in s | |
| def is_developer(username: str) -> bool: | |
| global _dev_cache | |
| if _dev_cache is None: | |
| sb = get_service_client() | |
| rows = sb.table("developers").select("username").execute() | |
| _dev_cache = {r["username"] for r in rows.data} | |
| return username in _dev_cache | |
| def _generate_salt(length=32): | |
| charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789" | |
| return ''.join(secrets.choice(charset) for _ in range(length)) | |
| def _hash_password(password, salt): | |
| return hashlib.sha256((salt + password).encode()).hexdigest() | |
| def _clamp_theme(theme_val): | |
| try: | |
| t = int(theme_val) | |
| except (TypeError, ValueError): | |
| return 5 | |
| return max(0, min(6, t)) | |
| def register_user(fn, ln, gender, uname, email, pw, theme=5): | |
| sb = get_service_client() | |
| existing = sb.table("users").select("username").eq("username", uname).execute() | |
| if existing.data: | |
| raise RuntimeError("Username already exists") | |
| avatar = "" | |
| if gender == "male": avatar = "https://urgedwtmlfsmpkomptmb.supabase.co/storage/v1/object/public/media/42ffc00b3c154a9c8a2001f54c174d77.png" | |
| elif gender == "female": avatar = "https://urgedwtmlfsmpkomptmb.supabase.co/storage/v1/object/public/media/ee2f12f4a6ce48ada053e499caea6019.png" | |
| else: avatar = "https://urgedwtmlfsmpkomptmb.supabase.co/storage/v1/object/public/media/950a997d996c4ce2a76f3d6d29597671.png" | |
| salt = _generate_salt() | |
| h = _hash_password(pw, salt) | |
| theme_i = _clamp_theme(theme) | |
| sb.table("users").insert({ | |
| "username": uname, "firstName": fn, "lastName": ln, "gender": gender, | |
| "email": email, "password_hash": h, "password_salt": salt, | |
| "bio": "", "followers": 0, "follows": 0, "avatar": avatar, "heroBanner": "", "theme": theme_i | |
| }).execute() | |
| return True | |
| def login(uname, pw): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("password_hash, password_salt").eq("username", uname).execute() | |
| if not rows.data: return False | |
| r = rows.data[0] | |
| return r["password_hash"] == _hash_password(pw, r["password_salt"]) | |
| def get_username(user_input): | |
| sb = get_service_client() | |
| user_input = user_input.strip() | |
| result = {} | |
| if "@" in user_input: | |
| rows = sb.table("users").select("username").eq("email", user_input).execute() | |
| usernames = ",".join(r["username"] for r in rows.data) | |
| result = {"email": user_input, "usernames": usernames, "count": str(len(rows.data))} | |
| else: | |
| result = {"email": "", "usernames": user_input, "count": "1"} | |
| return result | |
| def get_user_by_username(username): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("firstName, lastName, gender, username, email, bio, followers, follows, visibility, avatar, \"heroBanner\"").eq("username", username).execute() | |
| if not rows.data: return {} | |
| r = rows.data[0] | |
| return {k: str(v) for k, v in r.items()} | |
| def get_all_users(): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("username, visibility").neq("username", "BitAI").execute() | |
| return rows.data | |
| def user_exists(username): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("username").eq("username", username).execute() | |
| return len(rows.data) > 0 | |
| def get_all_users_detailed(): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("*").neq("username", "BitAI").order("username").execute() | |
| return rows.data | |
| def get_theme(username): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("theme").eq("username", username).execute() | |
| if not rows.data: return "5" | |
| return str(rows.data[0]["theme"]) | |
| def update_theme(username, theme): | |
| sb = get_service_client() | |
| sb.table("users").update({"theme": theme}).eq("username", username).execute() | |
| def get_avatar(username): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("avatar").eq("username", username).execute() | |
| if not rows.data: return "" | |
| return rows.data[0]["avatar"] or "" | |
| def update_avatar(username, avatar_base64): | |
| sb = get_service_client() | |
| if avatar_base64 and avatar_base64.startswith("data:"): | |
| avatar_base64 = upload_media_to_cloudinary(avatar_base64) | |
| sb.table("users").update({"avatar": avatar_base64}).eq("username", username).execute() | |
| def get_hero_banner(username): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("heroBanner").eq("username", username).execute() | |
| if not rows.data: return "" | |
| return rows.data[0]["heroBanner"] or "" | |
| def update_hero_banner(username, hero_banner_base64): | |
| sb = get_service_client() | |
| if hero_banner_base64 and hero_banner_base64.startswith("data:"): | |
| hero_banner_base64 = upload_media_to_cloudinary(hero_banner_base64) | |
| sb.table("users").update({"heroBanner": hero_banner_base64}).eq("username", username).execute() | |
| def update_user_basic(current_username, new_username, email, bio): | |
| sb = get_service_client() | |
| if new_username and new_username != current_username: | |
| existing = sb.table("users").select("username").eq("username", new_username).execute() | |
| if existing.data: return False | |
| updates = {} | |
| if new_username: updates["username"] = new_username | |
| if email: updates["email"] = email | |
| if bio: updates["bio"] = bio | |
| if updates: | |
| # ON UPDATE CASCADE handles updating related tables! | |
| sb.table("users").update(updates).eq("username", current_username).execute() | |
| return True | |
| def set_visibility(username, visibility): | |
| """0: Private, 1: Public""" | |
| sb = get_service_client() | |
| sb.table("users").update({"visibility": int(visibility)}).eq("username", username).execute() | |
| return True | |
| def update_password(username, old_pass, new_pass): | |
| sb = get_service_client() | |
| rows = sb.table("users").select("password_hash, password_salt").eq("username", username).execute() | |
| if not rows.data: return False | |
| r = rows.data[0] | |
| if r["password_hash"] != _hash_password(old_pass, r["password_salt"]): return False | |
| new_salt = _generate_salt() | |
| new_hash = _hash_password(new_pass, new_salt) | |
| sb.table("users").update({"password_hash": new_hash, "password_salt": new_salt}).eq("username", username).execute() | |
| return True | |
| def _parse_expires_at(val): | |
| if not val: | |
| return None | |
| s = str(val).replace("Z", "+00:00") | |
| try: | |
| return datetime.fromisoformat(s) | |
| except Exception: | |
| return None | |
| def create_password_reset_token(username: str): | |
| """Create a one-time reset token; returns plaintext token for email, or None if user missing.""" | |
| if not user_exists(username): | |
| return None | |
| sb = get_service_client() | |
| token_plain = secrets.token_urlsafe(32) | |
| token_hash = hashlib.sha256(token_plain.encode()).hexdigest() | |
| expires_at = datetime.now(timezone.utc) + timedelta(hours=1) | |
| try: | |
| sb.table("password_reset_tokens").insert( | |
| { | |
| "username": username, | |
| "token_hash": token_hash, | |
| "expires_at": expires_at.isoformat(), | |
| "used_at": None, | |
| } | |
| ).execute() | |
| except Exception as e: | |
| if _missing_password_reset_table(e): | |
| raise PasswordResetUnavailableError(RESET_TABLE_INSTRUCTIONS) from e | |
| # Check for RLS error 42501 | |
| if hasattr(e, 'code') and e.code == '42501': | |
| print("ERROR: Row-Level Security violation. Ensure your SUPABASE_SERVICE_KEY is a real 'service_role' key, not an 'anon' key.") | |
| raise | |
| return token_plain | |
| def validate_password_reset_token(token_plain: str): | |
| """ | |
| Returns dict with username if token is valid and unused, else None. | |
| """ | |
| if not token_plain or not isinstance(token_plain, str): | |
| return None | |
| sb = get_service_client() | |
| token_hash = hashlib.sha256(token_plain.encode()).hexdigest() | |
| try: | |
| rows = ( | |
| sb.table("password_reset_tokens") | |
| .select("username, expires_at, used_at") | |
| .eq("token_hash", token_hash) | |
| .execute() | |
| ) | |
| except Exception as e: | |
| if _missing_password_reset_table(e): | |
| raise PasswordResetUnavailableError(RESET_TABLE_INSTRUCTIONS) from e | |
| raise | |
| if not rows.data: | |
| return None | |
| r = rows.data[0] | |
| if r.get("used_at"): | |
| return None | |
| exp = _parse_expires_at(r.get("expires_at")) | |
| if exp is None: | |
| return None | |
| if exp.tzinfo is None: | |
| exp = exp.replace(tzinfo=timezone.utc) | |
| if datetime.now(timezone.utc) > exp: | |
| return None | |
| return {"username": r["username"]} | |
| def get_user_public_for_reset(username: str): | |
| """Safe fields for reset-password page.""" | |
| u = get_user_by_username(username) | |
| if not u: | |
| return None | |
| return { | |
| "username": u.get("username", ""), | |
| "firstName": u.get("firstName", ""), | |
| "lastName": u.get("lastName", ""), | |
| "email": u.get("email", ""), | |
| "avatar": get_avatar(username), | |
| } | |
| def confirm_password_reset(token_plain: str, new_password: str): | |
| if not new_password or len(new_password) < 1: | |
| return False | |
| info = validate_password_reset_token(token_plain) | |
| if not info: | |
| return False | |
| username = info["username"] | |
| sb = get_service_client() | |
| new_salt = _generate_salt() | |
| new_hash = _hash_password(new_password, new_salt) | |
| sb.table("users").update({"password_hash": new_hash, "password_salt": new_salt}).eq( | |
| "username", username | |
| ).execute() | |
| token_hash = hashlib.sha256(token_plain.encode()).hexdigest() | |
| now = datetime.now(timezone.utc).isoformat() | |
| try: | |
| sb.table("password_reset_tokens").update({"used_at": now}).eq("token_hash", token_hash).execute() | |
| except Exception as e: | |
| if _missing_password_reset_table(e): | |
| raise PasswordResetUnavailableError(RESET_TABLE_INSTRUCTIONS) from e | |
| raise | |
| return username | |
| def delete_user(username): | |
| sb = get_service_client() | |
| # Cascading deletes handle everything else! | |
| sb.table("users").delete().eq("username", username).execute() | |
| return True | |
| def follow_user(follower, following): | |
| if follower == following: return False | |
| sb = get_service_client() | |
| # Check target visibility | |
| target = sb.table("users").select("visibility").eq("username", following).execute() | |
| if not target.data: return False | |
| visibility = target.data[0]["visibility"] | |
| status = 1 if int(visibility) == 1 else 0 # 1 if public, 0 if private | |
| existing = sb.table("follows").select("id, status").eq("follower", follower).eq("following", following).execute() | |
| if existing.data: return False | |
| try: | |
| sb.table("follows").insert({"follower": follower, "following": following, "status": status}).execute() | |
| except Exception as e: | |
| return False | |
| return {"status": "success", "follow_status": status} | |
| def unfollow_user(follower, following): | |
| if follower == following: return False | |
| sb = get_service_client() | |
| sb.table("follows").delete().eq("follower", follower).eq("following", following).execute() | |
| return True | |
| def approve_follow(follower, following): | |
| sb = get_service_client() | |
| sb.table("follows").update({"status": 1}).eq("follower", follower).eq("following", following).execute() | |
| return True | |
| def reject_follow(follower, following): | |
| sb = get_service_client() | |
| sb.table("follows").delete().eq("follower", follower).eq("following", following).execute() | |
| return True | |
| def get_pending_follows(username): | |
| sb = get_service_client() | |
| rows = sb.table("follows").select("follower").eq("following", username).eq("status", 0).execute() | |
| followers = [r["follower"] for r in rows.data] | |
| if not followers: return [] | |
| profiles = sb.table("users").select("username, avatar").in_("username", followers).execute() | |
| return profiles.data | |
| def get_follow_status(follower, following): | |
| """Returns 1 (Following), 0 (Pending), -1 (Not Following)""" | |
| if not follower or not following: return -1 | |
| sb = get_service_client() | |
| rows = sb.table("follows").select("status").eq("follower", follower).eq("following", following).execute() | |
| if not rows.data: return -1 | |
| return int(rows.data[0]["status"]) | |
| def get_followers(username): | |
| sb = get_service_client() | |
| rows = sb.table("follows").select("follower").eq("following", username).eq("status", 1).execute() | |
| return [r["follower"] for r in rows.data] | |
| def get_following(username): | |
| sb = get_service_client() | |
| rows = sb.table("follows").select("following").eq("follower", username).eq("status", 1).execute() | |
| return [r["following"] for r in rows.data] | |
| # ======================== | |
| # BLOCKING | |
| # ======================== | |
| def block_user(blocker, blocked): | |
| if not blocker or not blocked or blocker == blocked: | |
| return False | |
| sb = get_service_client() | |
| # Check if already blocked | |
| existing = sb.table("blocks").select("id").eq("blocker", blocker).eq("blocked", blocked).execute() | |
| if existing.data: | |
| return True # Treat as success if already blocked | |
| try: | |
| # Let DB handle the timestamp via DEFAULT | |
| sb.table("blocks").insert({ | |
| "blocker": blocker, | |
| "blocked": blocked | |
| }).execute() | |
| # Remove follow relationships | |
| unfollow_user(blocker, blocked) | |
| unfollow_user(blocked, blocker) | |
| return True | |
| except Exception as e: | |
| return False | |
| def unblock_user(blocker, blocked): | |
| if not blocker or not blocked: return False | |
| sb = get_service_client() | |
| sb.table("blocks").delete().eq("blocker", blocker).eq("blocked", blocked).execute() | |
| return True | |
| def is_blocked(blocker, blocked): | |
| if not blocker or not blocked: return False | |
| sb = get_service_client() | |
| rows = sb.table("blocks").select("id").eq("blocker", blocker).eq("blocked", blocked).execute() | |
| return len(rows.data) > 0 | |
| def is_blocked_either_way(user1, user2): | |
| if not user1 or not user2: return False | |
| sb = get_service_client() | |
| res1 = sb.table("blocks").select("id").eq("blocker", user1).eq("blocked", user2).limit(1).execute() | |
| if res1.data: return True | |
| res2 = sb.table("blocks").select("id").eq("blocker", user2).eq("blocked", user1).limit(1).execute() | |
| return len(res2.data) > 0 | |
| def get_blocked_users(blocker): | |
| if not blocker: return [] | |
| sb = get_service_client() | |
| rows = sb.table("blocks").select("blocked").eq("blocker", blocker).execute() | |
| return [r["blocked"] for r in rows.data] | |
| def get_blockers(blocked): | |
| if not blocked: return [] | |
| sb = get_service_client() | |
| rows = sb.table("blocks").select("blocker").eq("blocked", blocked).execute() | |
| return [r["blocker"] for r in rows.data] | |
| def get_blocking_relationship_usernames(username): | |
| if not username: return set() | |
| blocked = get_blocked_users(username) | |
| blockers = get_blockers(username) | |
| return set(blocked) | set(blockers) | |