feat: implement auth, projects, and frontend foundation
This commit is contained in:
@@ -0,0 +1,36 @@
|
||||
[alembic]
|
||||
script_location = alembic
|
||||
prepend_sys_path = .
|
||||
sqlalchemy.url = postgresql+asyncpg://headquarter:headquarter@postgres:5432/headquarter
|
||||
|
||||
[loggers]
|
||||
keys = root,sqlalchemy,alembic
|
||||
|
||||
[handlers]
|
||||
keys = console
|
||||
|
||||
[formatters]
|
||||
keys = generic
|
||||
|
||||
[logger_root]
|
||||
level = WARN
|
||||
handlers = console
|
||||
|
||||
[logger_sqlalchemy]
|
||||
level = WARN
|
||||
handlers =
|
||||
qualname = sqlalchemy.engine
|
||||
|
||||
[logger_alembic]
|
||||
level = INFO
|
||||
handlers =
|
||||
qualname = alembic
|
||||
|
||||
[handler_console]
|
||||
class = StreamHandler
|
||||
args = (sys.stderr,)
|
||||
level = NOTSET
|
||||
formatter = generic
|
||||
|
||||
[formatter_generic]
|
||||
format = %(levelname)-5.5s [%(name)s] %(message)s
|
||||
@@ -0,0 +1,65 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from sqlalchemy import pool
|
||||
from sqlalchemy.engine import Connection
|
||||
from sqlalchemy.ext.asyncio import async_engine_from_config
|
||||
|
||||
from src.config import Settings
|
||||
from src.models import Base
|
||||
|
||||
config = context.config
|
||||
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name)
|
||||
|
||||
settings = Settings()
|
||||
config.set_main_option("sqlalchemy.url", settings.database_url)
|
||||
|
||||
target_metadata = Base.metadata
|
||||
|
||||
|
||||
def run_migrations_offline() -> None:
|
||||
context.configure(
|
||||
url=settings.database_url,
|
||||
target_metadata=target_metadata,
|
||||
literal_binds=True,
|
||||
dialect_opts={"paramstyle": "named"},
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
def do_run_migrations(connection: Connection) -> None:
|
||||
context.configure(connection=connection, target_metadata=target_metadata)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
async def run_async_migrations() -> None:
|
||||
connectable = async_engine_from_config(
|
||||
config.get_section(config.config_ini_section, {}),
|
||||
prefix="sqlalchemy.",
|
||||
poolclass=pool.NullPool,
|
||||
)
|
||||
|
||||
async with connectable.connect() as connection:
|
||||
await connection.run_sync(do_run_migrations)
|
||||
|
||||
await connectable.dispose()
|
||||
|
||||
|
||||
def run_migrations_online() -> None:
|
||||
import asyncio
|
||||
|
||||
asyncio.run(run_async_migrations())
|
||||
|
||||
|
||||
if context.is_offline_mode():
|
||||
run_migrations_offline()
|
||||
else:
|
||||
run_migrations_online()
|
||||
@@ -0,0 +1,25 @@
|
||||
"""${message}
|
||||
|
||||
Revision ID: ${up_revision}
|
||||
Revises: ${down_revision | comma,n}
|
||||
Create Date: ${create_date}
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
${imports if imports else ""}
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = ${repr(up_revision)}
|
||||
down_revision = ${repr(down_revision)}
|
||||
branch_labels = ${repr(branch_labels)}
|
||||
depends_on = ${repr(depends_on)}
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
${upgrades if upgrades else "pass"}
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
${downgrades if downgrades else "pass"}
|
||||
@@ -0,0 +1,106 @@
|
||||
"""initial schema
|
||||
|
||||
Revision ID: 0001_initial_schema
|
||||
Revises:
|
||||
Create Date: 2026-05-17 00:00:00.000000
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
revision = "0001_initial_schema"
|
||||
down_revision = None
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
TABLE_NAMES = [
|
||||
"users",
|
||||
"ssh_keys",
|
||||
"projects",
|
||||
"git_repositories",
|
||||
"user_configs",
|
||||
]
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"users",
|
||||
sa.Column("email", sa.String(length=255), nullable=False),
|
||||
sa.Column("name", sa.String(length=255), nullable=False),
|
||||
sa.Column("authentik_id", sa.String(length=255), nullable=False),
|
||||
sa.Column("avatar_url", sa.String(length=1024), nullable=True),
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("authentik_id"),
|
||||
sa.UniqueConstraint("email"),
|
||||
)
|
||||
op.create_index(op.f("ix_users_authentik_id"), "users", ["authentik_id"], unique=True)
|
||||
op.create_index(op.f("ix_users_email"), "users", ["email"], unique=True)
|
||||
|
||||
op.create_table(
|
||||
"ssh_keys",
|
||||
sa.Column("name", sa.String(length=255), nullable=False),
|
||||
sa.Column("public_key", sa.Text(), nullable=False),
|
||||
sa.Column("private_key_encrypted", sa.Text(), nullable=False),
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("project_id", postgresql.UUID(as_uuid=True), nullable=True),
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"]),
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"projects",
|
||||
sa.Column("name", sa.String(length=255), nullable=False),
|
||||
sa.Column("description", sa.Text(), nullable=True),
|
||||
sa.Column("owner_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("default_ssh_key_id", postgresql.UUID(as_uuid=True), nullable=True),
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.ForeignKeyConstraint(["default_ssh_key_id"], ["ssh_keys.id"]),
|
||||
sa.ForeignKeyConstraint(["owner_id"], ["users.id"]),
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"git_repositories",
|
||||
sa.Column("name", sa.String(length=255), nullable=False),
|
||||
sa.Column("path", sa.String(length=1024), nullable=False),
|
||||
sa.Column("project_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("owner_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("is_mirror", sa.Boolean(), nullable=False),
|
||||
sa.Column("remote_url", sa.String(length=1024), nullable=True),
|
||||
sa.Column("last_push", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.ForeignKeyConstraint(["owner_id"], ["users.id"]),
|
||||
sa.ForeignKeyConstraint(["project_id"], ["projects.id"]),
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"user_configs",
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("config", postgresql.JSONB(astext_type=sa.Text()), nullable=False),
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("user_id"),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"]),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("user_configs")
|
||||
op.drop_table("git_repositories")
|
||||
op.drop_table("projects")
|
||||
op.drop_table("ssh_keys")
|
||||
op.drop_index(op.f("ix_users_email"), table_name="users")
|
||||
op.drop_index(op.f("ix_users_authentik_id"), table_name="users")
|
||||
op.drop_table("users")
|
||||
@@ -0,0 +1,58 @@
|
||||
"""add refresh tokens table
|
||||
|
||||
Revision ID: 0002_refresh_tokens
|
||||
Revises: 0001_initial_schema
|
||||
Create Date: 2026-05-17 00:00:01.000000
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
revision = "0002_refresh_tokens"
|
||||
down_revision = "0001_initial_schema"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
connection = op.get_bind()
|
||||
inspector = sa.inspect(connection)
|
||||
|
||||
if not inspector.has_table("refresh_tokens"):
|
||||
op.create_table(
|
||||
"refresh_tokens",
|
||||
sa.Column("user_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("token_hash", sa.String(length=255), nullable=False),
|
||||
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("revoked_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("user_agent", sa.String(length=512), nullable=True),
|
||||
sa.Column("ip_address", sa.String(length=64), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("token_hash"),
|
||||
)
|
||||
|
||||
existing_indexes = {index["name"] for index in inspector.get_indexes("refresh_tokens")}
|
||||
user_index = op.f("ix_refresh_tokens_user_id")
|
||||
expires_index = op.f("ix_refresh_tokens_expires_at")
|
||||
if user_index not in existing_indexes:
|
||||
op.create_index(user_index, "refresh_tokens", ["user_id"], unique=False)
|
||||
if expires_index not in existing_indexes:
|
||||
op.create_index(expires_index, "refresh_tokens", ["expires_at"], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
connection = op.get_bind()
|
||||
inspector = sa.inspect(connection)
|
||||
if inspector.has_table("refresh_tokens"):
|
||||
existing_indexes = {index["name"] for index in inspector.get_indexes("refresh_tokens")}
|
||||
expires_index = op.f("ix_refresh_tokens_expires_at")
|
||||
user_index = op.f("ix_refresh_tokens_user_id")
|
||||
if expires_index in existing_indexes:
|
||||
op.drop_index(expires_index, table_name="refresh_tokens")
|
||||
if user_index in existing_indexes:
|
||||
op.drop_index(user_index, table_name="refresh_tokens")
|
||||
op.drop_table("refresh_tokens")
|
||||
@@ -26,3 +26,6 @@ dev = [
|
||||
"ruff>=0.1.0",
|
||||
"httpx>=0.25.0",
|
||||
]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
pythonpath = ["."]
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Headquarter API package."""
|
||||
@@ -0,0 +1,3 @@
|
||||
from src.api.auth import router as auth_router
|
||||
|
||||
__all__ = ["auth_router"]
|
||||
@@ -0,0 +1,190 @@
|
||||
from secrets import token_urlsafe
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import AsyncGenerator, Literal, cast
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Cookie, Depends, HTTPException, Response, status
|
||||
from fastapi.responses import RedirectResponse
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.cookies import build_cookie_options
|
||||
from src.auth.jwt_service import decode_access_token, mint_access_token
|
||||
from src.auth.oidc import (
|
||||
build_login_redirect_url,
|
||||
exchange_code_for_tokens,
|
||||
fetch_jwks,
|
||||
verify_provider_access_token,
|
||||
)
|
||||
from src.auth.refresh_store import create_refresh_token, revoke_refresh_token, rotate_refresh_token
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.models.user import User
|
||||
|
||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
|
||||
|
||||
async def get_db_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with SessionLocal() as session:
|
||||
yield session
|
||||
|
||||
|
||||
@router.get("/login")
|
||||
async def login() -> RedirectResponse:
|
||||
settings = Settings()
|
||||
redirect_uri = "http://localhost:8000/auth/callback"
|
||||
state = token_urlsafe(24)
|
||||
location = build_login_redirect_url(
|
||||
settings=settings,
|
||||
redirect_uri=redirect_uri,
|
||||
state=state,
|
||||
nonce=token_urlsafe(16),
|
||||
)
|
||||
response = RedirectResponse(location)
|
||||
response.set_cookie("auth_state", state, httponly=True, samesite="lax")
|
||||
return response
|
||||
|
||||
|
||||
@router.get("/callback")
|
||||
async def callback(
|
||||
code: str,
|
||||
state: str,
|
||||
response: Response,
|
||||
auth_state: str | None = Cookie(default=None),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict[str, str]:
|
||||
if auth_state is None or auth_state != state:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid state")
|
||||
|
||||
settings = Settings()
|
||||
redirect_uri = "http://localhost:8000/auth/callback"
|
||||
async with httpx.AsyncClient() as client:
|
||||
token_payload = await exchange_code_for_tokens(
|
||||
settings=settings,
|
||||
code=code,
|
||||
redirect_uri=redirect_uri,
|
||||
client=client,
|
||||
)
|
||||
jwks = await fetch_jwks(settings=settings, client=client)
|
||||
|
||||
provider_claims = verify_provider_access_token(
|
||||
settings=settings,
|
||||
token=token_payload["access_token"],
|
||||
jwks=jwks,
|
||||
)
|
||||
|
||||
authentik_id = str(provider_claims["sub"])
|
||||
email = str(provider_claims.get("email", f"{authentik_id}@authentik.local"))
|
||||
name = str(provider_claims.get("name", email))
|
||||
|
||||
user = await session.scalar(select(User).where(User.authentik_id == authentik_id))
|
||||
if user is None:
|
||||
user = User(email=email, name=name, authentik_id=authentik_id, avatar_url=None)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
else:
|
||||
user.email = email
|
||||
user.name = name
|
||||
await session.commit()
|
||||
|
||||
access_token = mint_access_token(
|
||||
settings=settings,
|
||||
subject=str(user.id),
|
||||
email=user.email,
|
||||
name=user.name,
|
||||
expires_at=datetime.now(UTC) + timedelta(minutes=settings.access_token_ttl_minutes),
|
||||
)
|
||||
refresh_token, _ = await create_refresh_token(
|
||||
session=session,
|
||||
user_id=user.id,
|
||||
expires_at=datetime.now(UTC) + timedelta(days=settings.refresh_token_ttl_days),
|
||||
user_agent=None,
|
||||
ip_address=None,
|
||||
)
|
||||
|
||||
cookie_options = build_cookie_options(settings)
|
||||
cookie_samesite = cast(Literal["lax", "strict", "none"], cookie_options["samesite"])
|
||||
cookie_secure = bool(cookie_options["secure"])
|
||||
response.set_cookie("access_token", access_token, httponly=True, samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.set_cookie("refresh_token", refresh_token, httponly=True, samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.delete_cookie("auth_state", samesite="lax")
|
||||
|
||||
return {"sub": str(user.id), "email": user.email, "name": user.name}
|
||||
|
||||
|
||||
@router.post("/refresh")
|
||||
async def refresh(
|
||||
response: Response,
|
||||
refresh_token: str | None = Cookie(default=None),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict[str, str]:
|
||||
if not refresh_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing refresh token")
|
||||
|
||||
settings = Settings()
|
||||
try:
|
||||
rotated_raw_token, rotated_record = await rotate_refresh_token(
|
||||
session=session,
|
||||
raw_token=refresh_token,
|
||||
user_agent=None,
|
||||
ip_address=None,
|
||||
)
|
||||
except ValueError as error:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=str(error)) from error
|
||||
|
||||
user = await session.get(User, rotated_record.user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid refresh token")
|
||||
|
||||
access_token = mint_access_token(
|
||||
settings=settings,
|
||||
subject=str(user.id),
|
||||
email=user.email,
|
||||
name=user.name,
|
||||
expires_at=datetime.now(UTC) + timedelta(minutes=settings.access_token_ttl_minutes),
|
||||
)
|
||||
|
||||
cookie_options = build_cookie_options(settings)
|
||||
cookie_samesite = cast(Literal["lax", "strict", "none"], cookie_options["samesite"])
|
||||
cookie_secure = bool(cookie_options["secure"])
|
||||
|
||||
response.set_cookie("access_token", access_token, httponly=True, samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.set_cookie("refresh_token", rotated_raw_token, httponly=True, samesite=cookie_samesite, secure=cookie_secure)
|
||||
|
||||
return {"sub": str(user.id), "email": user.email, "name": user.name}
|
||||
|
||||
|
||||
@router.post("/logout")
|
||||
async def logout(
|
||||
response: Response,
|
||||
refresh_token: str | None = Cookie(default=None),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> dict[str, str]:
|
||||
settings = Settings()
|
||||
cookie_options = build_cookie_options(settings)
|
||||
cookie_samesite = cast(Literal["lax", "strict", "none"], cookie_options["samesite"])
|
||||
cookie_secure = bool(cookie_options["secure"])
|
||||
|
||||
if refresh_token:
|
||||
try:
|
||||
await revoke_refresh_token(session=session, raw_token=refresh_token)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
response.delete_cookie("access_token", samesite=cookie_samesite, secure=cookie_secure)
|
||||
response.delete_cookie("refresh_token", samesite=cookie_samesite, secure=cookie_secure)
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
@router.get("/me")
|
||||
async def me(access_token: str | None = Cookie(default=None)) -> dict[str, str]:
|
||||
if not access_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing access token")
|
||||
|
||||
claims = decode_access_token(settings=Settings(), token=access_token)
|
||||
return {
|
||||
"sub": str(claims["sub"]),
|
||||
"email": str(claims["email"]),
|
||||
"name": str(claims["name"]),
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
import uuid
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Cookie, Depends, HTTPException, Response, status
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.auth.jwt_service import decode_access_token
|
||||
from src.config import Settings
|
||||
from src.database import SessionLocal
|
||||
from src.models.project import Project
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user import User
|
||||
|
||||
router = APIRouter(prefix="/projects", tags=["projects"])
|
||||
|
||||
|
||||
async def get_db_session():
|
||||
async with SessionLocal() as session:
|
||||
yield session
|
||||
|
||||
|
||||
async def get_current_user_id(
|
||||
access_token: Annotated[str | None, Cookie()] = None,
|
||||
) -> uuid.UUID:
|
||||
if not access_token:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing access token")
|
||||
|
||||
try:
|
||||
claims = decode_access_token(settings=Settings(), token=access_token)
|
||||
return uuid.UUID(str(claims["sub"]))
|
||||
except Exception:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid access token")
|
||||
|
||||
|
||||
async def _get_user(session: AsyncSession, user_id: uuid.UUID) -> User:
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||
return user
|
||||
|
||||
|
||||
class ProjectCreate(BaseModel):
|
||||
name: str
|
||||
description: str | None = None
|
||||
|
||||
|
||||
class ProjectUpdate(BaseModel):
|
||||
name: str | None = None
|
||||
description: str | None = None
|
||||
|
||||
|
||||
class ProjectResponse(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: uuid.UUID
|
||||
name: str
|
||||
description: str | None
|
||||
owner_id: uuid.UUID
|
||||
default_ssh_key_id: uuid.UUID | None
|
||||
|
||||
|
||||
class SetDefaultSSHKeyRequest(BaseModel):
|
||||
ssh_key_id: uuid.UUID
|
||||
|
||||
|
||||
@router.post("", response_model=ProjectResponse, status_code=status.HTTP_201_CREATED)
|
||||
async def create_project(
|
||||
data: ProjectCreate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Project:
|
||||
user = await _get_user(session, user_id)
|
||||
project = Project(
|
||||
name=data.name,
|
||||
description=data.description,
|
||||
owner_id=user.id,
|
||||
default_ssh_key_id=None,
|
||||
)
|
||||
session.add(project)
|
||||
await session.commit()
|
||||
await session.refresh(project)
|
||||
return project
|
||||
|
||||
|
||||
@router.get("", response_model=list[ProjectResponse])
|
||||
async def list_projects(
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> list[Project]:
|
||||
user = await _get_user(session, user_id)
|
||||
result = await session.execute(select(Project).where(Project.owner_id == user.id))
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def _get_owned_project(
|
||||
project_id: uuid.UUID,
|
||||
user_id: uuid.UUID,
|
||||
session: AsyncSession,
|
||||
) -> Project:
|
||||
project = await session.get(Project, project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="project not found")
|
||||
if project.owner_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="not project owner")
|
||||
return project
|
||||
|
||||
|
||||
@router.patch("/{project_id}", response_model=ProjectResponse)
|
||||
async def update_project(
|
||||
project_id: uuid.UUID,
|
||||
data: ProjectUpdate,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Project:
|
||||
await _get_user(session, user_id)
|
||||
project = await _get_owned_project(project_id, user_id, session)
|
||||
|
||||
if data.name is not None:
|
||||
project.name = data.name
|
||||
if data.description is not None:
|
||||
project.description = data.description
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(project)
|
||||
return project
|
||||
|
||||
|
||||
@router.delete("/{project_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
async def delete_project(
|
||||
project_id: uuid.UUID,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Response:
|
||||
await _get_user(session, user_id)
|
||||
project = await _get_owned_project(project_id, user_id, session)
|
||||
await session.delete(project)
|
||||
await session.commit()
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
@router.patch("/{project_id}/default-ssh-key", response_model=ProjectResponse)
|
||||
async def set_default_ssh_key(
|
||||
project_id: uuid.UUID,
|
||||
data: SetDefaultSSHKeyRequest,
|
||||
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> Project:
|
||||
user = await _get_user(session, user_id)
|
||||
project = await _get_owned_project(project_id, user_id, session)
|
||||
|
||||
ssh_key = await session.get(SSHKey, data.ssh_key_id)
|
||||
if ssh_key is None or ssh_key.user_id != user.id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="invalid ssh key",
|
||||
)
|
||||
|
||||
project.default_ssh_key_id = data.ssh_key_id
|
||||
await session.commit()
|
||||
await session.refresh(project)
|
||||
return project
|
||||
@@ -0,0 +1,12 @@
|
||||
from src.auth.cookies import build_cookie_options
|
||||
from src.auth.jwt_service import decode_access_token, mint_access_token
|
||||
from src.auth.oidc import build_login_redirect_url
|
||||
from src.auth.refresh_store import hash_refresh_token
|
||||
|
||||
__all__ = [
|
||||
"build_cookie_options",
|
||||
"build_login_redirect_url",
|
||||
"decode_access_token",
|
||||
"hash_refresh_token",
|
||||
"mint_access_token",
|
||||
]
|
||||
@@ -0,0 +1,9 @@
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def build_cookie_options(settings: Settings) -> dict[str, str | bool]:
|
||||
return {
|
||||
"httponly": True,
|
||||
"secure": settings.cookie_secure,
|
||||
"samesite": settings.cookie_samesite,
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
from datetime import datetime
|
||||
|
||||
from jose import jwt # type: ignore[import-untyped]
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def mint_access_token(
|
||||
*,
|
||||
settings: Settings,
|
||||
subject: str,
|
||||
email: str,
|
||||
name: str,
|
||||
expires_at: datetime,
|
||||
) -> str:
|
||||
payload = {
|
||||
"sub": subject,
|
||||
"email": email,
|
||||
"name": name,
|
||||
"exp": expires_at,
|
||||
}
|
||||
return jwt.encode(payload, settings.jwt_secret, algorithm=settings.jwt_algorithm)
|
||||
|
||||
|
||||
def decode_access_token(*, settings: Settings, token: str) -> dict[str, str | int]:
|
||||
claims = jwt.decode(token, settings.jwt_secret, algorithms=[settings.jwt_algorithm])
|
||||
return dict(claims)
|
||||
@@ -0,0 +1,77 @@
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import httpx
|
||||
from jose import jwt # type: ignore[import-untyped]
|
||||
|
||||
from src.config import Settings
|
||||
|
||||
|
||||
def build_login_redirect_url(
|
||||
*,
|
||||
settings: Settings,
|
||||
redirect_uri: str,
|
||||
state: str,
|
||||
nonce: str,
|
||||
) -> str:
|
||||
query = urlencode(
|
||||
{
|
||||
"response_type": "code",
|
||||
"client_id": settings.authentik_client_id,
|
||||
"redirect_uri": redirect_uri,
|
||||
"scope": "openid profile email",
|
||||
"state": state,
|
||||
"nonce": nonce,
|
||||
}
|
||||
)
|
||||
return f"{settings.authentik_authorize_url}?{query}"
|
||||
|
||||
|
||||
async def exchange_code_for_tokens(
|
||||
*,
|
||||
settings: Settings,
|
||||
code: str,
|
||||
redirect_uri: str,
|
||||
client: httpx.AsyncClient,
|
||||
) -> dict[str, str]:
|
||||
response = await client.post(
|
||||
settings.authentik_token_url,
|
||||
data={
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"redirect_uri": redirect_uri,
|
||||
"client_id": settings.authentik_client_id,
|
||||
"client_secret": settings.authentik_client_secret,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
return {
|
||||
"access_token": payload["access_token"],
|
||||
"refresh_token": payload["refresh_token"],
|
||||
}
|
||||
|
||||
|
||||
async def fetch_jwks(*, settings: Settings, client: httpx.AsyncClient) -> dict[str, list[dict[str, str]]]:
|
||||
response = await client.get(settings.authentik_jwks_url)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
return {"keys": payload["keys"]}
|
||||
|
||||
|
||||
def verify_provider_access_token(
|
||||
*,
|
||||
settings: Settings,
|
||||
token: str,
|
||||
jwks: dict[str, list[dict[str, str]]],
|
||||
) -> dict[str, str | int]:
|
||||
unverified_header = jwt.get_unverified_header(token)
|
||||
key_id = unverified_header["kid"]
|
||||
jwk_key = next(key for key in jwks["keys"] if key.get("kid") == key_id)
|
||||
claims = jwt.decode(
|
||||
token,
|
||||
jwk_key,
|
||||
algorithms=[jwk_key.get("alg", "HS256")],
|
||||
audience=settings.authentik_audience,
|
||||
issuer=settings.authentik_issuer,
|
||||
)
|
||||
return dict(claims)
|
||||
@@ -0,0 +1,79 @@
|
||||
from datetime import UTC, datetime
|
||||
from hashlib import sha256
|
||||
from secrets import token_urlsafe
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.models.refresh_token import RefreshToken
|
||||
|
||||
|
||||
def hash_refresh_token(raw_token: str) -> str:
|
||||
return sha256(raw_token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
async def create_refresh_token(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
user_id: object,
|
||||
expires_at: datetime,
|
||||
user_agent: str | None,
|
||||
ip_address: str | None,
|
||||
) -> tuple[str, RefreshToken]:
|
||||
raw_token = token_urlsafe(48)
|
||||
record = RefreshToken(
|
||||
user_id=user_id,
|
||||
token_hash=hash_refresh_token(raw_token),
|
||||
expires_at=expires_at,
|
||||
created_at=datetime.now(UTC),
|
||||
user_agent=user_agent,
|
||||
ip_address=ip_address,
|
||||
)
|
||||
session.add(record)
|
||||
await session.commit()
|
||||
await session.refresh(record)
|
||||
return raw_token, record
|
||||
|
||||
|
||||
async def rotate_refresh_token(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
raw_token: str,
|
||||
user_agent: str | None,
|
||||
ip_address: str | None,
|
||||
) -> tuple[str, RefreshToken]:
|
||||
existing_hash = hash_refresh_token(raw_token)
|
||||
existing = await session.scalar(
|
||||
select(RefreshToken).where(
|
||||
RefreshToken.token_hash == existing_hash,
|
||||
RefreshToken.revoked_at.is_(None),
|
||||
)
|
||||
)
|
||||
if existing is None:
|
||||
raise ValueError("refresh token not found")
|
||||
if existing.expires_at <= datetime.now(UTC):
|
||||
raise ValueError("refresh token expired")
|
||||
|
||||
existing.revoked_at = datetime.now(UTC)
|
||||
await session.flush()
|
||||
|
||||
return await create_refresh_token(
|
||||
session=session,
|
||||
user_id=existing.user_id,
|
||||
expires_at=existing.expires_at,
|
||||
user_agent=user_agent,
|
||||
ip_address=ip_address,
|
||||
)
|
||||
|
||||
|
||||
async def revoke_refresh_token(*, session: AsyncSession, raw_token: str) -> bool:
|
||||
token_hash = hash_refresh_token(raw_token)
|
||||
existing = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||
if existing is None:
|
||||
return False
|
||||
if existing.revoked_at is not None:
|
||||
return True
|
||||
|
||||
existing.revoked_at = datetime.now(UTC)
|
||||
await session.commit()
|
||||
return True
|
||||
@@ -0,0 +1,62 @@
|
||||
from pydantic import Field
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
def build_database_url(
|
||||
*,
|
||||
user: str,
|
||||
password: str,
|
||||
host: str,
|
||||
port: int,
|
||||
database: str,
|
||||
) -> str:
|
||||
return f"postgresql+asyncpg://{user}:{password}@{host}:{port}/{database}"
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
app_env: str = "development"
|
||||
database_url_override: str | None = Field(default=None, alias="DATABASE_URL")
|
||||
postgres_user: str = "headquarter"
|
||||
postgres_password: str = "headquarter"
|
||||
postgres_host: str = "postgres"
|
||||
postgres_port: int = 5432
|
||||
postgres_db: str = "headquarter"
|
||||
|
||||
authentik_client_id: str = "headquarter-web"
|
||||
authentik_client_secret: str = "change-me"
|
||||
authentik_authorize_url: str = "https://authentik.local/application/o/authorize/"
|
||||
authentik_token_url: str = "https://authentik.local/application/o/token/"
|
||||
authentik_jwks_url: str = "https://authentik.local/application/o/headquarter-web/jwks/"
|
||||
authentik_issuer: str = "https://authentik.local/application/o/headquarter-web/"
|
||||
authentik_audience: str = "headquarter-web"
|
||||
|
||||
jwt_secret: str = "change-me-jwt-secret"
|
||||
jwt_algorithm: str = "HS256"
|
||||
access_token_ttl_minutes: int = 15
|
||||
refresh_token_ttl_days: int = 7
|
||||
|
||||
model_config = SettingsConfigDict(env_file=".env", extra="ignore", populate_by_name=True)
|
||||
|
||||
@property
|
||||
def database_url(self) -> str:
|
||||
if self.database_url_override:
|
||||
return self.database_url_override
|
||||
|
||||
return build_database_url(
|
||||
user=self.postgres_user,
|
||||
password=self.postgres_password,
|
||||
host=self.postgres_host,
|
||||
port=self.postgres_port,
|
||||
database=self.postgres_db,
|
||||
)
|
||||
|
||||
@property
|
||||
def cookie_secure(self) -> bool:
|
||||
return self.app_env == "production"
|
||||
|
||||
@property
|
||||
def cookie_samesite(self) -> str:
|
||||
if self.app_env == "production":
|
||||
return "strict"
|
||||
|
||||
return "lax"
|
||||
@@ -0,0 +1,15 @@
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
from sqlalchemy.pool import NullPool
|
||||
|
||||
from src.config import Settings, build_database_url
|
||||
|
||||
|
||||
settings = Settings()
|
||||
engine = create_async_engine(
|
||||
settings.database_url,
|
||||
future=True,
|
||||
poolclass=NullPool,
|
||||
)
|
||||
SessionLocal = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
|
||||
|
||||
__all__ = ["SessionLocal", "build_database_url", "engine", "settings"]
|
||||
@@ -0,0 +1,8 @@
|
||||
from fastapi import FastAPI
|
||||
|
||||
from src.api.auth import router as auth_router
|
||||
from src.api.projects import router as projects_router
|
||||
|
||||
app = FastAPI(title="Headquarter API")
|
||||
app.include_router(auth_router)
|
||||
app.include_router(projects_router)
|
||||
@@ -0,0 +1,9 @@
|
||||
from src.models.base import Base
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.project import Project
|
||||
from src.models.refresh_token import RefreshToken
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user import User
|
||||
from src.models.user_config import UserConfig
|
||||
|
||||
__all__ = ["Base", "GitRepository", "Project", "RefreshToken", "SSHKey", "User", "UserConfig"]
|
||||
@@ -0,0 +1,24 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, func
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
|
||||
|
||||
class UUIDPrimaryKeyMixin:
|
||||
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
|
||||
|
||||
class TimestampMixin:
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now(), nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
onupdate=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
@@ -0,0 +1,28 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, ForeignKey, String
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.project import Project
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class GitRepository(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "git_repositories"
|
||||
|
||||
name: Mapped[str] = mapped_column(String(255))
|
||||
path: Mapped[str] = mapped_column(String(1024))
|
||||
project_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("projects.id"), nullable=False)
|
||||
owner_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
|
||||
is_mirror: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
remote_url: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||
last_push: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
project: Mapped["Project"] = relationship(back_populates="repositories")
|
||||
owner: Mapped["User"] = relationship()
|
||||
@@ -0,0 +1,31 @@
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import ForeignKey, String, Text
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class Project(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "projects"
|
||||
|
||||
name: Mapped[str] = mapped_column(String(255))
|
||||
description: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
owner_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
|
||||
default_ssh_key_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("ssh_keys.id"),
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
owner: Mapped["User"] = relationship(back_populates="projects")
|
||||
repositories: Mapped[list["GitRepository"]] = relationship(back_populates="project")
|
||||
default_ssh_key: Mapped["SSHKey | None"] = relationship(foreign_keys=[default_ssh_key_id])
|
||||
ssh_keys: Mapped[list["SSHKey"]] = relationship(back_populates="project", foreign_keys="SSHKey.project_id")
|
||||
@@ -0,0 +1,24 @@
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import DateTime, ForeignKey, String
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class RefreshToken(UUIDPrimaryKeyMixin, Base):
|
||||
__tablename__ = "refresh_tokens"
|
||||
|
||||
user_id: Mapped[str] = mapped_column(ForeignKey("users.id"), nullable=False, index=True)
|
||||
token_hash: Mapped[str] = mapped_column(String(255), unique=True, nullable=False)
|
||||
expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, index=True)
|
||||
revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
user_agent: Mapped[str | None] = mapped_column(String(512), nullable=True)
|
||||
ip_address: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
user: Mapped["User"] = relationship(back_populates="refresh_tokens")
|
||||
@@ -0,0 +1,25 @@
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import ForeignKey, String, Text
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.project import Project
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class SSHKey(UUIDPrimaryKeyMixin, Base):
|
||||
__tablename__ = "ssh_keys"
|
||||
|
||||
name: Mapped[str] = mapped_column(String(255))
|
||||
public_key: Mapped[str] = mapped_column(Text)
|
||||
private_key_encrypted: Mapped[str] = mapped_column(Text)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
|
||||
project_id: Mapped[uuid.UUID | None] = mapped_column(UUID(as_uuid=True), ForeignKey("projects.id"), nullable=True)
|
||||
|
||||
user: Mapped["User"] = relationship(back_populates="ssh_keys")
|
||||
project: Mapped["Project | None"] = relationship(back_populates="ssh_keys", foreign_keys=[project_id])
|
||||
@@ -0,0 +1,26 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import String
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.project import Project
|
||||
from src.models.refresh_token import RefreshToken
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user_config import UserConfig
|
||||
|
||||
|
||||
class User(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "users"
|
||||
|
||||
email: Mapped[str] = mapped_column(String(255), unique=True, index=True)
|
||||
name: Mapped[str] = mapped_column(String(255))
|
||||
authentik_id: Mapped[str] = mapped_column(String(255), unique=True, index=True)
|
||||
avatar_url: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||
|
||||
projects: Mapped[list["Project"]] = relationship(back_populates="owner")
|
||||
refresh_tokens: Mapped[list["RefreshToken"]] = relationship(back_populates="user")
|
||||
ssh_keys: Mapped[list["SSHKey"]] = relationship(back_populates="user")
|
||||
user_config: Mapped["UserConfig | None"] = relationship(back_populates="user", uselist=False)
|
||||
@@ -0,0 +1,20 @@
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import ForeignKey
|
||||
from sqlalchemy.dialects.postgresql import JSONB, UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
class UserConfig(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
__tablename__ = "user_configs"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False, unique=True)
|
||||
config: Mapped[dict[str, object]] = mapped_column(JSONB, default=dict, nullable=False)
|
||||
|
||||
user: Mapped["User"] = relationship(back_populates="user_config")
|
||||
@@ -0,0 +1 @@
|
||||
"""Utility scripts for the API package."""
|
||||
@@ -0,0 +1,48 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.database import SessionLocal
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
def build_seed_user() -> Mapping[str, str | None]:
|
||||
return {
|
||||
"email": "dev@headquarter.local",
|
||||
"name": "Development User",
|
||||
"authentik_id": "dev-authentik-user",
|
||||
"avatar_url": None,
|
||||
}
|
||||
|
||||
|
||||
async def seed_database(session: AsyncSession) -> User:
|
||||
payload = build_seed_user()
|
||||
existing_user = await session.scalar(select(User).where(User.email == payload["email"]))
|
||||
|
||||
if existing_user is not None:
|
||||
return existing_user
|
||||
|
||||
user = User(**payload)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def run() -> None:
|
||||
async with SessionLocal() as session:
|
||||
user = await seed_database(session)
|
||||
print({"user_id": str(user.id), "email": user.email})
|
||||
|
||||
|
||||
def main() -> None:
|
||||
import asyncio
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,218 @@
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
import asyncio
|
||||
import importlib
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
import pytest
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from src.auth.jwt_service import mint_access_token
|
||||
from src.config import Settings, build_database_url
|
||||
from src.models import Base
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def configure_local_database(monkeypatch) -> None:
|
||||
local_url = build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
monkeypatch.setenv("DATABASE_URL", local_url)
|
||||
|
||||
|
||||
def _prepare_auth_test_db() -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(text("TRUNCATE TABLE refresh_tokens, users RESTART IDENTITY CASCADE"))
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def _load_app():
|
||||
import src.database as database_module
|
||||
import src.api.auth as auth_module
|
||||
import src.main as main_module
|
||||
|
||||
importlib.reload(database_module)
|
||||
importlib.reload(auth_module)
|
||||
importlib.reload(main_module)
|
||||
return main_module.app
|
||||
|
||||
|
||||
def _insert_user_for_refresh(user_id: str) -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker
|
||||
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
async with session_factory() as session:
|
||||
user = User(
|
||||
id=uuid.UUID(user_id),
|
||||
email="refresh@headquarter.local",
|
||||
name="Refresh User",
|
||||
authentik_id="refresh-user",
|
||||
avatar_url=None,
|
||||
)
|
||||
await session.merge(user)
|
||||
await session.commit()
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_login_redirects_to_authentik_authorize_endpoint() -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
client = TestClient(app)
|
||||
response = client.get("/auth/login", follow_redirects=False)
|
||||
|
||||
assert response.status_code == 307
|
||||
assert "response_type=code" in response.headers["location"]
|
||||
|
||||
|
||||
def test_me_returns_401_without_access_cookie() -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
client = TestClient(app)
|
||||
response = client.get("/auth/me")
|
||||
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_me_returns_user_payload_with_valid_access_cookie() -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
settings = Settings()
|
||||
token = mint_access_token(
|
||||
settings=settings,
|
||||
subject=str(uuid.uuid4()),
|
||||
email="dev@headquarter.local",
|
||||
name="Dev User",
|
||||
expires_at=datetime.now(UTC) + timedelta(minutes=15),
|
||||
)
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", token)
|
||||
response = client.get("/auth/me")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["email"] == "dev@headquarter.local"
|
||||
|
||||
|
||||
def test_logout_clears_auth_cookies() -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("refresh_token", "opaque-token")
|
||||
response = client.post("/auth/logout")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert "access_token=" in response.headers.get("set-cookie", "")
|
||||
|
||||
|
||||
def test_callback_rejects_mismatched_state() -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("auth_state", "expected")
|
||||
response = client.get("/auth/callback?code=test-code&state=wrong")
|
||||
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_callback_sets_auth_cookies_after_success(monkeypatch) -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
async def fake_exchange_code_for_tokens(*, settings, code, redirect_uri, client):
|
||||
return {"access_token": "provider-access", "refresh_token": "provider-refresh"}
|
||||
|
||||
def fake_verify_provider_access_token(*, settings, token, jwks):
|
||||
return {"sub": "auth-sub-1", "email": "callback@headquarter.local", "name": "Callback User"}
|
||||
|
||||
async def fake_fetch_jwks(*, settings, client):
|
||||
return {"keys": []}
|
||||
|
||||
monkeypatch.setattr("src.api.auth.exchange_code_for_tokens", fake_exchange_code_for_tokens)
|
||||
monkeypatch.setattr("src.api.auth.verify_provider_access_token", fake_verify_provider_access_token)
|
||||
monkeypatch.setattr("src.api.auth.fetch_jwks", fake_fetch_jwks)
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("auth_state", "good-state")
|
||||
response = client.get("/auth/callback?code=valid-code&state=good-state")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["email"] == "callback@headquarter.local"
|
||||
set_cookie_header = response.headers.get("set-cookie", "")
|
||||
assert "access_token=" in set_cookie_header
|
||||
assert "refresh_token=" in set_cookie_header
|
||||
|
||||
|
||||
def test_refresh_rotates_cookie_and_returns_user_payload(monkeypatch) -> None:
|
||||
_prepare_auth_test_db()
|
||||
_insert_user_for_refresh("7f4b7ad8-c4ce-4d1b-8c83-7ce0f4f66dfb")
|
||||
app = _load_app()
|
||||
|
||||
async def fake_rotate_refresh_token(*, session, raw_token, user_agent, ip_address):
|
||||
class StoredToken:
|
||||
user_id = uuid.UUID("7f4b7ad8-c4ce-4d1b-8c83-7ce0f4f66dfb")
|
||||
|
||||
return "new-refresh-token", StoredToken()
|
||||
|
||||
monkeypatch.setattr("src.api.auth.rotate_refresh_token", fake_rotate_refresh_token)
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("refresh_token", "old-refresh-token")
|
||||
response = client.post("/auth/refresh")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["sub"] == "7f4b7ad8-c4ce-4d1b-8c83-7ce0f4f66dfb"
|
||||
assert "refresh_token=" in response.headers.get("set-cookie", "")
|
||||
|
||||
|
||||
def test_refresh_returns_401_for_invalid_refresh_token(monkeypatch) -> None:
|
||||
_prepare_auth_test_db()
|
||||
app = _load_app()
|
||||
|
||||
async def fake_rotate_refresh_token(*, session, raw_token, user_agent, ip_address):
|
||||
raise ValueError("refresh token not found")
|
||||
|
||||
monkeypatch.setattr("src.api.auth.rotate_refresh_token", fake_rotate_refresh_token)
|
||||
|
||||
client = TestClient(app)
|
||||
client.cookies.set("refresh_token", "invalid")
|
||||
response = client.post("/auth/refresh")
|
||||
|
||||
assert response.status_code == 401
|
||||
@@ -0,0 +1,211 @@
|
||||
from datetime import UTC, datetime, timedelta
|
||||
import base64
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
|
||||
from src.auth.cookies import build_cookie_options
|
||||
from src.auth.jwt_service import decode_access_token, mint_access_token
|
||||
from src.auth.oidc import build_login_redirect_url, exchange_code_for_tokens, verify_provider_access_token
|
||||
from src.auth.refresh_store import create_refresh_token, hash_refresh_token, revoke_refresh_token, rotate_refresh_token
|
||||
from src.config import Settings, build_database_url
|
||||
from src.models import Base
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
def test_cookie_options_follow_environment_defaults(monkeypatch) -> None:
|
||||
monkeypatch.setenv("APP_ENV", "development")
|
||||
dev_settings = Settings()
|
||||
dev_options = build_cookie_options(dev_settings)
|
||||
|
||||
monkeypatch.setenv("APP_ENV", "production")
|
||||
prod_settings = Settings()
|
||||
prod_options = build_cookie_options(prod_settings)
|
||||
|
||||
assert dev_options["httponly"] is True
|
||||
assert dev_options["secure"] is False
|
||||
assert dev_options["samesite"] == "lax"
|
||||
assert prod_options["secure"] is True
|
||||
assert prod_options["samesite"] == "strict"
|
||||
|
||||
|
||||
def test_login_redirect_url_contains_required_oidc_params() -> None:
|
||||
settings = Settings()
|
||||
|
||||
url = build_login_redirect_url(
|
||||
settings=settings,
|
||||
redirect_uri="http://localhost:8000/auth/callback",
|
||||
state="state-123",
|
||||
nonce="nonce-123",
|
||||
)
|
||||
|
||||
assert "response_type=code" in url
|
||||
assert "client_id=headquarter-web" in url
|
||||
assert "scope=openid+profile+email" in url
|
||||
assert "state=state-123" in url
|
||||
assert "nonce=nonce-123" in url
|
||||
|
||||
|
||||
def test_mint_and_decode_internal_access_token_round_trip() -> None:
|
||||
settings = Settings()
|
||||
expires_at = datetime.now(UTC) + timedelta(minutes=15)
|
||||
|
||||
token = mint_access_token(
|
||||
settings=settings,
|
||||
subject="user-123",
|
||||
email="dev@headquarter.local",
|
||||
name="Dev User",
|
||||
expires_at=expires_at,
|
||||
)
|
||||
|
||||
claims = decode_access_token(settings=settings, token=token)
|
||||
|
||||
assert claims["sub"] == "user-123"
|
||||
assert claims["email"] == "dev@headquarter.local"
|
||||
assert claims["name"] == "Dev User"
|
||||
assert "exp" in claims
|
||||
|
||||
|
||||
def test_refresh_token_hash_is_deterministic_and_non_reversible() -> None:
|
||||
raw_token = "refresh-token-abc"
|
||||
|
||||
first_hash = hash_refresh_token(raw_token)
|
||||
second_hash = hash_refresh_token(raw_token)
|
||||
|
||||
assert first_hash == second_hash
|
||||
assert first_hash != raw_token
|
||||
assert len(first_hash) == 64
|
||||
|
||||
|
||||
def test_decode_access_token_rejects_invalid_signature() -> None:
|
||||
settings = Settings()
|
||||
other_settings = Settings(jwt_secret="different-secret")
|
||||
expires_at = datetime.now(UTC) + timedelta(minutes=15)
|
||||
|
||||
token = mint_access_token(
|
||||
settings=other_settings,
|
||||
subject="user-123",
|
||||
email="dev@headquarter.local",
|
||||
name="Dev User",
|
||||
expires_at=expires_at,
|
||||
)
|
||||
|
||||
with pytest.raises(Exception):
|
||||
decode_access_token(settings=settings, token=token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_code_for_tokens_posts_expected_payload() -> None:
|
||||
settings = Settings()
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url == httpx.URL(settings.authentik_token_url)
|
||||
payload = dict(httpx.QueryParams(request.content.decode("utf-8")))
|
||||
assert payload["grant_type"] == "authorization_code"
|
||||
assert payload["code"] == "auth-code"
|
||||
assert payload["redirect_uri"] == "http://localhost:8000/auth/callback"
|
||||
return httpx.Response(200, json={"access_token": "provider-token", "refresh_token": "provider-refresh"})
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(transport=transport) as client:
|
||||
token_payload = await exchange_code_for_tokens(
|
||||
settings=settings,
|
||||
code="auth-code",
|
||||
redirect_uri="http://localhost:8000/auth/callback",
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert token_payload["access_token"] == "provider-token"
|
||||
|
||||
|
||||
def test_verify_provider_access_token_with_jwks_oct_key() -> None:
|
||||
settings = Settings(authentik_audience="headquarter-web", authentik_issuer="https://authentik.local/")
|
||||
shared_secret = b"shared-secret-123"
|
||||
jwks = {
|
||||
"keys": [
|
||||
{
|
||||
"kty": "oct",
|
||||
"alg": "HS256",
|
||||
"k": base64.urlsafe_b64encode(shared_secret).decode("utf-8").rstrip("="),
|
||||
"kid": "kid-1",
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
from jose import jwt # type: ignore[import-untyped]
|
||||
|
||||
token = jwt.encode(
|
||||
{
|
||||
"sub": "authentik-user",
|
||||
"iss": settings.authentik_issuer,
|
||||
"aud": settings.authentik_audience,
|
||||
"exp": int((datetime.now(UTC) + timedelta(minutes=5)).timestamp()),
|
||||
},
|
||||
shared_secret,
|
||||
algorithm="HS256",
|
||||
headers={"kid": "kid-1"},
|
||||
)
|
||||
|
||||
claims = verify_provider_access_token(settings=settings, token=token, jwks=jwks)
|
||||
|
||||
assert claims["sub"] == "authentik-user"
|
||||
|
||||
|
||||
TEST_DATABASE_URL = build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def db_session() -> AsyncSession:
|
||||
engine = create_async_engine(TEST_DATABASE_URL)
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
|
||||
async with session_factory() as session:
|
||||
await session.execute(text("TRUNCATE TABLE refresh_tokens, users RESTART IDENTITY CASCADE"))
|
||||
await session.commit()
|
||||
yield session
|
||||
await session.rollback()
|
||||
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_store_create_rotate_and_revoke(db_session: AsyncSession) -> None:
|
||||
user = User(email="dev-auth@headquarter.local", name="Dev Auth", authentik_id="auth-dev", avatar_url=None)
|
||||
db_session.add(user)
|
||||
await db_session.commit()
|
||||
await db_session.refresh(user)
|
||||
|
||||
raw_refresh_token, stored_token = await create_refresh_token(
|
||||
session=db_session,
|
||||
user_id=user.id,
|
||||
expires_at=datetime.now(UTC) + timedelta(days=7),
|
||||
user_agent="pytest",
|
||||
ip_address="127.0.0.1",
|
||||
)
|
||||
assert raw_refresh_token
|
||||
assert stored_token.revoked_at is None
|
||||
|
||||
rotated_raw, rotated_stored = await rotate_refresh_token(
|
||||
session=db_session,
|
||||
raw_token=raw_refresh_token,
|
||||
user_agent="pytest-rotated",
|
||||
ip_address="127.0.0.2",
|
||||
)
|
||||
assert rotated_raw != raw_refresh_token
|
||||
assert rotated_stored.revoked_at is None
|
||||
assert stored_token.revoked_at is not None
|
||||
|
||||
revoked = await revoke_refresh_token(session=db_session, raw_token=rotated_raw)
|
||||
assert revoked is True
|
||||
@@ -0,0 +1,59 @@
|
||||
from src.config import Settings
|
||||
from src.database import build_database_url
|
||||
|
||||
|
||||
def test_settings_default_database_url_uses_asyncpg() -> None:
|
||||
settings = Settings()
|
||||
|
||||
assert settings.database_url == "postgresql+asyncpg://headquarter:headquarter@postgres:5432/headquarter"
|
||||
|
||||
|
||||
def test_build_database_url_uses_explicit_values() -> None:
|
||||
url = build_database_url(
|
||||
user="user",
|
||||
password="pass",
|
||||
host="db",
|
||||
port=5433,
|
||||
database="app",
|
||||
)
|
||||
|
||||
assert url == "postgresql+asyncpg://user:pass@db:5433/app"
|
||||
|
||||
|
||||
def test_settings_prefers_explicit_database_url_env(monkeypatch) -> None:
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql+asyncpg://local:local@localhost:5432/localdb")
|
||||
|
||||
settings = Settings()
|
||||
|
||||
assert settings.database_url == "postgresql+asyncpg://local:local@localhost:5432/localdb"
|
||||
|
||||
|
||||
def test_auth_settings_have_secure_defaults() -> None:
|
||||
settings = Settings()
|
||||
|
||||
assert settings.authentik_client_id == "headquarter-web"
|
||||
assert settings.authentik_client_secret == "change-me"
|
||||
assert settings.authentik_authorize_url.endswith("/application/o/authorize/")
|
||||
assert settings.authentik_token_url.endswith("/application/o/token/")
|
||||
assert settings.authentik_jwks_url.endswith("/application/o/headquarter-web/jwks/")
|
||||
assert settings.jwt_algorithm == "HS256"
|
||||
assert settings.access_token_ttl_minutes == 15
|
||||
assert settings.refresh_token_ttl_days == 7
|
||||
|
||||
|
||||
def test_cookie_policy_is_strict_in_production(monkeypatch) -> None:
|
||||
monkeypatch.setenv("APP_ENV", "production")
|
||||
|
||||
settings = Settings()
|
||||
|
||||
assert settings.cookie_secure is True
|
||||
assert settings.cookie_samesite == "strict"
|
||||
|
||||
|
||||
def test_cookie_policy_is_relaxed_for_local_dev(monkeypatch) -> None:
|
||||
monkeypatch.setenv("APP_ENV", "development")
|
||||
|
||||
settings = Settings()
|
||||
|
||||
assert settings.cookie_secure is False
|
||||
assert settings.cookie_samesite == "lax"
|
||||
@@ -0,0 +1,35 @@
|
||||
from importlib.util import module_from_spec, spec_from_file_location
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def test_initial_migration_defines_all_core_tables() -> None:
|
||||
migration_path = Path(__file__).resolve().parents[1] / "alembic" / "versions" / "0001_initial_schema.py"
|
||||
spec = spec_from_file_location("initial_schema", migration_path)
|
||||
|
||||
assert spec is not None
|
||||
assert spec.loader is not None
|
||||
|
||||
module = module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
assert module.TABLE_NAMES == [
|
||||
"users",
|
||||
"ssh_keys",
|
||||
"projects",
|
||||
"git_repositories",
|
||||
"user_configs",
|
||||
]
|
||||
|
||||
|
||||
def test_refresh_tokens_migration_has_expected_revision_chain() -> None:
|
||||
migration_path = Path(__file__).resolve().parents[1] / "alembic" / "versions" / "0002_refresh_tokens.py"
|
||||
spec = spec_from_file_location("refresh_tokens", migration_path)
|
||||
|
||||
assert spec is not None
|
||||
assert spec.loader is not None
|
||||
|
||||
module = module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
assert module.revision == "0002_refresh_tokens"
|
||||
assert module.down_revision == "0001_initial_schema"
|
||||
@@ -0,0 +1,153 @@
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
|
||||
from src.config import build_database_url
|
||||
from src.models import Base
|
||||
from src.models.base import TimestampMixin, UUIDPrimaryKeyMixin
|
||||
from src.models.git_repository import GitRepository
|
||||
from src.models.project import Project
|
||||
from src.models.refresh_token import RefreshToken
|
||||
from src.models.ssh_key import SSHKey
|
||||
from src.models.user import User
|
||||
from src.models.user_config import UserConfig
|
||||
|
||||
|
||||
TEST_DATABASE_URL = build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def db_session() -> AsyncIterator[AsyncSession]:
|
||||
engine = create_async_engine(TEST_DATABASE_URL)
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
|
||||
async with session_factory() as session:
|
||||
table_rows = await session.execute(
|
||||
text(
|
||||
"SELECT tablename FROM pg_tables "
|
||||
"WHERE schemaname = 'public' "
|
||||
"AND tablename = ANY(:table_names)"
|
||||
),
|
||||
{
|
||||
"table_names": [
|
||||
"refresh_tokens",
|
||||
"user_configs",
|
||||
"git_repositories",
|
||||
"projects",
|
||||
"ssh_keys",
|
||||
"users",
|
||||
]
|
||||
},
|
||||
)
|
||||
existing_tables = [row[0] for row in table_rows]
|
||||
if existing_tables:
|
||||
await session.execute(text(f"TRUNCATE TABLE {', '.join(existing_tables)} RESTART IDENTITY CASCADE"))
|
||||
await session.commit()
|
||||
yield session
|
||||
await session.rollback()
|
||||
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def test_base_metadata_collects_declared_tables() -> None:
|
||||
assert isinstance(Base.metadata.tables, dict)
|
||||
|
||||
|
||||
def test_shared_mixins_define_expected_columns() -> None:
|
||||
assert "id" in UUIDPrimaryKeyMixin.__dict__
|
||||
assert "created_at" in TimestampMixin.__dict__
|
||||
assert "updated_at" in TimestampMixin.__dict__
|
||||
|
||||
|
||||
def test_expected_tables_are_registered() -> None:
|
||||
assert set(Base.metadata.tables) == {
|
||||
"refresh_tokens",
|
||||
"git_repositories",
|
||||
"projects",
|
||||
"ssh_keys",
|
||||
"user_configs",
|
||||
"users",
|
||||
}
|
||||
|
||||
|
||||
def test_user_table_has_required_columns() -> None:
|
||||
columns = User.__table__.columns
|
||||
|
||||
assert set(columns.keys()) == {
|
||||
"id",
|
||||
"email",
|
||||
"name",
|
||||
"authentik_id",
|
||||
"avatar_url",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
}
|
||||
assert columns["email"].unique is True
|
||||
assert columns["authentik_id"].unique is True
|
||||
assert columns["avatar_url"].nullable is True
|
||||
|
||||
|
||||
def test_project_relationships_point_to_owner_and_default_ssh_key() -> None:
|
||||
owner_fk = next(iter(Project.__table__.c.owner_id.foreign_keys))
|
||||
ssh_fk = next(iter(Project.__table__.c.default_ssh_key_id.foreign_keys))
|
||||
|
||||
assert owner_fk.target_fullname == "users.id"
|
||||
assert ssh_fk.target_fullname == "ssh_keys.id"
|
||||
assert Project.owner.property.mapper.class_ is User
|
||||
assert Project.default_ssh_key.property.mapper.class_ is SSHKey
|
||||
|
||||
|
||||
def test_repository_and_user_config_relationships_are_registered() -> None:
|
||||
project_fk = next(iter(GitRepository.__table__.c.project_id.foreign_keys))
|
||||
owner_fk = next(iter(GitRepository.__table__.c.owner_id.foreign_keys))
|
||||
user_config_fk = next(iter(UserConfig.__table__.c.user_id.foreign_keys))
|
||||
|
||||
assert project_fk.target_fullname == "projects.id"
|
||||
assert owner_fk.target_fullname == "users.id"
|
||||
assert user_config_fk.target_fullname == "users.id"
|
||||
assert GitRepository.project.property.mapper.class_ is Project
|
||||
assert GitRepository.owner.property.mapper.class_ is User
|
||||
assert UserConfig.user.property.mapper.class_ is User
|
||||
|
||||
|
||||
def test_refresh_token_table_has_required_columns_and_relationships() -> None:
|
||||
columns = RefreshToken.__table__.columns
|
||||
user_fk = next(iter(RefreshToken.__table__.c.user_id.foreign_keys))
|
||||
|
||||
assert set(columns.keys()) == {
|
||||
"id",
|
||||
"user_id",
|
||||
"token_hash",
|
||||
"expires_at",
|
||||
"revoked_at",
|
||||
"user_agent",
|
||||
"ip_address",
|
||||
"created_at",
|
||||
}
|
||||
assert columns["token_hash"].unique is True
|
||||
assert columns["revoked_at"].nullable is True
|
||||
assert user_fk.target_fullname == "users.id"
|
||||
assert RefreshToken.user.property.mapper.class_ is User
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_session_can_insert_and_load_user(db_session: AsyncSession) -> None:
|
||||
user = User(email="dev@headquarter.local", name="Dev User", authentik_id="dev-user", avatar_url=None)
|
||||
|
||||
db_session.add(user)
|
||||
await db_session.commit()
|
||||
await db_session.refresh(user)
|
||||
|
||||
loaded_user = await db_session.get(User, user.id)
|
||||
|
||||
assert loaded_user is not None
|
||||
assert loaded_user.email == "dev@headquarter.local"
|
||||
@@ -0,0 +1,270 @@
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
import asyncio
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
import pytest
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker
|
||||
|
||||
from src.auth.jwt_service import mint_access_token
|
||||
from src.config import Settings, build_database_url
|
||||
from src.models import Base
|
||||
from src.models.project import Project
|
||||
from src.models.user import User
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def configure_local_database(monkeypatch) -> None:
|
||||
local_url = build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
monkeypatch.setenv("DATABASE_URL", local_url)
|
||||
|
||||
|
||||
def _prepare_test_db() -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(text("TRUNCATE TABLE git_repositories, ssh_keys, projects, users RESTART IDENTITY CASCADE"))
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def _load_app():
|
||||
import importlib
|
||||
import src.database as database_module
|
||||
import src.api.auth as auth_module
|
||||
import src.api.projects as projects_module
|
||||
import src.main as main_module
|
||||
|
||||
# Dispose old engine connections before reload to prevent pool exhaustion
|
||||
if hasattr(database_module, 'engine'):
|
||||
import asyncio
|
||||
asyncio.run(database_module.engine.dispose())
|
||||
|
||||
importlib.reload(database_module)
|
||||
importlib.reload(auth_module)
|
||||
importlib.reload(projects_module)
|
||||
importlib.reload(main_module)
|
||||
return main_module.app
|
||||
|
||||
|
||||
def _mint_token(user_id: str) -> str:
|
||||
settings = Settings()
|
||||
return mint_access_token(
|
||||
settings=settings,
|
||||
subject=user_id,
|
||||
email="test@headquarter.local",
|
||||
name="Test User",
|
||||
expires_at=datetime.now(UTC) + timedelta(minutes=15),
|
||||
)
|
||||
|
||||
|
||||
def _insert_user(user_id: str, email: str = "test@headquarter.local") -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
async with session_factory() as session:
|
||||
user = User(
|
||||
id=uuid.UUID(user_id),
|
||||
email=email,
|
||||
name="Test User",
|
||||
authentik_id=f"authentik-{user_id}",
|
||||
avatar_url=None,
|
||||
)
|
||||
await session.merge(user)
|
||||
await session.commit()
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def _insert_project(project_id: str, owner_id: str, name: str = "Test Project") -> None:
|
||||
async def _run() -> None:
|
||||
engine = create_async_engine(
|
||||
build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
)
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
async with session_factory() as session:
|
||||
project = Project(
|
||||
id=uuid.UUID(project_id),
|
||||
name=name,
|
||||
description="A test project",
|
||||
owner_id=uuid.UUID(owner_id),
|
||||
default_ssh_key_id=None,
|
||||
)
|
||||
await session.merge(project)
|
||||
await session.commit()
|
||||
await engine.dispose()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_create_project_requires_authentication() -> None:
|
||||
_prepare_test_db()
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
|
||||
response = client.post("/projects", json={"name": "New Project", "description": "Description"})
|
||||
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
def test_create_project_successfully() -> None:
|
||||
_prepare_test_db()
|
||||
user_id = "11111111-1111-1111-1111-111111111111"
|
||||
_insert_user(user_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user_id))
|
||||
|
||||
response = client.post("/projects", json={"name": "New Project", "description": "Description"})
|
||||
|
||||
assert response.status_code == 201
|
||||
data = response.json()
|
||||
assert data["name"] == "New Project"
|
||||
assert data["description"] == "Description"
|
||||
assert data["owner_id"] == user_id
|
||||
assert "id" in data
|
||||
|
||||
|
||||
def test_list_projects_returns_only_owned_projects() -> None:
|
||||
_prepare_test_db()
|
||||
user1_id = "11111111-1111-1111-1111-111111111111"
|
||||
user2_id = "22222222-2222-2222-2222-222222222222"
|
||||
_insert_user(user1_id)
|
||||
_insert_user(user2_id, "other@headquarter.local")
|
||||
_insert_project("aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa", user1_id, "User1 Project")
|
||||
_insert_project("bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb", user2_id, "User2 Project")
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(user1_id))
|
||||
|
||||
response = client.get("/projects")
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert len(data) == 1
|
||||
assert data[0]["name"] == "User1 Project"
|
||||
|
||||
|
||||
def test_update_project_requires_ownership() -> None:
|
||||
_prepare_test_db()
|
||||
owner_id = "11111111-1111-1111-1111-111111111111"
|
||||
other_id = "22222222-2222-2222-2222-222222222222"
|
||||
project_id = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
_insert_user(owner_id)
|
||||
_insert_user(other_id, "other@headquarter.local")
|
||||
_insert_project(project_id, owner_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(other_id))
|
||||
|
||||
response = client.patch(f"/projects/{project_id}", json={"name": "Hacked"})
|
||||
|
||||
assert response.status_code == 403
|
||||
|
||||
|
||||
def test_update_project_successfully() -> None:
|
||||
_prepare_test_db()
|
||||
owner_id = "11111111-1111-1111-1111-111111111111"
|
||||
project_id = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
_insert_user(owner_id)
|
||||
_insert_project(project_id, owner_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(owner_id))
|
||||
|
||||
response = client.patch(f"/projects/{project_id}", json={"name": "Updated Name"})
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["name"] == "Updated Name"
|
||||
|
||||
|
||||
def test_delete_project_requires_ownership() -> None:
|
||||
_prepare_test_db()
|
||||
owner_id = "11111111-1111-1111-1111-111111111111"
|
||||
other_id = "22222222-2222-2222-2222-222222222222"
|
||||
project_id = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
_insert_user(owner_id)
|
||||
_insert_user(other_id, "other@headquarter.local")
|
||||
_insert_project(project_id, owner_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(other_id))
|
||||
|
||||
response = client.delete(f"/projects/{project_id}")
|
||||
|
||||
assert response.status_code == 403
|
||||
|
||||
|
||||
def test_delete_project_successfully() -> None:
|
||||
_prepare_test_db()
|
||||
owner_id = "11111111-1111-1111-1111-111111111111"
|
||||
project_id = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
_insert_user(owner_id)
|
||||
_insert_project(project_id, owner_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(owner_id))
|
||||
|
||||
response = client.delete(f"/projects/{project_id}")
|
||||
|
||||
assert response.status_code == 204
|
||||
|
||||
|
||||
def test_set_default_ssh_key_requires_ownership() -> None:
|
||||
_prepare_test_db()
|
||||
owner_id = "11111111-1111-1111-1111-111111111111"
|
||||
other_id = "22222222-2222-2222-2222-222222222222"
|
||||
project_id = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
_insert_user(owner_id)
|
||||
_insert_user(other_id, "other@headquarter.local")
|
||||
_insert_project(project_id, owner_id)
|
||||
|
||||
app = _load_app()
|
||||
client = TestClient(app)
|
||||
client.cookies.set("access_token", _mint_token(other_id))
|
||||
|
||||
response = client.patch(f"/projects/{project_id}/default-ssh-key", json={"ssh_key_id": "cccccccc-cccc-cccc-cccc-cccccccccccc"})
|
||||
|
||||
assert response.status_code == 403
|
||||
@@ -0,0 +1,55 @@
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from sqlalchemy import select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
|
||||
from src.config import build_database_url
|
||||
from src.models.user import User
|
||||
from src.scripts.seed import build_seed_user, seed_database
|
||||
|
||||
|
||||
TEST_DATABASE_URL = build_database_url(
|
||||
user="headquarter",
|
||||
password="headquarter",
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="headquarter",
|
||||
)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def db_session() -> AsyncIterator[AsyncSession]:
|
||||
engine = create_async_engine(TEST_DATABASE_URL)
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
|
||||
async with session_factory() as session:
|
||||
await session.execute(text("TRUNCATE TABLE git_repositories, projects, users RESTART IDENTITY CASCADE"))
|
||||
await session.commit()
|
||||
yield session
|
||||
await session.execute(text("TRUNCATE TABLE git_repositories, projects, users RESTART IDENTITY CASCADE"))
|
||||
await session.commit()
|
||||
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def test_build_seed_user_returns_deterministic_payload() -> None:
|
||||
payload = build_seed_user()
|
||||
|
||||
assert payload == {
|
||||
"email": "dev@headquarter.local",
|
||||
"name": "Development User",
|
||||
"authentik_id": "dev-authentik-user",
|
||||
"avatar_url": None,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_seed_database_creates_development_user(db_session: AsyncSession) -> None:
|
||||
await seed_database(db_session)
|
||||
|
||||
seeded_user = await db_session.scalar(select(User).where(User.email == "dev@headquarter.local"))
|
||||
|
||||
assert seeded_user is not None
|
||||
assert seeded_user.authentik_id == "dev-authentik-user"
|
||||
Reference in New Issue
Block a user