Files
ai-safety-platform/app/db/repositories/repositories.py
T
baozaotumao2025 14722be770 first commit
2026-07-18 21:00:26 +08:00

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)