596 lines
19 KiB
Python
596 lines
19 KiB
Python
from dataclasses import dataclass
|
||
from datetime import datetime, timezone
|
||
from typing import Optional, List
|
||
from uuid import UUID, uuid4
|
||
|
||
from cryptography.fernet import InvalidToken
|
||
from sqlalchemy import delete, select
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from app.core.encryption import (
|
||
get_database_encryptor,
|
||
get_encryptor,
|
||
is_database_encryption_configured,
|
||
is_encryption_configured,
|
||
)
|
||
from app.infra.db.metadb import models
|
||
|
||
|
||
def _normalize_postgres_dsn(dsn: str) -> str:
|
||
if not dsn or "://" not in dsn:
|
||
return dsn
|
||
scheme, rest = dsn.split("://", 1)
|
||
if scheme not in ("postgresql", "postgres", "postgresql+psycopg"):
|
||
return dsn
|
||
if "@" not in rest:
|
||
return dsn
|
||
userinfo, hostinfo = rest.rsplit("@", 1)
|
||
if ":" not in userinfo:
|
||
return dsn
|
||
username, password = userinfo.split(":", 1)
|
||
if "@" not in password:
|
||
return dsn
|
||
password = password.replace("@", "%40")
|
||
return f"{scheme}://{username}:{password}@{hostinfo}"
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ProjectDbRouting:
|
||
project_id: UUID
|
||
db_role: str
|
||
db_type: str
|
||
dsn: str
|
||
pool_min_size: int
|
||
pool_max_size: int
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ProjectGeoServerInfo:
|
||
project_id: UUID
|
||
gs_base_url: Optional[str]
|
||
gs_admin_user: Optional[str]
|
||
gs_admin_password: Optional[str]
|
||
gs_datastore_name: str
|
||
default_extent: Optional[dict]
|
||
srid: int
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ProjectSummary:
|
||
project_id: UUID
|
||
name: str
|
||
code: str
|
||
description: Optional[str]
|
||
gs_workspace: str
|
||
map_extent: Optional[dict]
|
||
status: str
|
||
project_role: str
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ProjectDetail:
|
||
project_id: UUID
|
||
name: str
|
||
code: str
|
||
description: Optional[str]
|
||
gs_workspace: str
|
||
map_extent: Optional[dict]
|
||
status: str
|
||
geoserver: Optional[ProjectGeoServerInfo]
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ProjectMemberSummary:
|
||
id: UUID
|
||
user_id: UUID
|
||
project_id: UUID
|
||
project_role: str
|
||
username: str
|
||
email: str
|
||
is_active: bool
|
||
|
||
|
||
def _utcnow() -> datetime:
|
||
return datetime.now(timezone.utc)
|
||
|
||
|
||
def _encrypt_database_secret(value: str) -> str:
|
||
if not is_database_encryption_configured():
|
||
raise ValueError("DATABASE_ENCRYPTION_KEY is not configured")
|
||
return get_database_encryptor().encrypt(value)
|
||
|
||
|
||
def _encrypt_general_secret(value: str) -> str:
|
||
if not is_encryption_configured():
|
||
raise ValueError("ENCRYPTION_KEY is not configured")
|
||
return get_encryptor().encrypt(value)
|
||
|
||
|
||
class MetadataRepository:
|
||
"""元数据访问层(system_hub)"""
|
||
|
||
def __init__(self, session: AsyncSession):
|
||
self.session = session
|
||
|
||
async def get_user_by_keycloak_id(self, keycloak_id: UUID) -> Optional[models.User]:
|
||
result = await self.session.execute(
|
||
select(models.User).where(models.User.keycloak_id == keycloak_id)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def get_user_by_username(self, username: str) -> Optional[models.User]:
|
||
result = await self.session.execute(
|
||
select(models.User).where(models.User.username == username)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def get_user_by_id(self, user_id: UUID) -> Optional[models.User]:
|
||
result = await self.session.execute(
|
||
select(models.User).where(models.User.id == user_id)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def list_users(self, skip: int = 0, limit: int = 100) -> List[models.User]:
|
||
result = await self.session.execute(
|
||
select(models.User)
|
||
.order_by(models.User.created_at.desc())
|
||
.offset(skip)
|
||
.limit(limit)
|
||
)
|
||
return list(result.scalars().all())
|
||
|
||
async def upsert_user_from_keycloak(
|
||
self,
|
||
*,
|
||
keycloak_id: UUID,
|
||
username: str,
|
||
email: str,
|
||
role: str,
|
||
is_active: bool,
|
||
) -> models.User:
|
||
user = await self.get_user_by_keycloak_id(keycloak_id)
|
||
if user is None:
|
||
user = models.User(
|
||
id=uuid4(),
|
||
keycloak_id=keycloak_id,
|
||
username=username,
|
||
email=email,
|
||
role=role,
|
||
is_active=is_active,
|
||
is_superuser=False,
|
||
)
|
||
self.session.add(user)
|
||
else:
|
||
user.username = username
|
||
user.email = email
|
||
user.role = role
|
||
user.is_active = is_active
|
||
await self.session.commit()
|
||
await self.session.refresh(user)
|
||
return user
|
||
|
||
async def refresh_user_keycloak_snapshot(
|
||
self,
|
||
user: models.User,
|
||
*,
|
||
username: str | None,
|
||
email: str | None,
|
||
last_login_at: datetime | None = None,
|
||
) -> models.User:
|
||
if username:
|
||
user.username = username
|
||
if email:
|
||
user.email = email
|
||
user.last_login_at = last_login_at or _utcnow()
|
||
user.updated_at = _utcnow()
|
||
await self.session.commit()
|
||
await self.session.refresh(user)
|
||
return user
|
||
|
||
async def update_user_admin(
|
||
self,
|
||
user_id: UUID,
|
||
*,
|
||
updates: dict,
|
||
) -> Optional[models.User]:
|
||
user = await self.get_user_by_id(user_id)
|
||
if user is None:
|
||
return None
|
||
if "role" in updates:
|
||
user.role = updates["role"]
|
||
if "is_active" in updates:
|
||
user.is_active = updates["is_active"]
|
||
await self.session.commit()
|
||
await self.session.refresh(user)
|
||
return user
|
||
|
||
async def get_project_by_id(self, project_id: UUID) -> Optional[models.Project]:
|
||
result = await self.session.execute(
|
||
select(models.Project).where(models.Project.id == project_id)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def get_project_by_code(self, code: str) -> Optional[models.Project]:
|
||
result = await self.session.execute(
|
||
select(models.Project).where(models.Project.code == code)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def list_project_records(self) -> List[models.Project]:
|
||
result = await self.session.execute(
|
||
select(models.Project).order_by(models.Project.name)
|
||
)
|
||
return list(result.scalars().all())
|
||
|
||
async def create_project(
|
||
self,
|
||
*,
|
||
name: str,
|
||
code: str,
|
||
description: str | None,
|
||
gs_workspace: str,
|
||
map_extent: dict | None,
|
||
status: str,
|
||
) -> models.Project:
|
||
project = models.Project(
|
||
id=uuid4(),
|
||
name=name,
|
||
code=code,
|
||
description=description,
|
||
gs_workspace=gs_workspace,
|
||
map_extent=map_extent,
|
||
status=status,
|
||
created_at=_utcnow(),
|
||
updated_at=_utcnow(),
|
||
)
|
||
self.session.add(project)
|
||
await self.session.commit()
|
||
await self.session.refresh(project)
|
||
return project
|
||
|
||
async def update_project(
|
||
self,
|
||
project_id: UUID,
|
||
*,
|
||
updates: dict,
|
||
) -> Optional[models.Project]:
|
||
project = await self.get_project_by_id(project_id)
|
||
if project is None:
|
||
return None
|
||
for field in (
|
||
"name",
|
||
"code",
|
||
"description",
|
||
"gs_workspace",
|
||
"map_extent",
|
||
"status",
|
||
):
|
||
if field in updates:
|
||
setattr(project, field, updates[field])
|
||
project.updated_at = _utcnow()
|
||
await self.session.commit()
|
||
await self.session.refresh(project)
|
||
return project
|
||
|
||
async def get_project_detail_by_code(self, code: str) -> Optional[ProjectDetail]:
|
||
project = await self.get_project_by_code(code)
|
||
if not project:
|
||
return None
|
||
|
||
geoserver = await self.get_geoserver_config(project.id)
|
||
|
||
return ProjectDetail(
|
||
project_id=project.id,
|
||
name=project.name,
|
||
code=project.code,
|
||
description=project.description,
|
||
gs_workspace=project.gs_workspace,
|
||
map_extent=project.map_extent,
|
||
status=project.status,
|
||
geoserver=geoserver
|
||
)
|
||
|
||
async def get_membership_role(
|
||
self, project_id: UUID, user_id: UUID
|
||
) -> Optional[str]:
|
||
result = await self.session.execute(
|
||
select(models.UserProjectMembership.project_role).where(
|
||
models.UserProjectMembership.project_id == project_id,
|
||
models.UserProjectMembership.user_id == user_id,
|
||
)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def list_project_members(
|
||
self, project_id: UUID
|
||
) -> List[ProjectMemberSummary]:
|
||
stmt = (
|
||
select(models.UserProjectMembership, models.User)
|
||
.join(models.User, models.User.id == models.UserProjectMembership.user_id)
|
||
.where(models.UserProjectMembership.project_id == project_id)
|
||
.order_by(models.User.username)
|
||
)
|
||
result = await self.session.execute(stmt)
|
||
return [
|
||
ProjectMemberSummary(
|
||
id=membership.id,
|
||
user_id=membership.user_id,
|
||
project_id=membership.project_id,
|
||
project_role=membership.project_role,
|
||
username=user.username,
|
||
email=user.email,
|
||
is_active=user.is_active,
|
||
)
|
||
for membership, user in result.all()
|
||
]
|
||
|
||
async def get_project_membership(
|
||
self, project_id: UUID, user_id: UUID
|
||
) -> Optional[models.UserProjectMembership]:
|
||
result = await self.session.execute(
|
||
select(models.UserProjectMembership).where(
|
||
models.UserProjectMembership.project_id == project_id,
|
||
models.UserProjectMembership.user_id == user_id,
|
||
)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def add_project_member(
|
||
self, project_id: UUID, user_id: UUID, project_role: str
|
||
) -> models.UserProjectMembership:
|
||
membership = models.UserProjectMembership(
|
||
id=uuid4(),
|
||
user_id=user_id,
|
||
project_id=project_id,
|
||
project_role=project_role,
|
||
)
|
||
self.session.add(membership)
|
||
await self.session.commit()
|
||
await self.session.refresh(membership)
|
||
return membership
|
||
|
||
async def update_project_member_role(
|
||
self, project_id: UUID, user_id: UUID, project_role: str
|
||
) -> Optional[models.UserProjectMembership]:
|
||
membership = await self.get_project_membership(project_id, user_id)
|
||
if membership is None:
|
||
return None
|
||
membership.project_role = project_role
|
||
await self.session.commit()
|
||
await self.session.refresh(membership)
|
||
return membership
|
||
|
||
async def remove_project_member(self, project_id: UUID, user_id: UUID) -> bool:
|
||
result = await self.session.execute(
|
||
delete(models.UserProjectMembership).where(
|
||
models.UserProjectMembership.project_id == project_id,
|
||
models.UserProjectMembership.user_id == user_id,
|
||
)
|
||
)
|
||
await self.session.commit()
|
||
return bool(result.rowcount)
|
||
|
||
async def list_project_databases(
|
||
self, project_id: UUID
|
||
) -> List[models.ProjectDatabase]:
|
||
result = await self.session.execute(
|
||
select(models.ProjectDatabase)
|
||
.where(models.ProjectDatabase.project_id == project_id)
|
||
.order_by(models.ProjectDatabase.db_role)
|
||
)
|
||
return list(result.scalars().all())
|
||
|
||
async def get_project_database_config(
|
||
self, project_id: UUID, db_role: str
|
||
) -> Optional[models.ProjectDatabase]:
|
||
result = await self.session.execute(
|
||
select(models.ProjectDatabase).where(
|
||
models.ProjectDatabase.project_id == project_id,
|
||
models.ProjectDatabase.db_role == db_role,
|
||
)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def upsert_project_database_config(
|
||
self,
|
||
project_id: UUID,
|
||
*,
|
||
db_role: str,
|
||
db_type: str,
|
||
dsn: str | None,
|
||
pool_min_size: int,
|
||
pool_max_size: int,
|
||
) -> models.ProjectDatabase:
|
||
record = await self.get_project_database_config(project_id, db_role)
|
||
if record is None:
|
||
if dsn is None:
|
||
raise ValueError("dsn is required when creating project database config")
|
||
record = models.ProjectDatabase(
|
||
id=uuid4(),
|
||
project_id=project_id,
|
||
db_role=db_role,
|
||
db_type=db_type,
|
||
dsn_encrypted=_encrypt_database_secret(dsn),
|
||
pool_min_size=pool_min_size,
|
||
pool_max_size=pool_max_size,
|
||
)
|
||
self.session.add(record)
|
||
else:
|
||
record.db_type = db_type
|
||
if dsn is not None:
|
||
record.dsn_encrypted = _encrypt_database_secret(dsn)
|
||
record.pool_min_size = pool_min_size
|
||
record.pool_max_size = pool_max_size
|
||
await self.session.commit()
|
||
await self.session.refresh(record)
|
||
return record
|
||
|
||
async def delete_project_database_config(
|
||
self, project_id: UUID, db_role: str
|
||
) -> bool:
|
||
result = await self.session.execute(
|
||
delete(models.ProjectDatabase).where(
|
||
models.ProjectDatabase.project_id == project_id,
|
||
models.ProjectDatabase.db_role == db_role,
|
||
)
|
||
)
|
||
await self.session.commit()
|
||
return bool(result.rowcount)
|
||
|
||
async def get_project_db_routing(
|
||
self, project_id: UUID, db_role: str
|
||
) -> Optional[ProjectDbRouting]:
|
||
result = await self.session.execute(
|
||
select(models.ProjectDatabase).where(
|
||
models.ProjectDatabase.project_id == project_id,
|
||
models.ProjectDatabase.db_role == db_role,
|
||
)
|
||
)
|
||
record = result.scalar_one_or_none()
|
||
if not record:
|
||
return None
|
||
if not is_database_encryption_configured():
|
||
raise ValueError("DATABASE_ENCRYPTION_KEY is not configured")
|
||
encryptor = get_database_encryptor()
|
||
try:
|
||
dsn = encryptor.decrypt(record.dsn_encrypted)
|
||
except InvalidToken:
|
||
raise ValueError(
|
||
"Failed to decrypt project DB DSN: DATABASE_ENCRYPTION_KEY mismatch "
|
||
"or invalid dsn_encrypted value"
|
||
)
|
||
dsn = _normalize_postgres_dsn(dsn)
|
||
return ProjectDbRouting(
|
||
project_id=record.project_id,
|
||
db_role=record.db_role,
|
||
db_type=record.db_type,
|
||
dsn=dsn,
|
||
pool_min_size=record.pool_min_size,
|
||
pool_max_size=record.pool_max_size,
|
||
)
|
||
|
||
async def get_geoserver_config(
|
||
self, project_id: UUID
|
||
) -> Optional[ProjectGeoServerInfo]:
|
||
result = await self.session.execute(
|
||
select(models.ProjectGeoServerConfig).where(
|
||
models.ProjectGeoServerConfig.project_id == project_id
|
||
)
|
||
)
|
||
record = result.scalar_one_or_none()
|
||
if not record:
|
||
return None
|
||
if record.gs_admin_password_encrypted:
|
||
if is_encryption_configured():
|
||
encryptor = get_encryptor()
|
||
password = encryptor.decrypt(record.gs_admin_password_encrypted)
|
||
else:
|
||
password = record.gs_admin_password_encrypted
|
||
else:
|
||
password = None
|
||
return ProjectGeoServerInfo(
|
||
project_id=record.project_id,
|
||
gs_base_url=record.gs_base_url,
|
||
gs_admin_user=record.gs_admin_user,
|
||
gs_admin_password=password,
|
||
gs_datastore_name=record.gs_datastore_name,
|
||
default_extent=record.default_extent,
|
||
srid=record.srid,
|
||
)
|
||
|
||
async def get_geoserver_config_record(
|
||
self, project_id: UUID
|
||
) -> Optional[models.ProjectGeoServerConfig]:
|
||
result = await self.session.execute(
|
||
select(models.ProjectGeoServerConfig).where(
|
||
models.ProjectGeoServerConfig.project_id == project_id
|
||
)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def upsert_geoserver_config(
|
||
self,
|
||
project_id: UUID,
|
||
*,
|
||
gs_base_url: str | None,
|
||
gs_admin_user: str | None,
|
||
gs_admin_password: str | None,
|
||
password_update_requested: bool,
|
||
gs_datastore_name: str,
|
||
default_extent: dict | None,
|
||
srid: int,
|
||
) -> models.ProjectGeoServerConfig:
|
||
record = await self.get_geoserver_config_record(project_id)
|
||
encrypted_password: str | None = None
|
||
if password_update_requested and gs_admin_password is not None:
|
||
encrypted_password = _encrypt_general_secret(gs_admin_password)
|
||
|
||
if record is None:
|
||
record = models.ProjectGeoServerConfig(
|
||
id=uuid4(),
|
||
project_id=project_id,
|
||
gs_base_url=gs_base_url,
|
||
gs_admin_user=gs_admin_user,
|
||
gs_admin_password_encrypted=encrypted_password,
|
||
gs_datastore_name=gs_datastore_name,
|
||
default_extent=default_extent,
|
||
srid=srid,
|
||
updated_at=_utcnow(),
|
||
)
|
||
self.session.add(record)
|
||
else:
|
||
record.gs_base_url = gs_base_url
|
||
record.gs_admin_user = gs_admin_user
|
||
if password_update_requested:
|
||
record.gs_admin_password_encrypted = encrypted_password
|
||
record.gs_datastore_name = gs_datastore_name
|
||
record.default_extent = default_extent
|
||
record.srid = srid
|
||
record.updated_at = _utcnow()
|
||
await self.session.commit()
|
||
await self.session.refresh(record)
|
||
return record
|
||
|
||
async def list_projects_for_user(self, user_id: UUID) -> List[ProjectSummary]:
|
||
stmt = (
|
||
select(models.Project, models.UserProjectMembership.project_role)
|
||
.join(
|
||
models.UserProjectMembership,
|
||
models.UserProjectMembership.project_id == models.Project.id,
|
||
)
|
||
.where(models.UserProjectMembership.user_id == user_id)
|
||
.order_by(models.Project.name)
|
||
)
|
||
result = await self.session.execute(stmt)
|
||
return [
|
||
ProjectSummary(
|
||
project_id=project.id,
|
||
name=project.name,
|
||
code=project.code,
|
||
description=project.description,
|
||
gs_workspace=project.gs_workspace,
|
||
map_extent=project.map_extent,
|
||
status=project.status,
|
||
project_role=role,
|
||
)
|
||
for project, role in result.all()
|
||
]
|
||
|
||
async def list_all_projects(self) -> List[ProjectSummary]:
|
||
result = await self.session.execute(
|
||
select(models.Project).order_by(models.Project.name)
|
||
)
|
||
return [
|
||
ProjectSummary(
|
||
project_id=project.id,
|
||
name=project.name,
|
||
code=project.code,
|
||
description=project.description,
|
||
gs_workspace=project.gs_workspace,
|
||
map_extent=project.map_extent,
|
||
status=project.status,
|
||
project_role="owner",
|
||
)
|
||
for project in result.scalars().all()
|
||
]
|