first commit
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
@@ -0,0 +1,19 @@
|
||||
from app.db.models.entities import (
|
||||
ModelProfile,
|
||||
ProviderConfig,
|
||||
RefreshToken,
|
||||
RevokedToken,
|
||||
TestRun,
|
||||
TestResult,
|
||||
User,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ModelProfile",
|
||||
"ProviderConfig",
|
||||
"RefreshToken",
|
||||
"RevokedToken",
|
||||
"TestRun",
|
||||
"TestResult",
|
||||
"User",
|
||||
]
|
||||
@@ -0,0 +1,116 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, Integer, String, Text, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.db.base import Base
|
||||
|
||||
|
||||
def utcnow():
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class ModelProfile(Base):
|
||||
__tablename__ = "model_profiles"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
profile_type: Mapped[str] = mapped_column(String(32), index=True)
|
||||
model_name: Mapped[str] = mapped_column(String(255))
|
||||
endpoint: Mapped[str] = mapped_column(String(1024))
|
||||
capabilities_json: Mapped[str] = mapped_column(Text, default="{}")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
|
||||
|
||||
class ProviderConfig(Base):
|
||||
__tablename__ = "provider_configs"
|
||||
__table_args__ = (UniqueConstraint("profile_type"),)
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
profile_type: Mapped[str] = mapped_column(String(32), index=True)
|
||||
model_name: Mapped[str] = mapped_column(String(255))
|
||||
endpoint: Mapped[str] = mapped_column(String(1024))
|
||||
chat_path: Mapped[str] = mapped_column(String(255), default="/chat/completions")
|
||||
models_path: Mapped[str] = mapped_column(String(255), default="/models")
|
||||
auth_type: Mapped[str] = mapped_column(String(32), default="none")
|
||||
api_key: Mapped[str] = mapped_column(Text, default="")
|
||||
auth_header: Mapped[str] = mapped_column(String(255), default="Authorization")
|
||||
auth_prefix: Mapped[str] = mapped_column(String(255), default="Bearer")
|
||||
verify_ssl: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
|
||||
|
||||
class User(Base):
|
||||
__tablename__ = "users"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
username: Mapped[str] = mapped_column(String(64), unique=True, index=True)
|
||||
password_hash: Mapped[str] = mapped_column(String(255))
|
||||
is_active: Mapped[bool] = mapped_column(default=True)
|
||||
is_admin: Mapped[bool] = mapped_column(default=False)
|
||||
token_version: Mapped[int] = mapped_column(Integer, default=0)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
|
||||
|
||||
class RevokedToken(Base):
|
||||
__tablename__ = "revoked_tokens"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
token_digest: Mapped[str] = mapped_column(String(64), unique=True, index=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
|
||||
|
||||
class RefreshToken(Base):
|
||||
# ponytail: retain rotation history for replay detection; add scheduled expiry cleanup when table growth matters.
|
||||
__tablename__ = "refresh_tokens"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
user_id: Mapped[int] = mapped_column(Integer, index=True)
|
||||
family_id: Mapped[str] = mapped_column(String(36), index=True)
|
||||
token_digest: Mapped[str] = mapped_column(String(64), unique=True, index=True)
|
||||
token_version: Mapped[int] = mapped_column(Integer)
|
||||
expires_at: Mapped[int] = mapped_column(Integer)
|
||||
consumed_at: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
idempotency_key: Mapped[str | None] = mapped_column(String(128), nullable=True)
|
||||
replacement_digest: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
access_expires_at: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
revoked_at: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
|
||||
|
||||
class TestRun(Base):
|
||||
__tablename__ = "test_runs"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
run_type: Mapped[str] = mapped_column(String(32), index=True)
|
||||
status: Mapped[str] = mapped_column(String(32), index=True)
|
||||
profile: Mapped[str] = mapped_column(String(32), default="smoke")
|
||||
selected_count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
idempotency_key: Mapped[str | None] = mapped_column(String(128), unique=True, nullable=True)
|
||||
request_fingerprint: Mapped[str] = mapped_column(String(128), default="")
|
||||
completed_count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
error_count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
summary_json: Mapped[str] = mapped_column(Text, default="{}")
|
||||
capabilities_json: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
|
||||
|
||||
class TestResult(Base):
|
||||
__tablename__ = "test_results"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("run_id", "execution_id", name="uq_test_results_run_execution"),
|
||||
)
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
run_id: Mapped[int] = mapped_column(Integer, index=True)
|
||||
execution_id: Mapped[str] = mapped_column(String(128), index=True)
|
||||
case_kind: Mapped[str] = mapped_column(String(32), index=True)
|
||||
interaction_mode: Mapped[str] = mapped_column(String(32), index=True)
|
||||
execution_status: Mapped[str] = mapped_column(String(32), index=True)
|
||||
verdict: Mapped[str | None] = mapped_column(String(32), index=True, nullable=True)
|
||||
model_response: Mapped[str] = mapped_column(Text, default="")
|
||||
judge_result_json: Mapped[str] = mapped_column(Text, default="{}")
|
||||
error_message: Mapped[str] = mapped_column(Text, default="")
|
||||
audit_context_json: Mapped[str] = mapped_column(Text, default="{}")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow)
|
||||
@@ -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)
|
||||
@@ -0,0 +1,19 @@
|
||||
from collections.abc import Generator
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
connect_args = {"check_same_thread": False} if settings.database_url.startswith("sqlite") else {}
|
||||
engine = create_engine(settings.database_url, connect_args=connect_args, future=True)
|
||||
SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False, class_=Session)
|
||||
|
||||
|
||||
def get_db() -> Generator[Session, None, None]:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
yield db
|
||||
finally:
|
||||
db.close()
|
||||
Reference in New Issue
Block a user