from datetime import datetime from sqlalchemy import and_, or_, select from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session from app.db.models import Activity, ActivityStatus, SyncRun, SyncRunStatus, SyncUser, utcnow class UserRepository: def __init__(self, session: Session) -> None: self.session = session def create(self, **values) -> SyncUser: user = SyncUser(**values) self.session.add(user) self.session.commit() return user def get(self, user_id: int) -> SyncUser | None: return self.session.get(SyncUser, user_id) def list_enabled(self) -> list[SyncUser]: return list(self.session.scalars(select(SyncUser).where(SyncUser.enabled.is_(True)).order_by(SyncUser.id))) def list_all(self) -> list[SyncUser]: return list(self.session.scalars(select(SyncUser).order_by(SyncUser.name))) def update(self, user: SyncUser, **values) -> SyncUser: for key, value in values.items(): setattr(user, key, value) self.session.commit() return user class ActivityRepository: def __init__(self, session: Session) -> None: self.session = session def _require(self, activity_id: int) -> Activity: activity = self.session.get(Activity, activity_id) if activity is None: raise ValueError(f"activity {activity_id} not found") return activity def get(self, activity_id: int) -> Activity | None: return self.session.get(Activity, activity_id) def get_or_create_discovered( self, *, user_id: int, mywhoosh_activity_id: str, activity_name: str, activity_timestamp: datetime | None, ) -> tuple[Activity, bool]: existing = self.session.scalar( select(Activity).where( Activity.user_id == user_id, Activity.mywhoosh_activity_id == mywhoosh_activity_id, ) ) if existing is not None: return existing, False activity = Activity( user_id=user_id, mywhoosh_activity_id=mywhoosh_activity_id, activity_name=activity_name, activity_timestamp=activity_timestamp, status=ActivityStatus.DISCOVERED, last_completed_stage=ActivityStatus.DISCOVERED, ) self.session.add(activity) try: self.session.commit() except IntegrityError: self.session.rollback() existing = self.session.scalar( select(Activity).where( Activity.user_id == user_id, Activity.mywhoosh_activity_id == mywhoosh_activity_id, ) ) if existing is None: raise return existing, False return activity, True def mark_downloaded(self, activity_id: int, path: str) -> Activity: activity = self._require(activity_id) activity.source_fit_path = path activity.status = ActivityStatus.DOWNLOADED activity.last_completed_stage = ActivityStatus.DOWNLOADED activity.last_error = None activity.retryable = True self.session.commit() return activity def mark_converted(self, activity_id: int, path: str) -> Activity: activity = self._require(activity_id) activity.converted_fit_path = path activity.status = ActivityStatus.CONVERTED activity.last_completed_stage = ActivityStatus.CONVERTED activity.last_error = None activity.retryable = True self.session.commit() return activity def mark_imported(self, activity_id: int, garmin_activity_id: str | None) -> Activity: activity = self._require(activity_id) activity.status = ActivityStatus.IMPORTED activity.last_completed_stage = ActivityStatus.IMPORTED activity.garmin_activity_id = garmin_activity_id activity.last_error = None activity.retryable = False self.session.commit() return activity def mark_duplicate(self, activity_id: int) -> Activity: activity = self._require(activity_id) activity.status = ActivityStatus.DUPLICATE activity.last_completed_stage = ActivityStatus.DUPLICATE activity.last_error = None activity.retryable = False self.session.commit() return activity def mark_failed(self, activity_id: int, error: str, *, retryable: bool) -> Activity: activity = self._require(activity_id) activity.status = ActivityStatus.FAILED activity.last_error = error[:2000] activity.retryable = retryable self.session.commit() return activity def list_pending_for_user(self, user_id: int) -> list[Activity]: return list( self.session.scalars( select(Activity).where( Activity.user_id == user_id, or_( Activity.status.in_([ActivityStatus.DISCOVERED, ActivityStatus.DOWNLOADED, ActivityStatus.CONVERTED]), and_(Activity.status == ActivityStatus.FAILED, Activity.retryable.is_(True)), ), ).order_by(Activity.id) ) ) class SyncRunRepository: def __init__(self, session: Session) -> None: self.session = session def start(self, user_id: int) -> SyncRun: sync_run = SyncRun(user_id=user_id, status=SyncRunStatus.RUNNING) self.session.add(sync_run) self.session.commit() return sync_run def get(self, sync_run_id: int) -> SyncRun | None: return self.session.get(SyncRun, sync_run_id) def finish( self, sync_run_id: int, *, status: SyncRunStatus, discovered: int, imported: int, skipped: int, failed: int, summary_error: str | None = None, ) -> SyncRun: sync_run = self.session.get(SyncRun, sync_run_id) if sync_run is None: raise ValueError(f"sync_run {sync_run_id} not found") sync_run.finished_at = utcnow() sync_run.status = status sync_run.discovered_count = discovered sync_run.imported_count = imported sync_run.skipped_count = skipped sync_run.failed_count = failed sync_run.summary_error = summary_error[:2000] if summary_error else None self.session.commit() return sync_run