92 lines
2.7 KiB
Python
92 lines
2.7 KiB
Python
# -*- coding:utf-8 -*-
|
|
import asyncio
|
|
import json
|
|
from ahserver import filedownload
|
|
from ahserver.webapp import webapp
|
|
from ahserver.serverenv import ServerEnv
|
|
from ahserver.configuredServer import add_startup
|
|
from longtasks.longtasks import LongTasks, schedule_once
|
|
from appPublic.log import debug, exception
|
|
|
|
|
|
class MediaTasks(LongTasks):
|
|
def __init__(self, *args, **kwargs):
|
|
super().__init__(*args, **kwargs)
|
|
self._handlers = {}
|
|
|
|
def register(self, task_type: str, handler):
|
|
self._handlers[task_type] = handler
|
|
debug(f'MediaTasks: registered handler for {task_type}')
|
|
|
|
async def process_task(self, payload: dict, workid: int = None):
|
|
if isinstance(payload, str):
|
|
payload = json.loads(payload)
|
|
task_type = payload.get('task_type', '')
|
|
debug(f'MediaTasks processing: type={task_type}')
|
|
handler = self._handlers.get(task_type)
|
|
if handler is None:
|
|
raise ValueError(f'Unknown task_type: {task_type}')
|
|
return await handler(payload, workid)
|
|
|
|
|
|
class GPULock:
|
|
def __init__(self, redis_client, total_gpus=8):
|
|
self.redis = redis_client
|
|
self.total_gpus = total_gpus
|
|
self.lock_prefix = 'gpu:lock:'
|
|
|
|
async def acquire(self, task_id: str, timeout=600):
|
|
for gpu_id in range(self.total_gpus):
|
|
key = f'{self.lock_prefix}{gpu_id}'
|
|
acquired = await self.redis.set(key, task_id, nx=True, ex=timeout)
|
|
if acquired:
|
|
return gpu_id
|
|
return None
|
|
|
|
async def release(self, gpu_id: int):
|
|
key = f'{self.lock_prefix}{gpu_id}'
|
|
await self.redis.delete(key)
|
|
|
|
async def status(self):
|
|
result = {}
|
|
for gpu_id in range(self.total_gpus):
|
|
key = f'{self.lock_prefix}{gpu_id}'
|
|
owner = await self.redis.get(key)
|
|
result[gpu_id] = {'busy': owner is not None, 'owner': owner}
|
|
return result
|
|
|
|
|
|
async def handle_ktv_pipeline(payload, workid=None):
|
|
from workers.ktv_pipeline import run_pipeline
|
|
pipeline_id = payload.get('pipeline_id', '')
|
|
debug(f'KTV pipeline handler: {pipeline_id}')
|
|
await run_pipeline(pipeline_id)
|
|
return {'pipeline_id': pipeline_id, 'status': 'completed'}
|
|
|
|
|
|
async def on_app_built(app):
|
|
env = ServerEnv()
|
|
longtasks = env.longtasks
|
|
if longtasks:
|
|
schedule_once(0.1, longtasks.run)
|
|
debug('longtasks worker started')
|
|
|
|
|
|
def init():
|
|
env = ServerEnv()
|
|
longtasks = MediaTasks(
|
|
'redis://127.0.0.1:6379',
|
|
'media',
|
|
worker_cnt=4,
|
|
stuck_seconds=1800,
|
|
max_age_hours=24
|
|
)
|
|
longtasks.register('ktv_pipeline', handle_ktv_pipeline)
|
|
env.longtasks = longtasks
|
|
|
|
add_startup(on_app_built)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
webapp(init)
|