Files
group_fqcd_jr/tests/integration/test_offsite_nl2sql_fields.py
T

304 lines
11 KiB
Python

"""单据 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 = ""