import asyncio import datetime import os import shutil from typing import Annotated, Optional, List, Tuple, Dict import modal from loguru import logger from modal import current_function_call_id import sentry_sdk from fastapi import APIRouter, Depends, UploadFile, HTTPException, File, Form from fastapi.responses import JSONResponse, RedirectResponse from starlette import status import boto3 from botocore.config import Config from typing_inspection.typing_objects import target from ..config import WorkerConfig from ..middleware.authorization import verify_token from ..models.media_model import (MediaSources, CacheResult, MediaSource, MediaCacheStatus, DownloadResult, UploadResultResponse, UploadBase64Request, UploadPresignRequest, UploadPresignResponse, UploadMultipartPresignRequest, UploadMultipartPresignResponse, MediaProtocol ) from ..models.web_model import SentryTransactionInfo, MonitorLiveRoomProductRequest, ModalTaskResponse, \ LiveRoomProductCachesResponse, CacheDeleteTaskResponse, ClusterCacheBatchRequest, CacheOperationType, \ ClusterCacheBatchResponse, CacheTaskResult from ..utils.KVCache import MediaSourceKVCache, LiveProductKVCache from ..utils.SentryUtils import SentryUtils config = WorkerConfig() client = boto3.client("s3", aws_access_key_id=os.environ.get("AWS_ACCESS_KEY_ID"), aws_secret_access_key=os.environ.get("AWS_SECRET_ACCESS_KEY"), region_name=config.S3_region, endpoint_url="https://s3-accelerate.amazonaws.com", config=Config( s3={'addressing_style': 'virtual'}, signature_version='s3v4', ) ) router = APIRouter(prefix="/cache", tags=['缓存'], ) modal_kv_cache = MediaSourceKVCache(kv_name=config.modal_kv_name, environment=config.modal_environment) modal_kv_product_cache = LiveProductKVCache(kv_name=config.modal_product_kv_name, environment=config.modal_environment) @router.post("/", summary="缓存视频文件", description="异步缓存视频文件到S3存储桶和Modal Dict(KV)", dependencies=[Depends(verify_token)]) async def cache(medias: MediaSources) -> CacheResult: fn_id = current_function_call_id() caches: MediaSources sentry_trace = SentryTransactionInfo(x_trace_id=sentry_sdk.get_traceparent(), x_baggage=sentry_sdk.get_baggage()) @SentryUtils.sentry_tracker(name="同步视频缓存", op="cache.get", fn_id=fn_id, sentry_trace_id=None, sentry_baggage=None) async def cache_handler(media: MediaSource): cache_span = sentry_sdk.get_current_span() cache_span.set_data("runner_id", fn_id) cache_span.set_data("cache.key", [media.urn]) cached_media = modal_kv_cache.get_cache(media.urn) cache_hit: bool = False if not cached_media: # start new download task with cache_span.start_child(name="视频缓存任务入队", op="queue.publish") as queue_publish_span: fn = modal.Function.from_name(config.modal_app_name, 'cache_submit') fn_task = fn.spawn(media, sentry_trace) queue_publish_span.set_data("cache.key", media.urn) queue_publish_span.set_data("messaging.message.id", fn_task.object_id) queue_publish_span.set_data("messaging.destination.name", "video-downloader.cache_submit") queue_publish_span.set_data("messaging.message.body.size", 0) media.status = MediaCacheStatus.downloading media.downloader_id = fn_task.object_id modal_kv_cache.set_cache(media) else: media = cached_media match media.status: case MediaCacheStatus.ready: cache_hit = True case MediaCacheStatus.downloading: # 下载任务已经在进行 cache_hit = True case _: # start new download task with cache_span.start_child(name="视频缓存任务入队", op="queue.publish") as queue_publish_span: fn = modal.Function.from_name(config.modal_app_name, 'cache_submit') fn_task = fn.spawn(media, sentry_trace) queue_publish_span.set_data("cache.key", media.urn) queue_publish_span.set_data("messaging.message.id", fn_task.object_id) queue_publish_span.set_data("messaging.destination.name", "video-downloader.cache_submit") queue_publish_span.set_data("messaging.message.body.size", 0) media.status = MediaCacheStatus.downloading media.downloader_id = fn_task.object_id modal_kv_cache.set_cache(media) cache_hit = False cache_span.set_data("cache.hit", cache_hit) return media async with asyncio.TaskGroup() as group: tasks = [group.create_task(cache_handler(media)) for media in medias.inputs] cache_task_result_dict = {} cache_task_result_list = [] for task in tasks: result = task.result() cache_task_result_dict[result.urn] = result.model_dump_json() cache_task_result_list.append(result) modal_kv_cache.batch_update_cloudflare_kv(cache_task_result_dict) return CacheResult(caches={media.urn: media for media in cache_task_result_list}) @router.delete("/", summary="清除指定的所有缓存", description="清除指定的所有缓存(包括KV记录和S3存储文件)", dependencies=[Depends(verify_token)]) async def purge_media_kv_file(medias: MediaSources) -> CacheDeleteTaskResponse: fn_id = current_function_call_id() fn = modal.Function.from_name(config.modal_app_name, "cache_delete", environment_name=config.modal_environment) @SentryUtils.sentry_tracker(name="清除媒体源缓存", op="cache.purge", fn_id=fn_id, sentry_trace_id=None, sentry_baggage=None) async def purge_handle(media: MediaSource) -> Tuple[Optional[str], int]: try: cache_media = modal_kv_cache.pop(media.urn) if cache_media: deleted_cache: MediaSource = await fn.remote.aio(cache_media) return deleted_cache.urn, 1 except KeyError as e: logger.exception(e) if media.local_exists: deleted_cache: MediaSource = await fn.remote.aio(media) if deleted_cache.status == MediaCacheStatus.missing: logger.warning(f"不存在s3挂载文件 {deleted_cache.urn}") return None, 0 else: logger.warning(f"s3挂载文件 {deleted_cache.urn} 已删除") return deleted_cache.urn, 0 return media.urn, -1 async with asyncio.TaskGroup() as group: tasks = [group.create_task(purge_handle(media)) for media in medias.inputs] keys: List[str] = [] non_kv_keys: List[str] = [] error_keys: List[str] = [] for task in tasks: urn, task_status = task.result() if urn: if task_status == 1: # 成功从kv和s3删除 keys.append(urn) elif task_status == 2: # 只从s3删除 non_kv_keys.append(urn) else: error_keys.append(urn) # keys = [task.result() for task in tasks] modal_kv_cache.batch_remove_cloudflare_kv(keys) return CacheDeleteTaskResponse(success=True, keys=keys, nonKVKeys=non_kv_keys, notFoundKeys=error_keys) # return JSONResponse(content={"success": True, "keys": keys, "nonKVKeys": non_kv_keys, "errorKeys": error_keys}) @router.post("/download", summary="批量获取下载地址, 返回的MediaSource类自带CDN访问URL, 不需要另外请求获取", deprecated=True, description="获取已缓存的视频下载地址", dependencies=[Depends(verify_token)]) @sentry_sdk.trace async def download_caches(medias: MediaSources) -> DownloadResult: cdn_endpoint = config.S3_cdn_endpoint urls = [] for media in medias.inputs: urls.append(f"{cdn_endpoint}/{media.get_cdn_url()}") return DownloadResult(urls=urls) @router.get("/download", summary="下载已缓存的视频", deprecated=True, description="通过CDN下载已缓存的视频文件, 不在提供通过此接口下载,请使用CDN URL直接下载文件") @sentry_sdk.trace async def download_cache(media: str) -> RedirectResponse: cdn_endpoint = config.S3_cdn_endpoint media = MediaSource.from_str(media) return RedirectResponse(url=f"{cdn_endpoint}/{media.get_cdn_url()}", status_code=status.HTTP_302_FOUND) @router.delete("/kv", summary="清除KV记录", description="清除当前环境下KV缓存过的所有数据(S3存储桶内的文件会保留)", dependencies=[Depends(verify_token)]) async def purge_kv_all(): parent = sentry_sdk.get_current_span() span = parent.start_child(name="清除缓存KV", op="purge.flush") modal_kv_cache.clear() span.set_data("cache.success", True) span.finish() return JSONResponse(content={"success": True}) @router.post("/kv", summary="删除对应的KV记录", description="删除请求中对应的视频缓存记录", dependencies=[Depends(verify_token)]) async def purge_kv(medias: MediaSources): try: for media in medias.inputs: modal_kv_cache.pop(media.urn) keys = [media.urn for media in medias.inputs] modal_kv_cache.batch_remove_cloudflare_kv(keys) return JSONResponse(content={"success": True, "keys": keys}) except Exception as e: return JSONResponse(content={"success": False, "error": str(e)}) @router.post("/batch", summary="批量操作集群S3缓存", description="批量操作集群S3缓存", dependencies=[Depends(verify_token)]) async def s3_copy(body: ClusterCacheBatchRequest) -> ClusterCacheBatchResponse: results: List[CacheTaskResult] = [] for task in body.tasks: try: if task.target: os.makedirs(os.path.dirname(task.target.local_mount_path), exist_ok=True) match task.type: case CacheOperationType.copy: shutil.copy(task.source.local_mount_path, task.target.local_mount_path) case CacheOperationType.delete: os.remove(task.source.local_mount_path) case CacheOperationType.move: shutil.copy(task.source.local_mount_path, task.target.local_mount_path) os.remove(task.source.local_mount_path) result = CacheTaskResult(**task.model_dump(), success=True) results.append(result) except Exception as e: logger.exception(e) result = CacheTaskResult(**task.model_dump(), success=False) results.append(result) return ClusterCacheBatchResponse(results=results) @router.post("/upload-s3", summary="上传文件到S3", description="上传文件到S3的文件必须小于200M", dependencies=[Depends(verify_token)]) async def s3_upload(file: Annotated[UploadFile, File(description="上传的文件")], prefix: Annotated[Optional[str], Form()] = None) -> UploadResultResponse: fn_id = current_function_call_id() if file.size > 200 * 1024 * 1024: # 上传文件不大于200M raise HTTPException(status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, detail="上传文件不可超过200MB") key = f"upload/{prefix}/{file.filename}" if prefix else f"upload/{file.filename}" local_path = f"{config.S3_mount_dir}/{key}" logger.info(f"s3上传到{key}, size={file.size}") os.makedirs(os.path.dirname(local_path), exist_ok=True) with open(local_path, 'wb') as f: f.write(file.file.read()) logger.info(f"{local_path} 保存成功") media_source = MediaSource.from_str(f"s3://{config.S3_region}/{config.S3_bucket_name}/{key}") media_source.status = MediaCacheStatus.ready media_source.downloader_id = fn_id return UploadResultResponse(media=media_source) @router.post('/upload-s3-b64', summary="基于Base64格式上传文件到S3", description="上传文件到S3当文件必须小于200M", dependencies=[Depends(verify_token)]) async def s3_upload_base64(body: UploadBase64Request) -> UploadResultResponse: fn_id = current_function_call_id() prefix = body.prefix file = body.file if file.size > 200 * 1024 * 1024: # 上传文件不大于200M raise HTTPException(status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, detail="上传文件不可超过200MB") key = f"upload/{prefix}/{file.filename}" if prefix else f"upload/{file.filename}" local_path = f"{config.S3_mount_dir}/{key}" logger.info(f"s3上传到{key}, size={file.size}") os.makedirs(os.path.dirname(local_path), exist_ok=True) with open(local_path, 'wb') as f: f.write(file.raw_content) logger.info(f"{local_path} 保存成功") media_source = MediaSource.from_str(f"s3://{config.S3_region}/{config.S3_bucket_name}/{key}") media_source.status = MediaCacheStatus.ready media_source.downloader_id = fn_id return UploadResultResponse(media=media_source) @router.post("/monitor_live_room_product_trigger", summary="触发监控直播间商品信息并缓存", description="触发监控直播间商品信息并缓存, 如果直播结束清除缓存, 触发间隔请控制在60s以上", dependencies=[Depends(verify_token)], deprecated=True) async def monitor_live_room_product(body: MonitorLiveRoomProductRequest) -> LiveRoomProductCachesResponse: fn = modal.Function.from_name(config.modal_app_name, "monitor_live_room_product_trigger", environment_name=config.modal_environment) status = await fn.remote.aio(body.cookie, body.room_id, body.author_id) if status == 0: product_list = modal_kv_product_cache.get_cache(body.room_id) return LiveRoomProductCachesResponse(status=status, cache_json=product_list.model_dump_json()) elif status == 1: return LiveRoomProductCachesResponse(status=status, message="直播已结束") elif status == 2: return LiveRoomProductCachesResponse(status=status, message="分配到风控IP, 请稍后重试") elif status == 3: return LiveRoomProductCachesResponse(status=status, message="请求Tikhub API出现错误") else: return LiveRoomProductCachesResponse(status=4, message="内部错误") @router.post('/upload-s3/simple/presign', summary="S3简单上传预签名", description="利用S3就近接入点上传", dependencies=[Depends(verify_token)]) async def s3_presign_upload(body: UploadPresignRequest) -> UploadPresignResponse: expires_in = 3600 expired_at = datetime.datetime.now() + datetime.timedelta(seconds=expires_in) signed_url = client.generate_presigned_url("put_object", Params={ 'Bucket': config.S3_bucket_name, 'Key': f"upload/{body.key}", "ContentType": body.content_type, }, ExpiresIn=expires_in, ) return UploadPresignResponse(url=signed_url, urn=f"s3://{config.S3_region}/{config.S3_bucket_name}/upload/{body.key}", expired_at=expired_at) @router.post("/upload-s3/multipart/presign", summary="S3分片上传预签名", description=""" 1. 本地按文件总大小分Chunk大小按urls链接内的顺序通过HTTP PUT请求上传文件分片; 并将上传完成后获得的返回头ETag值记录与PartNumber对应, PartNumber对应使用url在urls内的顺位, 以1开始 \n2. 所有分片上传完成后使用XML格式拼装出用于确认上传的body; 并通过HTTP POST complete_url确认上传, ContentType 需要确保为application/xml\n\n 1 "60575364b098a1a48765a28c3a48e0ef" 2 "a38691c31fd242faee5533c65b4501d7" 3 "7a7510cc83f98feea28f319567e4cf66" \n3. 如无法确认上传,可使用HTTP GET list_url debug当前分片上传状态,确认成功后list_url无法返回有效数据 """, dependencies=[Depends(verify_token)]) async def s3_presign_upload_multipart(body: UploadMultipartPresignRequest) -> UploadMultipartPresignResponse: chunk_count = body.parts_count multipart_upload_response = client.create_multipart_upload(Bucket=config.S3_bucket_name, Key=body.key, ContentType=body.content_type, ) upload_id = multipart_upload_response.get("UploadId") signed_urls = [] expires_in = 3600 expired_at = datetime.datetime.now() + datetime.timedelta(seconds=expires_in) for i in range(chunk_count): signed_url = client.generate_presigned_url("upload_part", Params={ 'Bucket': config.S3_bucket_name, 'Key': body.key, 'PartNumber': i + 1, 'UploadId': upload_id, }, ExpiresIn=expires_in) signed_urls.append(signed_url) signed_completed_url = client.generate_presigned_url("complete_multipart_upload", Params={ 'Bucket': config.S3_bucket_name, 'Key': body.key, 'UploadId': upload_id, }, ExpiresIn=expires_in) signed_list_url = client.generate_presigned_url("list_parts", Params={ 'Bucket': config.S3_bucket_name, 'Key': body.key, 'UploadId': upload_id, }, ExpiresIn=expires_in) return UploadMultipartPresignResponse(urls=signed_urls, urn=f"s3://{config.S3_region}/{config.S3_bucket_name}/upload/{body.key}", expired_at=expired_at, complete_url=signed_completed_url, list_url=signed_list_url)