271 lines
8.7 KiB
Python
271 lines
8.7 KiB
Python
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)
|