Files

185 lines
6.8 KiB
Python

"""NL2SQL 查询权限仓储。"""
from __future__ import annotations
from datetime import datetime
from sqlalchemy import delete, select
from model.nl2sql_permission import (
Nl2SqlQueryRole,
Nl2SqlRoleColumnPermission,
Nl2SqlRoleTablePermission,
Nl2SqlSensitiveField,
Nl2SqlQueryHistory,
)
from repositories.base import BaseRepository
class Nl2SqlPermissionRepo(BaseRepository):
async def list_roles(self, *, include_inactive: bool = False):
stmt = select(Nl2SqlQueryRole).order_by(Nl2SqlQueryRole.id)
if not include_inactive:
stmt = stmt.where(Nl2SqlQueryRole.status == "active")
return list((await self.db.scalars(stmt)).all())
async def get_role(self, role_id: int):
return await self.db.get(Nl2SqlQueryRole, role_id)
async def add_role(self, **kwargs):
return await self.add(Nl2SqlQueryRole(**kwargs))
async def update_role(self, role_id: int, **kwargs):
obj = await self.get_role(role_id)
if obj is None:
return None
for key, value in kwargs.items():
if value is not None:
setattr(obj, key, value)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def list_role_table_permissions(self, role_id: int, *, include_inactive: bool = False):
stmt = select(Nl2SqlRoleTablePermission).where(
Nl2SqlRoleTablePermission.role_id == role_id
)
if not include_inactive:
stmt = stmt.where(Nl2SqlRoleTablePermission.status == "active")
return list((await self.db.scalars(stmt.order_by(Nl2SqlRoleTablePermission.id))).all())
async def get_table_permission(self, permission_id: int):
return await self.db.get(Nl2SqlRoleTablePermission, permission_id)
async def add_table_permission(self, **kwargs):
return await self.add(Nl2SqlRoleTablePermission(**kwargs))
async def update_table_permission(self, permission_id: int, **kwargs):
obj = await self.get_table_permission(permission_id)
if obj is None:
return None
for key, value in kwargs.items():
if value is not None:
setattr(obj, key, value)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def list_role_column_permissions(self, role_id: int, *, include_inactive: bool = False):
stmt = select(Nl2SqlRoleColumnPermission).where(
Nl2SqlRoleColumnPermission.role_id == role_id
)
if not include_inactive:
stmt = stmt.where(Nl2SqlRoleColumnPermission.status == "active")
return list((await self.db.scalars(stmt.order_by(Nl2SqlRoleColumnPermission.id))).all())
async def get_column_permission(self, permission_id: int):
return await self.db.get(Nl2SqlRoleColumnPermission, permission_id)
async def add_column_permission(self, **kwargs):
return await self.add(Nl2SqlRoleColumnPermission(**kwargs))
async def update_column_permission(self, permission_id: int, **kwargs):
obj = await self.get_column_permission(permission_id)
if obj is None:
return None
for key, value in kwargs.items():
if value is not None:
setattr(obj, key, value)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def list_sensitive_fields_admin(self, *, include_inactive: bool = False):
stmt = select(Nl2SqlSensitiveField).order_by(Nl2SqlSensitiveField.id)
if not include_inactive:
stmt = stmt.where(Nl2SqlSensitiveField.status == "active")
return list((await self.db.scalars(stmt)).all())
async def get_sensitive_field(self, field_id: int):
return await self.db.get(Nl2SqlSensitiveField, field_id)
async def add_sensitive_field(self, **kwargs):
return await self.add(Nl2SqlSensitiveField(**kwargs))
async def update_sensitive_field(self, field_id: int, **kwargs):
obj = await self.get_sensitive_field(field_id)
if obj is None:
return None
for key, value in kwargs.items():
if value is not None:
setattr(obj, key, value)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def delete_history_before(self, before: datetime) -> int:
result = await self.db.execute(
delete(Nl2SqlQueryHistory).where(Nl2SqlQueryHistory.create_time < before)
)
await self.db.commit()
return int(result.rowcount or 0)
async def get_role_by_employee_role(self, employee_role: str):
return await self.db.scalar(
select(Nl2SqlQueryRole).where(
Nl2SqlQueryRole.employee_role == employee_role,
Nl2SqlQueryRole.status == "active",
)
)
async def list_table_permissions(self, role_id: int):
result = await self.db.scalars(
select(Nl2SqlRoleTablePermission).where(
Nl2SqlRoleTablePermission.role_id == role_id,
Nl2SqlRoleTablePermission.status == "active",
)
)
return list(result.all())
async def list_column_permissions(self, role_id: int):
result = await self.db.scalars(
select(Nl2SqlRoleColumnPermission).where(
Nl2SqlRoleColumnPermission.role_id == role_id,
Nl2SqlRoleColumnPermission.status == "active",
)
)
return list(result.all())
async def list_sensitive_fields(self):
result = await self.db.scalars(
select(Nl2SqlSensitiveField).where(Nl2SqlSensitiveField.status == "active")
)
return list(result.all())
async def list_query_history(
self,
user_id: int,
*,
limit: int = 20,
offset: int = 0,
status: str | None = None,
start_time: datetime | None = None,
end_time: datetime | None = None,
):
"""按员工自身范围分页读取查询历史。"""
stmt = select(Nl2SqlQueryHistory).where(Nl2SqlQueryHistory.user_id == user_id)
if status:
stmt = stmt.where(Nl2SqlQueryHistory.status == status)
if start_time:
stmt = stmt.where(Nl2SqlQueryHistory.create_time >= start_time)
if end_time:
stmt = stmt.where(Nl2SqlQueryHistory.create_time <= end_time)
result = await self.db.scalars(
stmt.order_by(Nl2SqlQueryHistory.create_time.desc()).limit(limit).offset(offset)
)
return list(result.all())
async def get_query_history(self, user_id: int, query_id: str):
"""按员工自身范围读取单条查询历史。"""
return await self.db.scalar(
select(Nl2SqlQueryHistory).where(
Nl2SqlQueryHistory.user_id == user_id,
Nl2SqlQueryHistory.query_id == query_id,
)
)