first commit

This commit is contained in:
baozaotumao2025
2026-07-18 21:00:26 +08:00
commit 14722be770
105 changed files with 25004 additions and 0 deletions
View File
+5
View File
@@ -0,0 +1,5 @@
from sqlalchemy.orm import DeclarativeBase
class Base(DeclarativeBase):
pass
+19
View File
@@ -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",
]
+116
View File
@@ -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)
View File
+270
View File
@@ -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)
+19
View File
@@ -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()