from __future__ import annotations import argparse import asyncio import os import sys from pathlib import Path from urllib.parse import urlsplit from sqlalchemy import select BACKEND_DIR = Path(__file__).resolve().parents[2] if str(BACKEND_DIR) not in sys.path: sys.path.insert(0, str(BACKEND_DIR)) from app.database import AsyncSessionLocal from app.models.activity import Activity from app.models.user import User CANONICAL_PREFIX = '/api/v1/static/' def _canonicalize_media_url(value: str | None) -> str | None: if not value: return None text = str(value).strip() if not text: return None if text.startswith(('data:', 'blob:', 'wxfile:', 'file:')): return text if text.startswith('//'): text = f'https:{text}' if text.startswith(('http:/', 'https:/')) and not text.startswith(('http://', 'https://')): text = text.replace(':/', '://', 1) if text.startswith('http://') or text.startswith('https://'): parsed = urlsplit(text) path = parsed.path or '/' if path.startswith('/api/v1/static/'): return path if path.startswith('/static/'): return f'/api/v1{path}' if path.startswith('/v1/static/'): return path.replace('/v1/', '/api/v1/', 1) return text if text.startswith('/api/v1/static/'): return text if text.startswith('/static/'): return f'/api/v1{text}' if text.startswith('/v1/static/'): return text.replace('/v1/', '/api/v1/', 1) if text.startswith('/'): return f'/api/v1/static/{text.lstrip('/')}' return f'/api/v1/static/{text.lstrip('/')}' async def normalize_activities(session) -> int: result = await session.execute(select(Activity)) items = result.scalars().all() changed = 0 for item in items: new_cover = _canonicalize_media_url(item.cover_image) new_share_qr = _canonicalize_media_url(item.share_qr_url) if new_cover != item.cover_image: item.cover_image = new_cover changed += 1 if new_share_qr != item.share_qr_url: item.share_qr_url = new_share_qr changed += 1 session.add(item) await session.commit() return changed async def normalize_users(session) -> int: result = await session.execute(select(User)) items = result.scalars().all() changed = 0 for item in items: new_avatar = _canonicalize_media_url(item.avatar_url) new_blur = _canonicalize_media_url(item.avatar_blur_url) if new_avatar != item.avatar_url: item.avatar_url = new_avatar changed += 1 if new_blur != item.avatar_blur_url: item.avatar_blur_url = new_blur changed += 1 if isinstance(item.profile_images, list): new_list = [_canonicalize_media_url(url) for url in item.profile_images] new_list = [url for url in new_list if url] if new_list != item.profile_images: item.profile_images = new_list changed += 1 session.add(item) await session.commit() return changed async def run() -> None: async with AsyncSessionLocal() as session: activity_changes = await normalize_activities(session) user_changes = await normalize_users(session) print(f'Normalized activities fields: {activity_changes}') print(f'Normalized user fields: {user_changes}') def _env_candidates() -> list[Path]: script_dir = Path(__file__).resolve().parent backend_dir = script_dir.parents[2] repo_root = backend_dir.parent return [backend_dir / '.env', repo_root / '.env', script_dir / '.env'] def _load_env_file_if_needed() -> None: if os.getenv('MYSQL_HOST') and os.getenv('MYSQL_DATABASE'): return try: from dotenv import load_dotenv except ImportError: return for path in _env_candidates(): if path.is_file(): load_dotenv(path) break def main() -> None: parser = argparse.ArgumentParser(description='Normalize historical media URLs in the database.') parser.add_argument('--confirm', action='store_true', help='确认执行批量修复') args = parser.parse_args() if not args.confirm: raise SystemExit('请添加 --confirm 参数后再执行,以避免误操作。') _load_env_file_if_needed() asyncio.run(run()) if __name__ == '__main__': main()