import uuid from copy import deepcopy from datetime import datetime from sqlalchemy import delete as sa_delete, select, update as sa_update from sqlalchemy import func from sqlalchemy.ext.asyncio import AsyncSession from app.models.study_setup_config import StudySetupConfig from app.models.study_setup_config_version import StudySetupConfigVersion from app.schemas.study_setup_config import StudySetupConfigData MAIN_BRANCH = "main" def _to_storage(data: StudySetupConfigData | dict) -> dict: if isinstance(data, StudySetupConfigData): return data.model_dump(mode="json") return deepcopy(data) def empty_draft_payload() -> dict: return StudySetupConfigData().model_dump(mode="json") def _build_release_branch_name(base_display_version: int) -> str: return f"release/v{base_display_version}" def _parse_release_branch_base_display_version(branch_name: str) -> int | None: prefix = "release/v" if not branch_name.startswith(prefix): return None suffix = branch_name[len(prefix) :] if not suffix.isdigit(): return None return int(suffix) def _build_version_label(branch_name: str, *, display_version: int, branch_seq: int) -> str: if branch_name == MAIN_BRANCH: return f"v{display_version}" return f"v{display_version}.{branch_seq}" async def _get_snapshot_by_version( db: AsyncSession, *, study_id: uuid.UUID, version: int, ) -> StudySetupConfigVersion | None: result = await db.execute( select(StudySetupConfigVersion).where( StudySetupConfigVersion.study_id == study_id, StudySetupConfigVersion.version == version, ) ) return result.scalar_one_or_none() async def _get_snapshot_by_id( db: AsyncSession, *, snapshot_id: uuid.UUID, ) -> StudySetupConfigVersion | None: result = await db.execute(select(StudySetupConfigVersion).where(StudySetupConfigVersion.id == snapshot_id)) return result.scalar_one_or_none() async def _get_latest_snapshot_in_branch( db: AsyncSession, *, study_id: uuid.UUID, branch_name: str, ) -> StudySetupConfigVersion | None: result = await db.execute( select(StudySetupConfigVersion) .where( StudySetupConfigVersion.study_id == study_id, StudySetupConfigVersion.branch_name == branch_name, ) .order_by( StudySetupConfigVersion.branch_seq.desc(), StudySetupConfigVersion.version.desc(), StudySetupConfigVersion.published_at.desc(), ) .limit(1) ) return result.scalar_one_or_none() async def _get_main_next_display_version(db: AsyncSession, *, study_id: uuid.UUID) -> int: result = await db.execute( select(func.max(StudySetupConfigVersion.display_version)).where( StudySetupConfigVersion.study_id == study_id, StudySetupConfigVersion.branch_name == MAIN_BRANCH, ) ) current_max = result.scalar_one_or_none() or 0 return current_max + 1 async def _get_branch_next_seq( db: AsyncSession, *, study_id: uuid.UUID, branch_name: str, ) -> int: result = await db.execute( select(func.max(StudySetupConfigVersion.branch_seq)).where( StudySetupConfigVersion.study_id == study_id, StudySetupConfigVersion.branch_name == branch_name, ) ) current_max = result.scalar_one_or_none() or 0 return current_max + 1 async def _resolve_current_published_snapshot_id( db: AsyncSession, *, study_id: uuid.UUID, published_config: dict | None, published_by: uuid.UUID | None, published_at: datetime | None, ) -> uuid.UUID | None: if not published_config: return None rows = await list_versions(db, study_id) best_score = -1 best_id: uuid.UUID | None = None for item in rows: if (item.config or {}) != (published_config or {}): continue score = 1 if published_at and item.published_at == published_at: score += 2 if published_by and item.published_by == published_by: score += 1 if score > best_score: best_score = score best_id = item.id return best_id async def get_by_study(db: AsyncSession, study_id: uuid.UUID) -> StudySetupConfig | None: result = await db.execute(select(StudySetupConfig).where(StudySetupConfig.study_id == study_id)) return result.scalar_one_or_none() async def upsert( db: AsyncSession, study_id: uuid.UUID, *, expected_version: int | None, data: StudySetupConfigData | dict, saved_by: uuid.UUID | None, force_draft: bool = False, ) -> tuple[StudySetupConfig | None, bool]: existing = await get_by_study(db, study_id) payload = _to_storage(data) if existing: if expected_version is not None and expected_version != existing.version: return None, True existing.version = existing.version + 1 existing.config = payload existing.saved_by = saved_by if force_draft: existing.publish_status = "DRAFT" elif existing.publish_status == "DRAFT": # Keep draft state sticky until explicit publish API is called. existing.publish_status = "DRAFT" else: existing.publish_status = "PUBLISHED" if existing.published_config == payload else "DRAFT" existing.updated_at = datetime.utcnow() await db.commit() await db.refresh(existing) return existing, False record = StudySetupConfig( study_id=study_id, version=1, config=payload, publish_status="DRAFT", current_branch_name=MAIN_BRANCH, saved_by=saved_by, ) db.add(record) await db.commit() await db.refresh(record) return record, False async def publish( db: AsyncSession, study_id: uuid.UUID, *, expected_version: int | None, published_by: uuid.UUID | None, force_create_snapshot: bool = False, auto_commit: bool = True, ) -> tuple[StudySetupConfig | None, bool]: existing = await get_by_study(db, study_id) if not existing: return None, False if expected_version is not None and expected_version != existing.version: return None, True publish_payload = _to_storage(existing.config or {}) next_version = existing.version + 1 branch_name = (existing.current_branch_name or MAIN_BRANCH).strip() or MAIN_BRANCH current_published_snapshot: StudySetupConfigVersion | None = None current_published_branch = "" if existing.current_published_version_id: current_published_snapshot = await _get_snapshot_by_id(db, snapshot_id=existing.current_published_version_id) if current_published_snapshot: current_published_branch = (current_published_snapshot.branch_name or MAIN_BRANCH).strip() or MAIN_BRANCH branch_changed = bool(current_published_branch and current_published_branch != branch_name) should_create_snapshot = force_create_snapshot or (existing.published_config or {}) != publish_payload or branch_changed existing.version = next_version existing.publish_status = "PUBLISHED" existing.published_config = _to_storage(publish_payload) existing.published_by = published_by existing.published_at = datetime.utcnow() existing.saved_by = published_by existing.updated_at = datetime.utcnow() if should_create_snapshot: parent_snapshot_id = existing.active_branch_base_version_id if parent_snapshot_id is None: latest_branch_snapshot = await _get_latest_snapshot_in_branch(db, study_id=study_id, branch_name=branch_name) parent_snapshot_id = latest_branch_snapshot.id if latest_branch_snapshot else None if branch_name == MAIN_BRANCH: next_display_version = await _get_main_next_display_version(db, study_id=study_id) next_branch_seq = await _get_branch_next_seq(db, study_id=study_id, branch_name=branch_name) else: next_branch_seq = await _get_branch_next_seq(db, study_id=study_id, branch_name=branch_name) branch_base = _parse_release_branch_base_display_version(branch_name) if branch_base is None: # Fall back to parent/main latest when branch name is customized. if parent_snapshot_id: parent_result = await db.execute( select(StudySetupConfigVersion.display_version).where(StudySetupConfigVersion.id == parent_snapshot_id) ) branch_base = parent_result.scalar_one_or_none() if branch_base is None: branch_base = await _get_main_next_display_version(db, study_id=study_id) - 1 next_display_version = branch_base snapshot = StudySetupConfigVersion( study_setup_config_id=existing.id, study_id=study_id, version=next_version, display_version=next_display_version, branch_name=branch_name, branch_seq=next_branch_seq, version_label=_build_version_label( branch_name, display_version=next_display_version, branch_seq=next_branch_seq, ), source_version=next_version - 1, parent_version_id=parent_snapshot_id, config=_to_storage(publish_payload), published_by=published_by, published_at=existing.published_at, ) db.add(snapshot) existing.current_published_version_id = snapshot.id existing.active_branch_base_version_id = snapshot.id else: if not existing.current_published_version_id: existing.current_published_version_id = await _resolve_current_published_snapshot_id( db, study_id=study_id, published_config=existing.published_config, published_by=existing.published_by, published_at=existing.published_at, ) if auto_commit: await db.commit() await db.refresh(existing) else: await db.flush() return existing, False async def list_versions(db: AsyncSession, study_id: uuid.UUID) -> list[StudySetupConfigVersion]: result = await db.execute( select(StudySetupConfigVersion) .where(StudySetupConfigVersion.study_id == study_id) .order_by(StudySetupConfigVersion.version.desc(), StudySetupConfigVersion.published_at.desc()) ) return list(result.scalars().all()) async def set_version_project_snapshot( db: AsyncSession, study_id: uuid.UUID, *, version: int, project_snapshot: dict | None, ) -> bool: target = await _get_snapshot_by_version(db, study_id=study_id, version=version) if not target: return False target.published_project_snapshot = _to_storage(project_snapshot) if project_snapshot else None await db.flush() return True async def rollback_to_version( db: AsyncSession, study_id: uuid.UUID, *, expected_version: int | None, target_version: int, saved_by: uuid.UUID | None, ) -> tuple[StudySetupConfig | None, bool, bool]: existing = await get_by_study(db, study_id) if not existing: return None, False, False if expected_version is not None and expected_version != existing.version: return None, True, False target = await _get_snapshot_by_version(db, study_id=study_id, version=target_version) if not target: return None, False, True target_display_version = target.display_version or 0 target_branch = (target.branch_name or MAIN_BRANCH).strip() or MAIN_BRANCH if target_branch == MAIN_BRANCH: branch_name = _build_release_branch_name(target_display_version) else: branch_name = target_branch now = datetime.utcnow() existing.version = existing.version + 1 target_config = _to_storage(target.config or {}) existing.config = target_config existing.published_config = _to_storage(target.config or {}) existing.published_project_snapshot = ( _to_storage(target.published_project_snapshot) if target.published_project_snapshot else None ) existing.current_branch_name = branch_name existing.active_branch_base_version_id = target.id existing.current_published_version_id = target.id existing.saved_by = saved_by existing.published_by = saved_by existing.published_at = now existing.publish_status = "PUBLISHED" existing.updated_at = now await db.commit() await db.refresh(existing) return existing, False, False async def checkout_branch_draft_from_version( db: AsyncSession, study_id: uuid.UUID, *, expected_version: int | None, target_version: int, saved_by: uuid.UUID | None, ) -> tuple[StudySetupConfig | None, bool, bool]: existing = await get_by_study(db, study_id) if not existing: return None, False, False if expected_version is not None and expected_version != existing.version: return None, True, False target = await _get_snapshot_by_version(db, study_id=study_id, version=target_version) if not target: return None, False, True target_display_version = target.display_version or 0 target_branch = (target.branch_name or MAIN_BRANCH).strip() or MAIN_BRANCH if target_branch == MAIN_BRANCH: checkout_branch = _build_release_branch_name(target_display_version) else: checkout_branch = target_branch now = datetime.utcnow() existing.version = existing.version + 1 existing.config = _to_storage(target.config or {}) existing.current_branch_name = checkout_branch existing.active_branch_base_version_id = target.id existing.saved_by = saved_by existing.publish_status = "PUBLISHED" if (existing.published_config or {}) == (existing.config or {}) else "DRAFT" existing.updated_at = now await db.commit() await db.refresh(existing) return existing, False, False async def clear_draft( db: AsyncSession, study_id: uuid.UUID, *, expected_version: int | None, saved_by: uuid.UUID | None, ) -> tuple[StudySetupConfig | None, bool]: existing = await get_by_study(db, study_id) if not existing: return None, False if expected_version is not None and expected_version != existing.version: return None, True existing.version = existing.version + 1 existing.config = empty_draft_payload() existing.saved_by = saved_by existing.active_branch_base_version_id = None existing.publish_status = "DRAFT" existing.updated_at = datetime.utcnow() await db.commit() await db.refresh(existing) return existing, False async def refill_draft_from_published( db: AsyncSession, study_id: uuid.UUID, *, expected_version: int | None, saved_by: uuid.UUID | None, ) -> tuple[StudySetupConfig | None, bool, bool]: existing = await get_by_study(db, study_id) if not existing: return None, False, False if expected_version is not None and expected_version != existing.version: return None, True, False if not existing.published_config: return None, False, True existing.version = existing.version + 1 existing.config = _to_storage(existing.published_config or {}) existing.saved_by = saved_by if existing.current_published_version_id: snapshot = await _get_snapshot_by_id(db, snapshot_id=existing.current_published_version_id) if snapshot: if not (existing.current_branch_name or "").strip(): existing.current_branch_name = snapshot.branch_name or MAIN_BRANCH if existing.active_branch_base_version_id is None: existing.active_branch_base_version_id = snapshot.id existing.publish_status = "PUBLISHED" if (existing.published_config or {}) == (existing.config or {}) else "DRAFT" existing.updated_at = datetime.utcnow() await db.commit() await db.refresh(existing) return existing, False, False async def merge_to_main( db: AsyncSession, study_id: uuid.UUID, *, expected_version: int | None, source_version: int, merged_by: uuid.UUID | None, auto_commit: bool = True, ) -> tuple[StudySetupConfig | None, bool, bool]: existing = await get_by_study(db, study_id) if not existing: return None, False, False if expected_version is not None and expected_version != existing.version: return None, True, False source = await _get_snapshot_by_version(db, study_id=study_id, version=source_version) if not source: return None, False, True payload = _to_storage(source.config or {}) next_version = existing.version + 1 now = datetime.utcnow() latest_main_snapshot = await _get_latest_snapshot_in_branch(db, study_id=study_id, branch_name=MAIN_BRANCH) parent_snapshot_id = latest_main_snapshot.id if latest_main_snapshot else None next_main_display_version = await _get_main_next_display_version(db, study_id=study_id) next_main_branch_seq = await _get_branch_next_seq(db, study_id=study_id, branch_name=MAIN_BRANCH) snapshot = StudySetupConfigVersion( study_setup_config_id=existing.id, study_id=study_id, version=next_version, display_version=next_main_display_version, branch_name=MAIN_BRANCH, branch_seq=next_main_branch_seq, version_label=_build_version_label( MAIN_BRANCH, display_version=next_main_display_version, branch_seq=next_main_branch_seq, ), source_version=source.version, parent_version_id=parent_snapshot_id, merged_from_version_id=source.id, config=payload, published_project_snapshot=_to_storage(source.published_project_snapshot) if source.published_project_snapshot else None, published_by=merged_by, published_at=now, ) db.add(snapshot) existing.version = next_version existing.config = payload existing.publish_status = "PUBLISHED" existing.published_config = _to_storage(payload) existing.published_project_snapshot = ( _to_storage(source.published_project_snapshot) if source.published_project_snapshot else None ) existing.published_by = merged_by existing.published_at = now existing.saved_by = merged_by existing.current_branch_name = MAIN_BRANCH existing.active_branch_base_version_id = snapshot.id existing.current_published_version_id = snapshot.id existing.updated_at = now if auto_commit: await db.commit() await db.refresh(existing) else: await db.flush() return existing, False, False async def delete_version( db: AsyncSession, study_id: uuid.UUID, *, target_version: int, ) -> bool: target = await _get_snapshot_by_version(db, study_id=study_id, version=target_version) if not target: return False await db.execute( sa_update(StudySetupConfig) .where(StudySetupConfig.study_id == study_id, StudySetupConfig.active_branch_base_version_id == target.id) .values(active_branch_base_version_id=None) ) await db.execute( sa_update(StudySetupConfig) .where(StudySetupConfig.study_id == study_id, StudySetupConfig.current_published_version_id == target.id) .values(current_published_version_id=None) ) await db.delete(target) await db.commit() return True async def delete_by_study(db: AsyncSession, study_id: uuid.UUID) -> None: await db.execute(sa_delete(StudySetupConfig).where(StudySetupConfig.study_id == study_id)) await db.execute(sa_delete(StudySetupConfigVersion).where(StudySetupConfigVersion.study_id == study_id)) await db.commit()