"""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, ) )