147 lines
4.4 KiB
Python
147 lines
4.4 KiB
Python
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()
|