Skip to content

Commit 8edcb8c

Browse files
authored
Merge pull request #3 from adeptvin1/feature/fix-group-status-update
fix race condition
2 parents 086740b + 0070d09 commit 8edcb8c

1 file changed

Lines changed: 100 additions & 12 deletions

File tree

client/api/endpoints.py

Lines changed: 100 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,8 @@ def default(self, obj):
4141

4242
# Сет для отслеживания активных экспериментов
4343
active_experiments: Set[str] = set()
44+
# Блокировка для предотвращения race conditions при работе с active_experiments
45+
active_experiments_lock = asyncio.Lock()
4446

4547

4648
# Функция для преобразования документа MongoDB в JSON-сериализуемый формат
@@ -396,21 +398,52 @@ async def check_and_update_group_status(experiment_id: str, db):
396398

397399
async def run_experiment(experiment_id: str):
398400
"""Фоновая задача для выполнения эксперимента с динамической нагрузкой."""
399-
if experiment_id in active_experiments:
400-
logger.warning(f"Experiment {experiment_id} is already running")
401-
return
402-
403-
active_experiments.add(experiment_id)
401+
# Используем блокировку для предотвращения race conditions
402+
async with active_experiments_lock:
403+
if experiment_id in active_experiments:
404+
logger.warning(f"Experiment {experiment_id} is already running")
405+
return
406+
407+
# Проверяем статус эксперимента в БД перед запуском
408+
try:
409+
db = await get_database()
410+
experiment = await db.experiments.find_one({"_id": ObjectId(experiment_id)})
411+
if not experiment:
412+
logger.error(f"Experiment {experiment_id} not found")
413+
return
414+
415+
# Если эксперимент уже завершен или провалился, не запускаем его снова
416+
exp_state = experiment.get("state")
417+
if exp_state in [ExperimentState.COMPLETED, ExperimentState.FAILED]:
418+
logger.info(f"Experiment {experiment_id} is already {exp_state}, skipping")
419+
return
420+
421+
# Если эксперимент уже выполняется (по статусу в БД), не запускаем его снова
422+
if exp_state == ExperimentState.RUNNING:
423+
logger.warning(f"Experiment {experiment_id} is already running (status in DB)")
424+
return
425+
426+
except Exception as e:
427+
logger.error(f"Error checking experiment {experiment_id} status: {str(e)}")
428+
return
429+
430+
# Добавляем в активные эксперименты только после всех проверок
431+
active_experiments.add(experiment_id)
432+
433+
# Инициализация переменных вне блока блокировки
404434
semaphore = asyncio.Semaphore(MAX_CONCURRENT_TASKS)
405435
save_semaphore = asyncio.Semaphore(1) # Семафор для сохранения результатов
406-
db = None # Инициализируем для доступа в finally блоке
407436

408437
try:
438+
# Получаем эксперимент еще раз для работы
409439
db = await get_database()
410-
experiment = await with_retry(get_experiment_by_id, experiment_id, db)
440+
experiment = await with_retry(_get_experiment_from_db, experiment_id, db)
411441

412442
if not experiment:
413443
logger.error(f"Experiment {experiment_id} not found")
444+
# Удаляем из активных, если эксперимент не найден
445+
async with active_experiments_lock:
446+
active_experiments.discard(experiment_id)
414447
return
415448

416449
settings = experiment["settings"]
@@ -598,7 +631,9 @@ async def run_experiment(experiment_id: str):
598631
# Логируем ошибку, но не прерываем выполнение
599632
logger.error(f"Failed to check group status in finally block: {str(group_check_error)}")
600633

601-
active_experiments.remove(experiment_id)
634+
# Удаляем эксперимент из активных с блокировкой
635+
async with active_experiments_lock:
636+
active_experiments.discard(experiment_id) # Используем discard вместо remove для безопасности
602637

603638

604639
@router.get("/experiment_stats")
@@ -852,20 +887,73 @@ async def manage_group(
852887
if not group:
853888
raise HTTPException(status_code=404, detail="Group not found")
854889

890+
current_group_state = group.get("state", ExperimentState.PENDING)
891+
855892
new_state = ""
856893
if state == "start":
894+
# Проверяем, не завершена ли уже группа
895+
if current_group_state == ExperimentState.COMPLETED:
896+
logger.warning(f"Group {group_id} is already completed, cannot start")
897+
raise HTTPException(
898+
status_code=400,
899+
detail="Group is already completed and cannot be started again"
900+
)
901+
902+
# Проверяем, не запущена ли уже группа
903+
if current_group_state == ExperimentState.RUNNING:
904+
logger.warning(f"Group {group_id} is already running")
905+
# Возвращаем текущий статус без повторного запуска
906+
return {"_id": group_id, "state": current_group_state, "message": "Group is already running"}
907+
857908
new_state = ExperimentState.RUNNING
858-
# Запускаем все эксперименты в группе
859-
for exp_id in group["experiment_ids"]:
860-
asyncio.create_task(run_experiment(exp_id))
909+
910+
# Запускаем только те эксперименты, которые еще не завершены и не запущены
911+
experiment_ids = group.get("experiment_ids", [])
912+
started_count = 0
913+
skipped_count = 0
914+
915+
for exp_id in experiment_ids:
916+
try:
917+
# Проверяем статус эксперимента перед запуском
918+
experiment = await db.experiments.find_one({"_id": ObjectId(exp_id)})
919+
if not experiment:
920+
logger.warning(f"Experiment {exp_id} not found, skipping")
921+
skipped_count += 1
922+
continue
923+
924+
exp_state = experiment.get("state")
925+
926+
# Пропускаем уже завершенные или провалившиеся эксперименты
927+
if exp_state in [ExperimentState.COMPLETED, ExperimentState.FAILED]:
928+
logger.debug(f"Experiment {exp_id} is already {exp_state}, skipping")
929+
skipped_count += 1
930+
continue
931+
932+
# Проверяем, не запущен ли уже эксперимент
933+
async with active_experiments_lock:
934+
if exp_id in active_experiments:
935+
logger.debug(f"Experiment {exp_id} is already running, skipping")
936+
skipped_count += 1
937+
continue
938+
939+
# Запускаем эксперимент
940+
asyncio.create_task(run_experiment(exp_id))
941+
started_count += 1
942+
943+
except Exception as e:
944+
logger.error(f"Error starting experiment {exp_id}: {str(e)}")
945+
skipped_count += 1
946+
947+
logger.info(f"Group {group_id}: started {started_count} experiments, skipped {skipped_count}")
948+
861949
elif state == "pause":
862950
new_state = ExperimentState.PAUSED
863951
elif state == "stop":
864952
new_state = ExperimentState.COMPLETED
865953

866954
await db.groups.update_one(
867955
{"_id": ObjectId(group_id)},
868-
{"$set": {"state": new_state}}
956+
{"$set": {"state": new_state, "updated_at": datetime.now()}}
869957
)
870958

871959
return {"_id": group_id, "state": new_state}

0 commit comments

Comments
 (0)