Files

149 lines
4.6 KiB
Python

"""NL2SQL 查询缓存适配器,Redis 不可用时自动降级。"""
from __future__ import annotations
import hashlib
import json
import logging
from typing import Any
from sqlglot import parse_one
from nl2sql.contracts import DataQueryResult
logger = logging.getLogger("nl2sql.cache")
DEFAULT_CACHE_TTL = 300
def normalize_sql(sql: str) -> str:
"""统一 SQL 的关键字、空白和末尾分号。"""
return parse_one(sql, read="mysql").sql(dialect="mysql")
def _permission_payload(permission: dict[str, Any]) -> dict[str, Any]:
"""将权限集合转换成稳定、可哈希的结构。"""
return {
"tables": sorted(permission.get("tables", set())),
"columns": {
table: sorted(columns)
for table, columns in sorted(permission.get("columns", {}).items())
},
}
def build_cache_key(
sql: str,
*,
permission: dict[str, Any],
data_version: str = "",
) -> str:
"""生成包含标准化 SQL、权限范围和数据版本的缓存键。"""
payload = {
"sql": normalize_sql(sql),
"permission": _permission_payload(permission),
"data_version": data_version,
}
digest = hashlib.sha256(
json.dumps(payload, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode()
).hexdigest()
return f"nl2sql:{digest}"
def build_question_cache_key(
question: str,
*,
permission: dict[str, Any],
data_scope: dict[str, Any] | None = None,
page: int = 1,
page_size: int = 100,
sort_by: str | None = None,
sort_order: str = "asc",
) -> str:
"""生成自然语言查询结果缓存键,覆盖问题、权限和数据范围。"""
payload = {
"question": " ".join((question or "").split()),
"permission": _permission_payload(permission),
"data_scope": data_scope or {},
"page": page,
"page_size": page_size,
"sort_by": sort_by or "",
"sort_order": sort_order,
}
digest = hashlib.sha256(
json.dumps(payload, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode()
).hexdigest()
return f"nl2sql:question:{digest}"
async def cache_get(redis, key: str) -> dict[str, Any] | None:
"""读取缓存,Redis 异常或内容损坏时返回空结果。"""
try:
value = await redis.get(key)
if not value:
return None
return json.loads(value)
except Exception: # noqa: BLE001 缓存故障必须降级直连数据库
logger.warning("NL2SQL 缓存读取失败", exc_info=True)
return None
async def cache_set(
redis,
key: str,
value: dict[str, Any],
*,
ttl: int,
access_tables: set[str] | None = None,
) -> bool:
"""写入缓存并建立访问表索引,Redis 异常时返回 False。"""
try:
await redis.set(key, json.dumps(value, ensure_ascii=False), ex=ttl)
for table_name in access_tables or set():
await redis.sadd(f"nl2sql:table-keys:{table_name}", key)
except Exception: # noqa: BLE001 缓存故障不能阻断直连查询
logger.warning("NL2SQL 缓存写入失败", exc_info=True)
return False
return True
async def invalidate_tables(redis, table_names: set[str]) -> int:
"""删除指定表关联的全部缓存 key。"""
deleted = 0
try:
for table_name in table_names:
index_key = f"nl2sql:table-keys:{table_name}"
keys = await redis.smembers(index_key)
if keys:
await redis.delete(*keys)
deleted += len(keys)
await redis.delete(index_key)
except Exception: # noqa: BLE001 缓存失效失败只记录并返回已处理数量
logger.warning("NL2SQL 缓存失效失败", exc_info=True)
return deleted
def result_to_cache(result: DataQueryResult) -> dict[str, Any]:
"""将统一查询结果转换为可安全序列化的缓存载荷。"""
return {
"query_id": result.query_id,
"trace_id": result.trace_id,
"columns": result.columns,
"rows": result.rows,
"row_count": result.row_count,
"truncated": result.truncated,
"summary": result.summary,
"answer": result.answer,
"markdown": result.markdown,
"chart": result.chart,
"metric_definitions": result.metric_definitions,
"query_plan": result.query_plan,
"sql": result.sql,
"elapsed_ms": result.elapsed_ms,
"warnings": result.warnings,
}
def result_from_cache(payload: dict[str, Any]) -> DataQueryResult:
"""将缓存载荷恢复为统一查询结果。"""
return DataQueryResult(**payload)