185 lines
6.8 KiB
Python
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,
|
|
)
|
|
)
|