first commit
This commit is contained in:
@@ -0,0 +1,270 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user