947 lines
32 KiB
Python
947 lines
32 KiB
Python
import asyncio
|
||
import json
|
||
import re
|
||
from collections.abc import AsyncIterator
|
||
from copy import deepcopy
|
||
from datetime import datetime, timezone
|
||
from typing import Any
|
||
|
||
import httpx
|
||
from fastapi import HTTPException, UploadFile
|
||
|
||
from config.cubeReport import get_cube_report_settings
|
||
from db.sqlserver import execute_parameterized_scalar, execute_parameterized_sql
|
||
from service.cube_metadata import cube_member_candidates, get_cube_member_titles
|
||
|
||
|
||
def _require_dify_key() -> str:
|
||
key = get_cube_report_settings().dify_api_key
|
||
if not key:
|
||
raise HTTPException(
|
||
status_code=503,
|
||
detail="尚未配置 DIFY_DATABASE_ASSISTANT_API_KEY",
|
||
)
|
||
return key
|
||
|
||
|
||
def _headers(token: str, content_type: str | None = "application/json") -> dict[str, str]:
|
||
headers: dict[str, str] = {}
|
||
if token:
|
||
headers["Authorization"] = f"Bearer {token}"
|
||
if content_type:
|
||
headers["Content-Type"] = content_type
|
||
return headers
|
||
|
||
|
||
def _timeout() -> httpx.Timeout:
|
||
seconds = get_cube_report_settings().request_timeout_seconds
|
||
return httpx.Timeout(seconds, connect=15.0)
|
||
|
||
|
||
def _raise_dify_error(response: httpx.Response) -> None:
|
||
if not response.is_error:
|
||
return
|
||
try:
|
||
payload = response.json()
|
||
detail = payload.get("message") or payload.get("detail") or response.text
|
||
except (ValueError, TypeError):
|
||
detail = response.text
|
||
raise HTTPException(
|
||
status_code=response.status_code,
|
||
detail=f"Dify 请求失败:{str(detail)[:500]}",
|
||
)
|
||
|
||
|
||
def _iso_timestamp(value: Any) -> str | None:
|
||
try:
|
||
return datetime.fromtimestamp(float(value), tz=timezone.utc).isoformat()
|
||
except (TypeError, ValueError, OSError):
|
||
return None
|
||
|
||
|
||
def _as_bool(value: Any) -> bool:
|
||
if isinstance(value, bool):
|
||
return value
|
||
return str(value).strip().lower() in {"1", "true", "yes", "on"}
|
||
|
||
|
||
def _decode_json_objects(text: str) -> list[Any]:
|
||
"""从混合文本中提取完整 JSON,兼容工作流误输出连续多个对象。"""
|
||
decoder = json.JSONDecoder()
|
||
values: list[Any] = []
|
||
index = 0
|
||
while index < len(text):
|
||
object_start = text.find("{", index)
|
||
array_start = text.find("[", index)
|
||
starts = [item for item in (object_start, array_start) if item >= 0]
|
||
if not starts:
|
||
break
|
||
start = min(starts)
|
||
try:
|
||
value, end = decoder.raw_decode(text, start)
|
||
except json.JSONDecodeError:
|
||
index = start + 1
|
||
continue
|
||
values.append(value)
|
||
index = end
|
||
return values
|
||
|
||
|
||
def normalize_cube_response(value: Any) -> dict[str, Any] | None:
|
||
"""解析工作流输出;多个有效 JSON 连续出现时,以最后一个为准。"""
|
||
current = value
|
||
if isinstance(current, str):
|
||
text = current.strip()
|
||
if text.startswith("```"):
|
||
text = re.sub(r"^```(?:json)?\s*", "", text, flags=re.IGNORECASE)
|
||
text = re.sub(r"\s*```$", "", text)
|
||
try:
|
||
current = json.loads(text)
|
||
except json.JSONDecodeError:
|
||
candidates = [
|
||
item
|
||
for item in _decode_json_objects(text)
|
||
if isinstance(item, dict) and isinstance(item.get("query"), dict)
|
||
]
|
||
if not candidates:
|
||
return None
|
||
current = candidates[-1]
|
||
if not isinstance(current, dict) or not isinstance(current.get("query"), dict):
|
||
return None
|
||
return {
|
||
"query": current["query"],
|
||
"message": str(current.get("message") or "数据请求已生成"),
|
||
"status": _as_bool(current.get("status", True)),
|
||
"limitSource": str(current.get("limit_source") or "system"),
|
||
}
|
||
|
||
|
||
async def get_dify_parameters() -> dict[str, Any]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.get(
|
||
f"{settings.dify_api_base}/parameters",
|
||
headers=_headers(key),
|
||
)
|
||
_raise_dify_error(response)
|
||
payload = response.json()
|
||
questions = payload.get("suggested_questions", [])
|
||
questions = questions if isinstance(questions, list) else []
|
||
return {
|
||
"openingStatement": str(
|
||
payload.get("opening_statement") or payload.get("introduction") or ""
|
||
),
|
||
"suggestedQuestions": [str(item) for item in questions if item],
|
||
"suggestedQuestionsAfterAnswer": payload.get(
|
||
"suggested_questions_after_answer", {"enabled": False}
|
||
),
|
||
"fileUpload": payload.get("file_upload", {}),
|
||
"systemParameters": payload.get("system_parameters", {}),
|
||
"userInputForm": payload.get("user_input_form", []),
|
||
}
|
||
|
||
|
||
def normalize_conversation(item: dict[str, Any]) -> dict[str, Any]:
|
||
inputs = item.get("inputs") if isinstance(item.get("inputs"), dict) else {}
|
||
return {
|
||
"id": str(item.get("id") or ""),
|
||
"title": str(item.get("name") or "新数据对话"),
|
||
"tenantId": str(inputs.get("tenant_id") or ""),
|
||
"tenantName": str(inputs.get("tenant_name") or "未指定租户"),
|
||
"createdAt": _iso_timestamp(item.get("created_at")),
|
||
"updatedAt": _iso_timestamp(item.get("updated_at")),
|
||
"inputs": inputs,
|
||
}
|
||
|
||
|
||
async def list_dify_conversations(
|
||
user_id: str, *, max_items: int = 500
|
||
) -> list[dict[str, Any]]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
conversations: list[dict[str, Any]] = []
|
||
last_id: str | None = None
|
||
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
while len(conversations) < max_items:
|
||
params: dict[str, Any] = {
|
||
"user": user_id,
|
||
"limit": min(100, max_items - len(conversations)),
|
||
"sort_by": "-updated_at",
|
||
}
|
||
if last_id:
|
||
params["last_id"] = last_id
|
||
response = await client.get(
|
||
f"{settings.dify_api_base}/conversations",
|
||
headers=_headers(key),
|
||
params=params,
|
||
)
|
||
_raise_dify_error(response)
|
||
payload = response.json()
|
||
batch = payload.get("data", [])
|
||
if not isinstance(batch, list) or not batch:
|
||
break
|
||
conversations.extend(item for item in batch if isinstance(item, dict))
|
||
if not payload.get("has_more"):
|
||
break
|
||
last_item = batch[-1]
|
||
last_id = (
|
||
str(last_item.get("id") or "")
|
||
if isinstance(last_item, dict)
|
||
else ""
|
||
)
|
||
if not last_id:
|
||
break
|
||
return [normalize_conversation(item) for item in conversations]
|
||
|
||
|
||
async def list_dify_conversations_page(
|
||
user_id: str,
|
||
*,
|
||
tenant_id: str,
|
||
cursor: str | None,
|
||
page_size: int,
|
||
keyword: str = "",
|
||
max_scanned: int = 500,
|
||
) -> dict[str, Any]:
|
||
"""按 Dify 游标向后扫描,收集当前租户的一页会话。"""
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
items: list[dict[str, Any]] = []
|
||
last_id = cursor
|
||
has_more = True
|
||
scanned = 0
|
||
normalized_keyword = keyword.strip().lower()
|
||
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
while has_more and len(items) < page_size and scanned < max_scanned:
|
||
response = await client.get(
|
||
f"{settings.dify_api_base}/conversations",
|
||
headers=_headers(key),
|
||
params={
|
||
"user": user_id,
|
||
"limit": min(100, max_scanned - scanned),
|
||
"sort_by": "-updated_at",
|
||
**({"last_id": last_id} if last_id else {}),
|
||
},
|
||
)
|
||
_raise_dify_error(response)
|
||
payload = response.json()
|
||
batch = payload.get("data", [])
|
||
if not isinstance(batch, list) or not batch:
|
||
has_more = False
|
||
break
|
||
|
||
stopped_inside_batch = False
|
||
for index, raw_item in enumerate(batch):
|
||
if not isinstance(raw_item, dict):
|
||
continue
|
||
scanned += 1
|
||
last_id = str(raw_item.get("id") or last_id or "")
|
||
item = normalize_conversation(raw_item)
|
||
if item["tenantId"] != tenant_id:
|
||
continue
|
||
if normalized_keyword and normalized_keyword not in item[
|
||
"title"
|
||
].lower():
|
||
continue
|
||
items.append(item)
|
||
if len(items) >= page_size:
|
||
stopped_inside_batch = index < len(batch) - 1
|
||
break
|
||
|
||
has_more = stopped_inside_batch or bool(payload.get("has_more"))
|
||
if not last_id:
|
||
has_more = False
|
||
|
||
return {
|
||
"items": items,
|
||
"nextCursor": last_id if has_more else None,
|
||
"hasMore": has_more,
|
||
}
|
||
|
||
|
||
async def get_dify_conversation(
|
||
conversation_id: str, user_id: str
|
||
) -> dict[str, Any]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
last_id: str | None = None
|
||
scanned = 0
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
while scanned < 5000:
|
||
response = await client.get(
|
||
f"{settings.dify_api_base}/conversations",
|
||
headers=_headers(key),
|
||
params={
|
||
"user": user_id,
|
||
"limit": 100,
|
||
"sort_by": "-updated_at",
|
||
**({"last_id": last_id} if last_id else {}),
|
||
},
|
||
)
|
||
_raise_dify_error(response)
|
||
payload = response.json()
|
||
batch = payload.get("data", [])
|
||
if not isinstance(batch, list) or not batch:
|
||
break
|
||
for raw_item in batch:
|
||
if not isinstance(raw_item, dict):
|
||
continue
|
||
scanned += 1
|
||
if str(raw_item.get("id") or "") == conversation_id:
|
||
return normalize_conversation(raw_item)
|
||
if not payload.get("has_more"):
|
||
break
|
||
last_id = str(batch[-1].get("id") or "")
|
||
if not last_id:
|
||
break
|
||
raise HTTPException(status_code=404, detail="Dify 会话不存在或无权访问")
|
||
|
||
|
||
def _message_files(item: dict[str, Any]) -> list[dict[str, Any]]:
|
||
files = item.get("message_files")
|
||
if not isinstance(files, list):
|
||
return []
|
||
return [
|
||
{
|
||
"id": str(file.get("id") or ""),
|
||
"name": str(file.get("name") or file.get("id") or "附件"),
|
||
"type": str(file.get("type") or "custom"),
|
||
"belongsTo": str(file.get("belongs_to") or "user"),
|
||
}
|
||
for file in files
|
||
if isinstance(file, dict) and file.get("id")
|
||
]
|
||
|
||
|
||
async def get_dify_message_records(
|
||
conversation_id: str, user_id: str
|
||
) -> list[dict[str, Any]]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
data: list[dict[str, Any]] = []
|
||
first_id: str | None = None
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
for _ in range(10):
|
||
params: dict[str, Any] = {
|
||
"conversation_id": conversation_id,
|
||
"user": user_id,
|
||
"limit": 100,
|
||
}
|
||
if first_id:
|
||
params["first_id"] = first_id
|
||
response = await client.get(
|
||
f"{settings.dify_api_base}/messages",
|
||
headers=_headers(key),
|
||
params=params,
|
||
)
|
||
_raise_dify_error(response)
|
||
payload = response.json()
|
||
batch = payload.get("data", [])
|
||
if not isinstance(batch, list) or not batch:
|
||
break
|
||
data.extend(item for item in batch if isinstance(item, dict))
|
||
if not payload.get("has_more"):
|
||
break
|
||
first_id = str(batch[0].get("id") or "")
|
||
if not first_id:
|
||
break
|
||
|
||
unique_data = {str(item.get("id")): item for item in data if item.get("id")}
|
||
data = sorted(unique_data.values(), key=lambda item: item.get("created_at", 0))
|
||
|
||
return data
|
||
|
||
|
||
async def get_dify_messages(
|
||
conversation_id: str, user_id: str
|
||
) -> list[dict[str, Any]]:
|
||
data = await get_dify_message_records(conversation_id, user_id)
|
||
messages: list[dict[str, Any]] = []
|
||
for item in data:
|
||
if not isinstance(item, dict):
|
||
continue
|
||
message_id = str(item.get("id") or "")
|
||
created_at = _iso_timestamp(item.get("created_at"))
|
||
query = str(item.get("query") or "")
|
||
if query or _message_files(item):
|
||
messages.append(
|
||
{
|
||
"id": f"{message_id}-user",
|
||
"role": "user",
|
||
"content": query,
|
||
"status": "completed",
|
||
"difyMessageId": message_id,
|
||
"createdAt": created_at,
|
||
"files": [
|
||
file
|
||
for file in _message_files(item)
|
||
if file["belongsTo"] == "user"
|
||
],
|
||
}
|
||
)
|
||
raw_answer = str(item.get("answer") or "")
|
||
cube_response = normalize_cube_response(raw_answer)
|
||
answer = cube_response["message"] if cube_response else raw_answer
|
||
if answer:
|
||
messages.append(
|
||
{
|
||
"id": f"{message_id}-assistant",
|
||
"role": "assistant",
|
||
"content": answer,
|
||
"status": "completed",
|
||
"difyMessageId": message_id,
|
||
"createdAt": created_at,
|
||
"files": [
|
||
file
|
||
for file in _message_files(item)
|
||
if file["belongsTo"] == "assistant"
|
||
],
|
||
}
|
||
)
|
||
return messages
|
||
|
||
|
||
async def get_latest_cube_response(
|
||
conversation_id: str, user_id: str
|
||
) -> dict[str, Any] | None:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.get(
|
||
f"{settings.dify_api_base}/messages",
|
||
headers=_headers(key),
|
||
params={
|
||
"conversation_id": conversation_id,
|
||
"user": user_id,
|
||
"limit": 20,
|
||
},
|
||
)
|
||
_raise_dify_error(response)
|
||
records = response.json().get("data", [])
|
||
records = records if isinstance(records, list) else []
|
||
records.sort(
|
||
key=lambda item: item.get("created_at", 0) if isinstance(item, dict) else 0,
|
||
reverse=True,
|
||
)
|
||
for item in records:
|
||
if not isinstance(item, dict):
|
||
continue
|
||
cube_response = normalize_cube_response(item.get("answer"))
|
||
if cube_response and cube_response["status"]:
|
||
return cube_response
|
||
return None
|
||
|
||
|
||
async def get_dify_conversation_variables(
|
||
conversation_id: str, user_id: str
|
||
) -> dict[str, Any]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.get(
|
||
f"{settings.dify_api_base}/conversations/{conversation_id}/variables",
|
||
headers=_headers(key),
|
||
params={"user": user_id, "limit": 100},
|
||
)
|
||
_raise_dify_error(response)
|
||
payload = response.json()
|
||
return {
|
||
item["name"]: item.get("value")
|
||
for item in payload.get("data", [])
|
||
if isinstance(item, dict) and item.get("name")
|
||
}
|
||
|
||
|
||
async def stream_dify_chat(
|
||
*,
|
||
user_id: str,
|
||
content: str,
|
||
conversation_id: str | None,
|
||
inputs: dict[str, Any],
|
||
files: list[dict[str, Any]],
|
||
) -> AsyncIterator[dict[str, Any]]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
payload = {
|
||
"query": content or "请分析上传的文件",
|
||
"inputs": inputs,
|
||
"files": [
|
||
{
|
||
"type": item.get("type", "custom"),
|
||
"transfer_method": "local_file",
|
||
"upload_file_id": item["id"],
|
||
}
|
||
for item in files
|
||
],
|
||
"response_mode": "streaming",
|
||
"conversation_id": conversation_id or "",
|
||
"user": user_id,
|
||
"auto_generate_name": True,
|
||
}
|
||
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
async with client.stream(
|
||
"POST",
|
||
f"{settings.dify_api_base}/chat-messages",
|
||
headers=_headers(key),
|
||
json=payload,
|
||
) as response:
|
||
if response.is_error:
|
||
body = (await response.aread()).decode("utf-8", errors="replace")
|
||
raise HTTPException(
|
||
status_code=response.status_code,
|
||
detail=f"Dify 请求失败:{body[:500]}",
|
||
)
|
||
async for line in response.aiter_lines():
|
||
if not line.startswith("data: "):
|
||
continue
|
||
try:
|
||
event = json.loads(line[6:])
|
||
except json.JSONDecodeError:
|
||
continue
|
||
if isinstance(event, dict):
|
||
yield event
|
||
|
||
|
||
async def stop_dify_task(task_id: str, user_id: str) -> dict[str, Any]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.post(
|
||
f"{settings.dify_api_base}/chat-messages/{task_id}/stop",
|
||
headers=_headers(key),
|
||
json={"user": user_id},
|
||
)
|
||
_raise_dify_error(response)
|
||
return response.json()
|
||
|
||
|
||
async def upload_dify_file(file: UploadFile, user_id: str) -> dict[str, Any]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
content = await file.read()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.post(
|
||
f"{settings.dify_api_base}/files/upload",
|
||
headers=_headers(key, None),
|
||
data={"user": user_id},
|
||
files={
|
||
"file": (
|
||
file.filename or "upload.xlsx",
|
||
content,
|
||
file.content_type or "application/octet-stream",
|
||
)
|
||
},
|
||
)
|
||
_raise_dify_error(response)
|
||
return response.json()
|
||
|
||
|
||
async def rename_dify_conversation(
|
||
conversation_id: str, user_id: str, name: str
|
||
) -> dict[str, Any]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.post(
|
||
f"{settings.dify_api_base}/conversations/{conversation_id}/name",
|
||
headers=_headers(key),
|
||
json={"name": name, "auto_generate": False, "user": user_id},
|
||
)
|
||
_raise_dify_error(response)
|
||
return normalize_conversation(response.json())
|
||
|
||
|
||
async def delete_dify_conversation(conversation_id: str, user_id: str) -> None:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.request(
|
||
"DELETE",
|
||
f"{settings.dify_api_base}/conversations/{conversation_id}",
|
||
headers=_headers(key),
|
||
json={"user": user_id},
|
||
)
|
||
_raise_dify_error(response)
|
||
|
||
|
||
async def preview_dify_file(
|
||
file_id: str, *, as_attachment: bool = False
|
||
) -> tuple[bytes, dict[str, str]]:
|
||
settings = get_cube_report_settings()
|
||
key = _require_dify_key()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.get(
|
||
f"{settings.dify_api_base}/files/{file_id}/preview",
|
||
headers=_headers(key),
|
||
params={"as_attachment": str(as_attachment).lower()},
|
||
)
|
||
_raise_dify_error(response)
|
||
headers = {
|
||
name: value
|
||
for name, value in response.headers.items()
|
||
if name.lower() in {"content-disposition", "content-length", "cache-control"}
|
||
}
|
||
headers["content-type"] = response.headers.get(
|
||
"content-type", "application/octet-stream"
|
||
)
|
||
return response.content, headers
|
||
|
||
|
||
def normalize_cube_load(value: Any) -> dict[str, Any] | None:
|
||
response = normalize_cube_response(value)
|
||
if response:
|
||
return response["query"]
|
||
current = value
|
||
for _ in range(3):
|
||
if isinstance(current, str):
|
||
text = current.strip()
|
||
if not text:
|
||
return None
|
||
try:
|
||
current = json.loads(text)
|
||
except json.JSONDecodeError as error:
|
||
raise ValueError(f"Dify 返回的 Cube Load 不是合法 JSON:{error.msg}") from error
|
||
continue
|
||
break
|
||
if isinstance(current, dict) and isinstance(current.get("query"), dict):
|
||
current = current["query"]
|
||
return current if isinstance(current, dict) else None
|
||
|
||
|
||
def _walk_filter_nodes(filters: Any):
|
||
if not isinstance(filters, list):
|
||
return
|
||
for item in filters:
|
||
if not isinstance(item, dict):
|
||
continue
|
||
if "member" in item:
|
||
yield item
|
||
for key in ("and", "or"):
|
||
children = item.get(key)
|
||
if isinstance(children, list):
|
||
yield from _walk_filter_nodes(children)
|
||
|
||
|
||
def ensure_tenant_filter(query: dict[str, Any], tenant_id: str) -> None:
|
||
tenant_member = None
|
||
for item in _walk_filter_nodes(query.get("filters", [])):
|
||
member = str(item.get("member", ""))
|
||
values = [str(value) for value in item.get("values", [])]
|
||
operator = str(item.get("operator") or "").lower()
|
||
if (
|
||
re.search(r"tenant|租户", member, re.IGNORECASE)
|
||
and operator == "equals"
|
||
and tenant_id in values
|
||
):
|
||
tenant_member = member
|
||
break
|
||
if not tenant_member:
|
||
raise HTTPException(
|
||
status_code=422,
|
||
detail="Cube Query 缺少当前租户的 equals 过滤条件,已阻止数据查询",
|
||
)
|
||
|
||
# 顶层 filters 数组按 AND 组合。即使 Dify 将原租户条件放进 OR,
|
||
# 这里也会额外增加不可绕过的租户约束。
|
||
filters = query.get("filters")
|
||
if not isinstance(filters, list):
|
||
filters = []
|
||
query["filters"] = filters
|
||
filters.append(
|
||
{
|
||
"member": tenant_member,
|
||
"operator": "equals",
|
||
"values": [tenant_id],
|
||
}
|
||
)
|
||
|
||
|
||
def _mask_value(key: str, title: str, value: Any) -> Any:
|
||
if value is None:
|
||
return value
|
||
text = str(value)
|
||
field = f"{key} {title}"
|
||
if re.search(r"idcard|身份证", field, re.IGNORECASE) and len(text) >= 10:
|
||
return f"{text[:6]}********{text[-4:]}"
|
||
if re.search(r"phone|mobile|手机号", field, re.IGNORECASE) and len(text) == 11:
|
||
return f"{text[:3]}****{text[-4:]}"
|
||
return value
|
||
|
||
|
||
def _positive_int(value: Any) -> int | None:
|
||
if isinstance(value, bool):
|
||
return None
|
||
if isinstance(value, (int, float)):
|
||
return max(int(value), 0)
|
||
return None
|
||
|
||
|
||
def _ensure_page_order(query: dict[str, Any]) -> None:
|
||
if query.get("order"):
|
||
return
|
||
dimensions = query.get("dimensions")
|
||
if isinstance(dimensions, list) and dimensions:
|
||
query["order"] = {str(dimensions[0]): "asc"}
|
||
return
|
||
time_dimensions = query.get("timeDimensions")
|
||
if isinstance(time_dimensions, list):
|
||
for item in time_dimensions:
|
||
if isinstance(item, dict) and item.get("dimension"):
|
||
query["order"] = {str(item["dimension"]): "asc"}
|
||
return
|
||
|
||
|
||
def _validate_read_only_sql(sql: str) -> None:
|
||
normalized = re.sub(r"\s+", " ", sql).strip().lower()
|
||
if not (normalized.startswith("select ") or normalized.startswith("with ")):
|
||
raise HTTPException(status_code=502, detail="Cube 返回了非查询 SQL")
|
||
if ";" in normalized.rstrip(";") or re.search(
|
||
r"\b(insert|update|delete|drop|alter|create|truncate|merge|exec|execute)\b",
|
||
normalized,
|
||
):
|
||
raise HTTPException(status_code=502, detail="Cube 返回的 SQL 未通过只读校验")
|
||
|
||
|
||
async def _generate_cube_sql(query: dict[str, Any]) -> dict[str, Any]:
|
||
settings = get_cube_report_settings()
|
||
async with httpx.AsyncClient(timeout=_timeout()) as client:
|
||
response = await client.post(
|
||
f"{settings.cube_api_base}/sql",
|
||
headers=_headers(settings.cube_api_token),
|
||
json={"query": query},
|
||
)
|
||
if response.is_error:
|
||
raise HTTPException(
|
||
status_code=502,
|
||
detail=f"Cube SQL 生成失败:{response.text[:500]}",
|
||
)
|
||
payload = response.json()
|
||
sql_node = payload.get("sql") if isinstance(payload, dict) else None
|
||
sql_tuple = sql_node.get("sql") if isinstance(sql_node, dict) else None
|
||
if (
|
||
not isinstance(sql_tuple, list)
|
||
or not sql_tuple
|
||
or not isinstance(sql_tuple[0], str)
|
||
):
|
||
raise HTTPException(status_code=502, detail="Cube 没有返回可执行 SQL")
|
||
raw_params = sql_tuple[1] if len(sql_tuple) > 1 else []
|
||
if raw_params is None:
|
||
raw_params = []
|
||
elif not isinstance(raw_params, list):
|
||
raw_params = [raw_params]
|
||
aliases = sql_node.get("aliasNameToMember", {})
|
||
return {
|
||
"sql": sql_tuple[0].strip().rstrip(";"),
|
||
"params": raw_params,
|
||
"aliases": aliases if isinstance(aliases, dict) else {},
|
||
}
|
||
|
||
|
||
def _prepare_cube_sql(raw_sql: str, values: list[Any]) -> tuple[str, dict[str, Any]]:
|
||
bindings: dict[str, Any] = {}
|
||
|
||
def replace(match: re.Match[str]) -> str:
|
||
index = int(match.group(1))
|
||
if index < 1 or index > len(values):
|
||
raise HTTPException(status_code=502, detail="Cube SQL 参数数量不匹配")
|
||
name = f"cube_param_{index}"
|
||
bindings[name] = values[index - 1]
|
||
return f":{name}"
|
||
|
||
sql = re.sub(r"@_(\d+)", replace, raw_sql)
|
||
_validate_read_only_sql(sql)
|
||
return sql, bindings
|
||
|
||
|
||
def _sql_literal(value: Any) -> str:
|
||
if value is None:
|
||
return "NULL"
|
||
if isinstance(value, bool):
|
||
return "1" if value else "0"
|
||
if isinstance(value, (int, float)):
|
||
return str(value)
|
||
text = str(value).replace("'", "''")
|
||
return f"N'{text}'"
|
||
|
||
|
||
def _display_sql(raw_sql: str, values: list[Any]) -> str:
|
||
def replace(match: re.Match[str]) -> str:
|
||
index = int(match.group(1))
|
||
return (
|
||
_sql_literal(values[index - 1])
|
||
if 0 < index <= len(values)
|
||
else match.group(0)
|
||
)
|
||
|
||
return re.sub(r"@_(\d+)", replace, raw_sql)
|
||
|
||
|
||
def _column_title(
|
||
column: str,
|
||
aliases: dict[str, Any],
|
||
member_titles: dict[str, str],
|
||
) -> str:
|
||
member = str(aliases.get(column) or column)
|
||
for candidate in cube_member_candidates(member, column):
|
||
title = member_titles.get(candidate)
|
||
if title:
|
||
return title
|
||
return member
|
||
|
||
|
||
async def query_cube_page(
|
||
*,
|
||
base_load: dict[str, Any],
|
||
tenant_id: str | None,
|
||
page: int,
|
||
page_size: int,
|
||
limit_source: str | None,
|
||
) -> dict[str, Any]:
|
||
settings = get_cube_report_settings()
|
||
page_size = min(max(page_size, 1), settings.max_page_size)
|
||
page_offset = (page - 1) * page_size
|
||
canonical_query = deepcopy(base_load)
|
||
if not tenant_id:
|
||
raise HTTPException(status_code=422, detail="当前会话缺少租户信息")
|
||
ensure_tenant_filter(canonical_query, tenant_id)
|
||
|
||
is_user_limit = str(limit_source).lower() == "user"
|
||
semantic_limit = (
|
||
_positive_int(canonical_query.get("limit")) if is_user_limit else None
|
||
)
|
||
semantic_offset = (
|
||
_positive_int(canonical_query.get("offset")) or 0 if is_user_limit else 0
|
||
)
|
||
if semantic_limit is not None and page_offset >= semantic_limit:
|
||
return {
|
||
"columns": [],
|
||
"rows": [],
|
||
"page": page,
|
||
"pageSize": page_size,
|
||
"total": semantic_limit,
|
||
"sql": "",
|
||
"limitSource": "user",
|
||
}
|
||
|
||
page_query = deepcopy(canonical_query)
|
||
page_query["offset"] = semantic_offset + page_offset
|
||
page_query["limit"] = (
|
||
min(page_size, semantic_limit - page_offset)
|
||
if semantic_limit is not None
|
||
else page_size
|
||
)
|
||
page_query.pop("total", None)
|
||
_ensure_page_order(page_query)
|
||
|
||
count_query = deepcopy(canonical_query)
|
||
for key in ("limit", "offset", "order", "total"):
|
||
count_query.pop(key, None)
|
||
|
||
page_sql, count_sql, member_titles = await asyncio.gather(
|
||
_generate_cube_sql(page_query),
|
||
_generate_cube_sql(count_query),
|
||
get_cube_member_titles(),
|
||
)
|
||
executable_sql, bindings = _prepare_cube_sql(page_sql["sql"], page_sql["params"])
|
||
count_executable, count_bindings = _prepare_cube_sql(
|
||
count_sql["sql"], count_sql["params"]
|
||
)
|
||
wrapped_count_sql = (
|
||
"SELECT COUNT_BIG(1) AS __cube_total FROM (\n"
|
||
f"{count_executable}\n"
|
||
") AS __cube_count"
|
||
)
|
||
(keys, rows), raw_total = await asyncio.gather(
|
||
asyncio.to_thread(execute_parameterized_sql, executable_sql, bindings),
|
||
asyncio.to_thread(
|
||
execute_parameterized_scalar, wrapped_count_sql, count_bindings
|
||
),
|
||
)
|
||
columns = [
|
||
{
|
||
"key": key,
|
||
"title": _column_title(
|
||
key,
|
||
page_sql["aliases"],
|
||
member_titles,
|
||
),
|
||
"type": "text",
|
||
}
|
||
for key in keys
|
||
]
|
||
normalized_rows = []
|
||
for row_index, row in enumerate(rows):
|
||
normalized = {"__rowKey": f"{page}-{row_index}"}
|
||
for column in columns:
|
||
key = column["key"]
|
||
normalized[key] = _mask_value(key, column["title"], row.get(key))
|
||
normalized_rows.append(normalized)
|
||
|
||
try:
|
||
total = max(int(raw_total or 0) - semantic_offset, 0)
|
||
except (TypeError, ValueError):
|
||
total = len(normalized_rows)
|
||
if semantic_limit is not None:
|
||
total = min(total, semantic_limit)
|
||
return {
|
||
"columns": columns,
|
||
"rows": normalized_rows,
|
||
"page": page,
|
||
"pageSize": page_size,
|
||
"total": total,
|
||
"sql": _display_sql(page_sql["sql"], page_sql["params"]),
|
||
"limitSource": "user" if is_user_limit else "system",
|
||
}
|
||
|
||
|
||
async def query_cube_export(
|
||
*,
|
||
base_load: dict[str, Any],
|
||
tenant_id: str | None,
|
||
limit_source: str | None,
|
||
column_keys: list[str] | None = None,
|
||
) -> dict[str, Any]:
|
||
query = deepcopy(base_load)
|
||
if not tenant_id:
|
||
raise HTTPException(status_code=422, detail="当前会话缺少租户信息")
|
||
ensure_tenant_filter(query, tenant_id)
|
||
if str(limit_source).lower() != "user":
|
||
query.pop("limit", None)
|
||
query.pop("offset", None)
|
||
query.pop("total", None)
|
||
|
||
sql_info, member_titles = await asyncio.gather(
|
||
_generate_cube_sql(query),
|
||
get_cube_member_titles(),
|
||
)
|
||
executable_sql, bindings = _prepare_cube_sql(sql_info["sql"], sql_info["params"])
|
||
keys, rows = await asyncio.to_thread(
|
||
execute_parameterized_sql, executable_sql, bindings
|
||
)
|
||
headers = {
|
||
key: _column_title(key, sql_info["aliases"], member_titles)
|
||
for key in keys
|
||
}
|
||
if column_keys:
|
||
selected_keys = list(dict.fromkeys(str(key).strip() for key in column_keys))
|
||
invalid_keys = [key for key in selected_keys if key not in headers]
|
||
if invalid_keys:
|
||
raise HTTPException(
|
||
status_code=422,
|
||
detail=f"导出列不存在或已失效:{', '.join(invalid_keys[:5])}",
|
||
)
|
||
headers = {key: headers[key] for key in selected_keys}
|
||
rows = [
|
||
{key: row.get(key) for key in selected_keys}
|
||
for row in rows
|
||
]
|
||
return {
|
||
"headers": headers,
|
||
"rows": rows,
|
||
"sql": _display_sql(sql_info["sql"], sql_info["params"]),
|
||
}
|