"""单据 NL2SQL 返回字段查询接口的集成测试。 覆盖:申购单据查询成功时返回查询侧字段,赎回单不查询的字段标记为未查询, 查询失败时明确标记 query_failed,尚未核对时标记 pending,单据不存在、越权的处理, 以及查询动作必须留下审计记录。 统一在单事件循环内调用应用(httpx.ASGITransport),避免连接池里的连接被 跨事件循环复用。 """ from collections.abc import AsyncIterator from datetime import UTC, datetime from uuid import uuid4 import httpx import pytest from sqlalchemy import delete, select from sqlalchemy.ext.asyncio import AsyncSession from app.api.dependencies.auth import build_request_context from app.api.dependencies.database import get_session from app.core.contracts import RequestContext from app.infrastructure.db import SessionFactory, engine from app.main import app from app.model.audit import InteractionAudit from app.model.offsite_fund import ( OffsiteFundDocument, OffsiteQueryRecord, OffsiteRuleResult, ) FIELDS_PATH = "/api/v1/offsite-fund/documents/{task_id}/nl2sql-fields" TRACE_ID = "" @pytest.fixture(autouse=True) async def _dispose_engine_after_test() -> AsyncIterator[None]: """用例结束后释放连接池,避免连接绑定在已关闭的事件循环上。""" yield await engine.dispose() async def _override_session() -> AsyncIterator[AsyncSession]: async with SessionFactory() as session: yield session def _install_context(roles: tuple[str, ...], permissions: tuple[str, ...]) -> None: async def override_context() -> RequestContext: return RequestContext( user_id="1", trace_id=TRACE_ID, roles=roles, permissions=permissions, data_scope="all", ) app.dependency_overrides[build_request_context] = override_context app.dependency_overrides[get_session] = _override_session async def _get_fields(task_id: str) -> httpx.Response: transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: return await client.get(FIELDS_PATH.format(task_id=task_id)) async def _seed_document(document_type: str = "subscription") -> str: task_id = f"T{uuid4().hex[:14]}" now = datetime.now(UTC).replace(tzinfo=None) async with SessionFactory() as session, session.begin(): session.add(OffsiteFundDocument( task_id=task_id, mail_id=f"M{uuid4().hex[:12]}", attachment_id=f"A{uuid4().hex[:14]}", document_type=document_type, fund_code="15911", account_identifier="10001", status="planned", operator_decision="未处理", created_at=now, updated_at=now, )) return task_id async def _seed_success_queries(task_id: str) -> None: now = datetime.now(UTC).replace(tzinfo=None) async with SessionFactory() as session, session.begin(): session.add(OffsiteQueryRecord( task_id=task_id, rule_code="subscription_holding_ratio", natural_language_request="基金代码为15911,查询基金最新总份额、最新净值和申请前持有份额", script_path="nl2sql_yc.py", result_summary={ "status": "success", "data": {"total": 1, "rows": [{ "nav": "1.250000", "total_fund_shares": "10000000000.0000", "total_quantity": "100000000.0000", }]}, }, status="success", error_message=None, created_at=now, )) session.add(OffsiteRuleResult( task_id=task_id, rule_code="subscription_holding_ratio", rule_name="申购后单一投资者持有比例", result="正常", document_value={"申购金额元": "200000000.00"}, database_value={ "最新净值": "1.250000", "基金最新总份额": "10000000000.0000", "申请前持有份额": "100000000.0000", }, calculation={"申购后持有比例": "0.026"}, created_at=now, )) session.add(OffsiteRuleResult( task_id=task_id, rule_code="subscription_minimum_amount", rule_name="申购最低金额", result="正常", document_value={"申购金额元": "200000000.00"}, database_value={}, calculation={"判断口径": "标准化申购金额 <= 1 元为异常"}, created_at=now, )) async def _seed_failed_queries(task_id: str) -> None: now = datetime.now(UTC).replace(tzinfo=None) async with SessionFactory() as session, session.begin(): session.add(OffsiteQueryRecord( task_id=task_id, rule_code="subscription_holding_ratio", natural_language_request="基金代码为15911,查询基金最新总份额、最新净值和申请前持有份额", script_path="nl2sql_yc.py", result_summary={"status": "query_failed", "message": "查询无可用数据"}, status="query_failed", error_message="查询无可用数据", created_at=now, )) session.add(OffsiteRuleResult( task_id=task_id, rule_code="subscription_holding_ratio", rule_name="申购后单一投资者持有比例", result="无法判断", document_value={}, database_value={}, calculation={"原因": "依赖的 NL2SQL 查询失败或无可用数据"}, created_at=now, )) async def _count_audit(action_type: str) -> int: async with SessionFactory() as session: rows = await session.execute(select(InteractionAudit).where( InteractionAudit.action_type == action_type, InteractionAudit.detail["trace_id"].as_string() == TRACE_ID, )) return len(rows.scalars().all()) async def _cleanup(task_id: str) -> None: async with SessionFactory() as session, session.begin(): if TRACE_ID: await session.execute(delete(InteractionAudit).where( InteractionAudit.detail["trace_id"].as_string() == TRACE_ID )) if task_id: await session.execute(delete(OffsiteQueryRecord).where( OffsiteQueryRecord.task_id == task_id )) await session.execute(delete(OffsiteRuleResult).where( OffsiteRuleResult.task_id == task_id )) await session.execute(delete(OffsiteFundDocument).where( OffsiteFundDocument.task_id == task_id )) @pytest.mark.integration async def test_nl2sql_fields_return_queried_values_for_subscription() -> None: global TRACE_ID TRACE_ID = f"trace-nl2sql-fields-{uuid4()}" task_id = "" try: task_id = await _seed_document("subscription") await _seed_success_queries(task_id) _install_context(("operator",), ("offsite:read",)) response = await _get_fields(task_id) assert response.status_code == 200 body = response.json() assert body["code"] == 0 data = body["data"] assert data["task_id"] == task_id assert data["document_type"] == "subscription" assert data["fund_code"] == "15911" assert data["fields"]["最新净值"] == "1.250000" assert data["fields"]["基金最新总份额"] == "10000000000.0000" assert data["fields"]["申请前持有份额"] == "100000000.0000" # 申购单不查询可用份额,必须标为未查询而不是查询失败。 assert data["fields"]["当前最新可用份额"] is None assert data["field_status"]["当前最新可用份额"] == "not_queried" assert data["field_status"]["最新净值"] == "success" assert data["missing_fields"] == ["当前最新可用份额"] assert data["queries"][0]["rule_code"] == "subscription_holding_ratio" assert data["queries"][0]["status"] == "success" assert data["queries"][0]["row_count"] == 1 assert await _count_audit("offsite.nl2sql_fields_viewed") == 1 finally: await _cleanup(task_id) app.dependency_overrides.clear() TRACE_ID = "" @pytest.mark.integration async def test_nl2sql_fields_mark_query_failed_when_query_blocked() -> None: global TRACE_ID TRACE_ID = f"trace-nl2sql-fields-{uuid4()}" task_id = "" try: task_id = await _seed_document("subscription") await _seed_failed_queries(task_id) _install_context(("operator",), ("offsite:read",)) response = await _get_fields(task_id) data = response.json()["data"] assert data["fields"]["最新净值"] is None assert data["field_status"]["最新净值"] == "query_failed" assert data["field_status"]["基金最新总份额"] == "query_failed" assert data["field_status"]["申请前持有份额"] == "query_failed" assert data["field_status"]["当前最新可用份额"] == "not_queried" assert data["queries"][0]["status"] == "query_failed" assert data["queries"][0]["row_count"] == 0 finally: await _cleanup(task_id) app.dependency_overrides.clear() TRACE_ID = "" @pytest.mark.integration async def test_nl2sql_fields_mark_pending_before_verification() -> None: """尚未触发核对时字段为空,必须标为 pending,不能伪装成查询失败。""" global TRACE_ID TRACE_ID = f"trace-nl2sql-fields-{uuid4()}" task_id = "" try: task_id = await _seed_document("redemption") _install_context(("operator",), ("offsite:read",)) response = await _get_fields(task_id) data = response.json()["data"] assert data["fields"]["基金最新总份额"] is None assert data["field_status"]["基金最新总份额"] == "pending" assert data["field_status"]["当前最新可用份额"] == "pending" assert data["field_status"]["最新净值"] == "not_queried" assert data["queries"] == [] finally: await _cleanup(task_id) app.dependency_overrides.clear() TRACE_ID = "" @pytest.mark.integration async def test_nl2sql_fields_unknown_document_returns_not_found() -> None: global TRACE_ID TRACE_ID = f"trace-nl2sql-fields-{uuid4()}" try: _install_context(("operator",), ("offsite:read",)) response = await _get_fields("T000000000000") assert response.status_code == 200 assert response.json()["code"] == 404 assert response.json()["message"] == "单据不存在" finally: app.dependency_overrides.clear() TRACE_ID = "" @pytest.mark.integration async def test_nl2sql_fields_requires_offsite_permission() -> None: global TRACE_ID TRACE_ID = f"trace-nl2sql-fields-{uuid4()}" task_id = "" try: task_id = await _seed_document("subscription") _install_context(("operator",), ()) missing_permission = await _get_fields(task_id) _install_context(("customer",), ("offsite:read",)) wrong_role = await _get_fields(task_id) assert missing_permission.json()["code"] == 403 assert wrong_role.json()["code"] == 403 finally: await _cleanup(task_id) app.dependency_overrides.clear() TRACE_ID = ""