import json from datetime import datetime, timezone from sqlalchemy import delete, select, update from sqlalchemy.orm import Session from app.db.models import ProviderConfig, RefreshToken, RevokedToken, TestResult, TestRun, User from app.services.auth_service import hash_password class UserRepository: def __init__(self, db: Session): self.db = db def create(self, *, username: str, password: str, is_admin: bool = False) -> User: user = User(username=username, password_hash=hash_password(password), is_admin=is_admin) self.db.add(user) self.db.commit() self.db.refresh(user) return user def get_active_by_username(self, username: str) -> User | None: user = self.db.scalar(select(User).where(User.username == username)) return user if user and user.is_active else None def get_by_username(self, username: str) -> User | None: return self.db.scalar(select(User).where(User.username == username)) def get_active(self, user_id: int) -> User | None: user = self.db.get(User, user_id) return user if user and user.is_active else None def get(self, user_id: int) -> User | None: return self.db.get(User, user_id) def list(self) -> list[User]: return list(self.db.scalars(select(User).order_by(User.id))) def set_active(self, user: User, is_active: bool) -> User: user.is_active = is_active self.db.commit() self.db.refresh(user) return user def reset_password(self, user: User, password: str) -> None: user.password_hash = hash_password(password) user.token_version += 1 self.db.commit() def delete(self, user: User) -> None: self.db.delete(user) self.db.commit() def revoke_token(self, token_digest: str) -> None: if not self.db.scalar( select(RevokedToken.id).where(RevokedToken.token_digest == token_digest) ): self.db.add(RevokedToken(token_digest=token_digest)) self.db.commit() def is_token_revoked(self, token_digest: str) -> bool: return ( self.db.scalar(select(RevokedToken.id).where(RevokedToken.token_digest == token_digest)) is not None ) def create_refresh_token( self, *, user_id: int, family_id: str, token_digest: str, token_version: int, expires_at: int, access_expires_at: int, ) -> RefreshToken: token = RefreshToken( user_id=user_id, family_id=family_id, token_digest=token_digest, token_version=token_version, expires_at=expires_at, access_expires_at=access_expires_at, ) self.db.add(token) self.db.commit() self.db.refresh(token) return token def get_refresh_token(self, digest: str) -> RefreshToken | None: return self.db.scalar(select(RefreshToken).where(RefreshToken.token_digest == digest)) def rotate_refresh_token( self, token: RefreshToken, *, idempotency_key: str, replacement_digest: str, access_expires_at: int, refresh_expires_at: int, consumed_at: int, ) -> bool: claimed = self.db.execute( update(RefreshToken) .where( RefreshToken.id == token.id, RefreshToken.consumed_at.is_(None), RefreshToken.revoked_at.is_(None), ) .values( consumed_at=consumed_at, idempotency_key=idempotency_key, replacement_digest=replacement_digest, access_expires_at=access_expires_at, expires_at=refresh_expires_at, ) ).rowcount if not claimed: self.db.rollback() return False self.db.add( RefreshToken( user_id=token.user_id, family_id=token.family_id, token_digest=replacement_digest, token_version=token.token_version, expires_at=refresh_expires_at, access_expires_at=access_expires_at, ) ) self.db.commit() return True def revoke_refresh_family(self, family_id: str, revoked_at: int) -> None: self.db.execute( update(RefreshToken) .where(RefreshToken.family_id == family_id, RefreshToken.revoked_at.is_(None)) .values(revoked_at=revoked_at) ) self.db.commit() def is_refresh_family_active(self, family_id: str) -> bool: return ( self.db.scalar( select(RefreshToken.id) .where(RefreshToken.family_id == family_id, RefreshToken.revoked_at.is_(None)) .limit(1) ) is not None ) class ProviderConfigRepository: def __init__(self, db: Session): self.db = db def create(self, **values) -> ProviderConfig: row = ProviderConfig(**values) self.db.add(row) self.db.commit() self.db.refresh(row) return row def get_by_type(self, profile_type: str) -> ProviderConfig | None: return self.db.scalar( select(ProviderConfig).where(ProviderConfig.profile_type == profile_type) ) def list(self) -> list[ProviderConfig]: return list(self.db.scalars(select(ProviderConfig).order_by(ProviderConfig.id))) def update(self, row: ProviderConfig, **values) -> ProviderConfig: for key, value in values.items(): setattr(row, key, value) row.updated_at = datetime.now(timezone.utc) self.db.commit() self.db.refresh(row) return row def delete(self, row: ProviderConfig) -> None: self.db.delete(row) self.db.commit() class TestRunRepository: def __init__(self, db: Session): self.db = db def create( self, *, run_type: str, profile: str, selected_count: int, idempotency_key: str | None = None, request_fingerprint: str = "", ) -> TestRun: row = TestRun( run_type=run_type, status="pending", profile=profile, selected_count=selected_count, idempotency_key=idempotency_key, request_fingerprint=request_fingerprint, ) self.db.add(row) self.db.commit() self.db.refresh(row) return row def get(self, run_id: int) -> TestRun | None: return self.db.get(TestRun, run_id) def get_by_idempotency_key(self, key: str) -> TestRun | None: return self.db.scalar(select(TestRun).where(TestRun.idempotency_key == key)) def list(self, limit: int = 50) -> list[TestRun]: return list(self.db.scalars(select(TestRun).order_by(TestRun.id.desc()).limit(limit))) def update_status(self, run: TestRun, status: str, *, summary: dict | None = None) -> None: run.status = status run.updated_at = datetime.now(timezone.utc) if summary is not None: run.summary_json = json.dumps(summary, ensure_ascii=False) self.db.commit() def save_capabilities(self, run: TestRun, capabilities: dict) -> None: run.capabilities_json = json.dumps(capabilities, ensure_ascii=False) self.db.commit() def delete(self, run: TestRun) -> None: self.db.execute(delete(TestResult).where(TestResult.run_id == run.id)) self.db.delete(run) self.db.commit() class ResultRepository: def __init__(self, db: Session): self.db = db def add(self, **kwargs) -> TestResult: row = TestResult(**kwargs) self.db.add(row) self.db.commit() self.db.refresh(row) return row def list_by_run(self, run_id: int, limit: int = 1000) -> list[TestResult]: stmt = select(TestResult).where(TestResult.run_id == run_id).limit(limit) return list(self.db.scalars(stmt)) def execution_ids_by_run(self, run_id: int) -> set[str]: stmt = select(TestResult.execution_id).where(TestResult.run_id == run_id) return set(self.db.scalars(stmt)) def delete_errors_by_run(self, run_id: int) -> int: result = self.db.execute( delete(TestResult).where( TestResult.run_id == run_id, TestResult.execution_status == "error" ) ) return result.rowcount or 0 def delete_by_run_and_execution(self, run_id: int, execution_id: str) -> bool: result = self.db.execute( delete(TestResult).where( TestResult.run_id == run_id, TestResult.execution_id == execution_id ) ) return bool(result.rowcount)