"""登录客户(CUSTOMER)的 NL2SQL 权限快照服务。 与员工路径(permission_service.load_query_permission,按 nl2sql_query_role 配置)不同,客户权限不落库、不做管理后台:表白名单通过 sys_config 配置管理, 且只能从内置白名单中做"减法",行级范围由服务端强制注入 customer_ids, 保证客户永远只能查询自己的数据。 列级校验:快照时从 information_schema 加载白名单表的真实列清单写入 columns,validate_select_sql 据此在执行前拦截 LLM 幻觉列(避免把 Unknown column 错误漏到执行期)。 """ from __future__ import annotations import logging from inspect import isawaitable from sqlalchemy import bindparam, text logger = logging.getLogger(__name__) # 客户可查询的内置表白名单(配置只能在其中做减法,不能新增表) CUSTOMER_DEFAULT_TABLES: tuple[str, ...] = ( "fin_holdings", "fin_transaction", "fin_product", "fund_nav_history", "fund_performance", ) # 行级隔离:出现这些表的 SQL 会被强制注入 customer_id IN (<登录用户>) 条件 CUSTOMER_ROW_SCOPES: dict[str, dict[str, str]] = { "fin_holdings": {"type": "customer_ids", "column": "customer_id"}, "fin_transaction": {"type": "customer_ids", "column": "customer_id"}, } # 客户路径首期不开放敏感档案表;后续开放时在此配置 (table, column) -> mask_type CUSTOMER_MASKS: dict[tuple[str, str], str] = {} _TRUTHY = {"1", "true", "yes", "on"} _COLUMNS_SQL = text( "SELECT TABLE_NAME AS table_name, COLUMN_NAME AS column_name " "FROM information_schema.columns " "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME IN :tables " "ORDER BY TABLE_NAME, ORDINAL_POSITION" ).bindparams(bindparam("tables", expanding=True)) async def _load_real_columns(db, tables: set[str]) -> dict[str, set[str]] | None: """加载白名单表的真实列名;db 为 None 时返回 None(仅测试路径)。""" if db is None: return None result = await db.execute(_COLUMNS_SQL, {"tables": sorted(tables)}) columns: dict[str, set[str]] = {} for row in result.mappings(): table_name = str(row["table_name"] or "").strip() column_name = str(row["column_name"] or "").strip() if table_name and column_name: columns.setdefault(table_name, set()).add(column_name) return columns def _denied_permission() -> dict: return { "can_query": False, "role": "customer_self", "tables": set(), "columns": None, "masks": {}, "row_scopes": {}, "max_rows": 0, "daily_quota": 0, } async def _config(config_getter, key: str, default: str) -> str: value = config_getter(key, default) if isawaitable(value): value = await value if value is None or str(value).strip() == "": return default return str(value) def _parse_allowed_tables(raw: str) -> set[str]: """解析表白名单配置;非法表名直接忽略,只允许内置白名单的子集。""" known = set(CUSTOMER_DEFAULT_TABLES) names = { item.strip().lower() for item in str(raw).replace(";", ",").replace(";", ",").split(",") if item.strip() } tables = names & known return tables async def load_customer_query_permission( db, user_id: int, *, config_getter, ) -> dict: """每次请求重建客户权限快照。 客户身份已由 API 层(require_customer)和会话归属校验保证, 快照不依赖数据库中的角色配置;db 用于加载白名单表的真实列清单 (传入 None 时跳过列清单,columns 保持 None,仅限测试路径)。 """ del user_id # 权限与具体请求上下文无关,签名对齐 execute_query 的 permission_loader if config_getter is None: return _denied_permission() enabled = (await _config(config_getter, "nl2sql.customer.enabled", "true")).lower() if enabled not in _TRUTHY: return _denied_permission() raw_tables = await _config( config_getter, "nl2sql.customer.allowed_tables", ",".join(CUSTOMER_DEFAULT_TABLES), ) tables = _parse_allowed_tables(raw_tables) if not tables: return _denied_permission() try: max_rows = int(await _config(config_getter, "nl2sql.customer.max_rows", "200")) daily_quota = int( await _config(config_getter, "nl2sql.customer.daily_quota", "20") ) except ValueError: max_rows, daily_quota = 200, 20 max_rows = max(1, max_rows) daily_quota = max(0, daily_quota) # 列级校验用真实列清单:拦截 LLM 幻觉列,避免执行期 Unknown column。 # 信息读取失败时按"拒绝"处理(fail-closed),不让无列校验的快照放行。 try: real_columns = await _load_real_columns(db, tables) except Exception: logger.exception("load customer nl2sql columns failed") return _denied_permission() return { "can_query": True, "role": "customer_self", "tables": tables, # 真实列清单(db=None 的测试路径保持 None = 不限列) "columns": real_columns, "masks": dict(CUSTOMER_MASKS), "row_scopes": { table: dict(scope) for table, scope in CUSTOMER_ROW_SCOPES.items() if table in tables }, "max_rows": max_rows, "daily_quota": daily_quota, }