295 lines
11 KiB
Python
295 lines
11 KiB
Python
"""`query_customer_profile` 工具 + 画像字段投影的单元测试。
|
||||
|
|
|
|||
|
|
覆盖四条要求,两个方向都要有牙:
|
|||
|
|
|
|||
|
|
1. **字段策略**:只投影白名单字段,**绝不返回 PII**
|
|||
|
|
(`real_name`/`birth_date`/`mobile_masked`);且白名单是**白名单**——
|
|||
|
|
快照里新增未知字段不会自动外泄。
|
|||
|
|
2. **测评有效期实时判定**:`assessment_expired` 按**当前时间**重算,不信快照里的布尔值。
|
|||
|
|
3. **失败关闭**:无当前画像 / 快照非法 → **抛错**,不返回 `{}`
|
|||
|
|
(空画像会被下游误读成"该客户无偏好")。
|
|||
|
|
4. **越范围不泄露**:他人客户一律按"不存在"处理,且拒绝时不得暴露目标客户是否存在。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
from datetime import UTC, datetime, timedelta
|
|||
|
|
from typing import Any
|
|||
|
|
|
|||
|
|
import pytest
|
|||
|
|
|
|||
|
|
from app.core.contracts import RequestContext
|
|||
|
|
from app.core.errors import ForbiddenAgentError, GenericResourceNotFoundError
|
|||
|
|
from app.core.profile_projection import project_profile
|
|||
|
|
from app.service.customer_profile_service import (
|
|||
|
|
ALLOWED_FIELDS,
|
|||
|
|
AUDIT_ACTION,
|
|||
|
|
CustomerProfileQuery,
|
|||
|
|
CustomerProfileService,
|
|||
|
|
query_customer_profile_tool,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
NOW = datetime(2026, 9, 10, 12, 0, tzinfo=UTC)
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
# 替身
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
class FakeSession:
|
|||
|
|
"""只实现本服务用到的两件事:读快照与写审计。"""
|
|||
|
|
|
|||
|
|
def __init__(self) -> None:
|
|||
|
|
self.added: list[Any] = []
|
|||
|
|
|
|||
|
|
async def __aenter__(self) -> "FakeSession":
|
|||
|
|
return self
|
|||
|
|
|
|||
|
|
async def __aexit__(self, *args: object) -> None:
|
|||
|
|
return None
|
|||
|
|
|
|||
|
|
def begin(self) -> "FakeSession":
|
|||
|
|
return self
|
|||
|
|
|
|||
|
|
def add(self, item: Any) -> None:
|
|||
|
|
self.added.append(item)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def context(
|
|||
|
|
user_id: str = "9101",
|
|||
|
|
*,
|
|||
|
|
scopes: dict[str, str] | None = None,
|
|||
|
|
customer_ids: tuple[str, ...] = (),
|
|||
|
|
permissions: tuple[str, ...] = ("memory:read:self", "memory:read:customer"),
|
|||
|
|
) -> RequestContext:
|
|||
|
|
return RequestContext(
|
|||
|
|
user_id=user_id,
|
|||
|
|
trace_id="trace-1",
|
|||
|
|
roles=("customer",),
|
|||
|
|
permissions=permissions,
|
|||
|
|
permission_scopes=scopes or {"memory:read:self": "self", "memory:read:customer": "self"},
|
|||
|
|
customer_ids=customer_ids,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def snapshot(**overrides: Any) -> dict[str, Any]:
|
|||
|
|
base: dict[str, Any] = {
|
|||
|
|
"investor_type": "C5",
|
|||
|
|
"investment_horizon": "long_term",
|
|||
|
|
"trading_frequency": "high",
|
|||
|
|
"preferred_asset_class": ["equity_fund"],
|
|||
|
|
"risk_tags": ["aggressive"],
|
|||
|
|
"customer_tier": "diamond",
|
|||
|
|
"behavior_score": 88,
|
|||
|
|
"total_asset": "12800000.00",
|
|||
|
|
"assessment_valid_until": (NOW + timedelta(days=100)).isoformat(),
|
|||
|
|
"assessment_expired": False,
|
|||
|
|
}
|
|||
|
|
base.update(overrides)
|
|||
|
|
return base
|
|||
|
|
|
|||
|
|
|
|||
|
|
def wire(monkeypatch: pytest.MonkeyPatch, rows: list[dict[str, Any]]) -> FakeSession:
|
|||
|
|
"""把 `PlatformRepository.rows` 换成返回给定行的替身。"""
|
|||
|
|
session = FakeSession()
|
|||
|
|
|
|||
|
|
class FakeRepo:
|
|||
|
|
def __init__(self, _session: Any) -> None:
|
|||
|
|
pass
|
|||
|
|
|
|||
|
|
async def rows(self, *_a: Any, **_kw: Any) -> list[dict[str, Any]]:
|
|||
|
|
return rows
|
|||
|
|
|
|||
|
|
monkeypatch.setattr("app.service.customer_profile_service.PlatformRepository", FakeRepo)
|
|||
|
|
monkeypatch.setattr(
|
|||
|
|
"app.service.customer_profile_service.SessionFactory", lambda: session
|
|||
|
|
)
|
|||
|
|
return session
|
|||
|
|
|
|||
|
|
|
|||
|
|
def row(payload: dict[str, Any] | None = None, version: int = 1) -> dict[str, Any]:
|
|||
|
|
return {"snapshot": payload if payload is not None else snapshot(), "version": version}
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
# 1. 字段策略
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_projection_only_exposes_whitelisted_fields() -> None:
|
|||
|
|
"""白名单:快照里混入 PII 与未知字段都不得外泄。"""
|
|||
|
|
dirty = snapshot(
|
|||
|
|
real_name="陈宏远",
|
|||
|
|
birth_date="1972-06-15",
|
|||
|
|
mobile_masked="138****0101",
|
|||
|
|
trade_account="TA91010001",
|
|||
|
|
some_future_field="会外泄吗",
|
|||
|
|
)
|
|||
|
|
projected = project_profile(dirty, now=NOW)
|
|||
|
|
|
|||
|
|
assert set(projected) <= ALLOWED_FIELDS, "投影结果出现了白名单外的字段"
|
|||
|
|
leaks = ("real_name", "birth_date", "mobile_masked", "trade_account", "some_future_field")
|
|||
|
|
for leaked in leaks:
|
|||
|
|
assert leaked not in projected, f"{leaked} 不该外泄"
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_projection_returns_empty_for_malformed_snapshot() -> None:
|
|||
|
|
assert project_profile(None, now=NOW) == {}
|
|||
|
|
assert project_profile([], now=NOW) == {}
|
|||
|
|
assert project_profile("not-a-dict", now=NOW) == {}
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
# 2. 测评有效期实时判定
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_assessment_expired_is_recomputed_not_trusted() -> None:
|
|||
|
|
"""**反向断言**:快照谎报 `assessment_expired=False`,但有效期已过 → 必须判为过期。"""
|
|||
|
|
lying = snapshot(
|
|||
|
|
assessment_valid_until=(NOW - timedelta(days=1)).isoformat(),
|
|||
|
|
assessment_expired=False, # ← 快照里的值是错的
|
|||
|
|
)
|
|||
|
|
projected = project_profile(lying, now=NOW)
|
|||
|
|
assert projected["assessment_expired"] is True, "不能信快照里的布尔值"
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_assessment_valid_until_just_now_counts_as_expired() -> None:
|
|||
|
|
"""边界:有效期正好等于当前时间视为**已过期**(与 SuitabilityService 口径一致)。"""
|
|||
|
|
projected = project_profile(
|
|||
|
|
snapshot(assessment_valid_until=NOW.isoformat()), now=NOW
|
|||
|
|
)
|
|||
|
|
assert projected["assessment_expired"] is True
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_unparseable_valid_until_keeps_field_but_does_not_claim_expired() -> None:
|
|||
|
|
"""有效期无法解析时不臆造"过期"判定(也不删除原始值)。"""
|
|||
|
|
projected = project_profile(snapshot(assessment_valid_until="昨天"), now=NOW)
|
|||
|
|
assert projected["assessment_valid_until"] == "昨天"
|
|||
|
|
assert "assessment_expired" not in projected or projected["assessment_expired"] is False
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
# 3. 失败关闭
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_missing_snapshot_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
"""无当前画像必须**抛错**,不能返回空画像。"""
|
|||
|
|
wire(monkeypatch, [])
|
|||
|
|
with pytest.raises(GenericResourceNotFoundError, match="暂无当前画像"):
|
|||
|
|
await query_customer_profile_tool(CustomerProfileQuery(customer_id="9101"), context())
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_empty_snapshot_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
"""快照存在但内容非法(投影为空)同样失败关闭。"""
|
|||
|
|
wire(monkeypatch, [row(payload={})])
|
|||
|
|
with pytest.raises(GenericResourceNotFoundError):
|
|||
|
|
await query_customer_profile_tool(CustomerProfileQuery(customer_id="9101"), context())
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
# 4. 鉴权与数据范围
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_customer_can_read_own_profile(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
session = wire(monkeypatch, [row()])
|
|||
|
|
out = await query_customer_profile_tool(
|
|||
|
|
CustomerProfileQuery(customer_id="9101"), context("9101")
|
|||
|
|
)
|
|||
|
|
assert out["profile"]["investor_type"] == "C5"
|
|||
|
|
assert out["customer_id"] == "9101"
|
|||
|
|
# 审计与读取同事务落库
|
|||
|
|
assert [a.action_type for a in session.added] == [AUDIT_ACTION]
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_reading_another_customer_without_scope_is_denied_as_not_found(
|
|||
|
|
monkeypatch: pytest.MonkeyPatch,
|
|||
|
|
) -> None:
|
|||
|
|
"""越范围一律按"不存在"处理,不泄露目标客户是否存在。"""
|
|||
|
|
wire(monkeypatch, [row()])
|
|||
|
|
with pytest.raises(GenericResourceNotFoundError, match="客户不可访问"):
|
|||
|
|
await query_customer_profile_tool(
|
|||
|
|
CustomerProfileQuery(customer_id="9103"), context("9101")
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_own_customers_scope_allows_assigned_customer(
|
|||
|
|
monkeypatch: pytest.MonkeyPatch,
|
|||
|
|
) -> None:
|
|||
|
|
"""`own_customers` 且客户确实归属调用者时才放行。"""
|
|||
|
|
wire(monkeypatch, [row()])
|
|||
|
|
scopes = {"memory:read:self": "self", "memory:read:customer": "own_customers"}
|
|||
|
|
out = await query_customer_profile_tool(
|
|||
|
|
CustomerProfileQuery(customer_id="9103"),
|
|||
|
|
context("9002", scopes=scopes, customer_ids=("9103",)),
|
|||
|
|
)
|
|||
|
|
assert out["customer_id"] == "9103"
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_all_scope_allows_any_customer(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
wire(monkeypatch, [row()])
|
|||
|
|
out = await query_customer_profile_tool(
|
|||
|
|
CustomerProfileQuery(customer_id="9999"),
|
|||
|
|
context("9003", scopes={"memory:read:self": "self", "memory:read:customer": "all"}),
|
|||
|
|
)
|
|||
|
|
assert out["customer_id"] == "9999"
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_missing_permission_is_denied(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
"""连权限都没有:直接 403 语义(`ForbiddenAgentError`),并且会写权限拒绝审计。"""
|
|||
|
|
wire(monkeypatch, [row()])
|
|||
|
|
with pytest.raises(ForbiddenAgentError):
|
|||
|
|
await query_customer_profile_tool(
|
|||
|
|
CustomerProfileQuery(customer_id="9101"), context("9101", permissions=())
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
# 5. 字段筛选
|
|||
|
|
# ---------------------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_requested_fields_filter_the_projection(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
wire(monkeypatch, [row()])
|
|||
|
|
out = await query_customer_profile_tool(
|
|||
|
|
CustomerProfileQuery(customer_id="9101", fields=("investor_type", "assessment_expired")),
|
|||
|
|
context("9101"),
|
|||
|
|
)
|
|||
|
|
assert set(out["profile"]) == {"investor_type", "assessment_expired"}
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_requesting_a_pii_field_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
"""Agent 不能通过 `fields` 参数把 PII 要出来。"""
|
|||
|
|
wire(monkeypatch, [row()])
|
|||
|
|
with pytest.raises(ForbiddenAgentError, match="不支持查询画像字段"):
|
|||
|
|
await query_customer_profile_tool(
|
|||
|
|
CustomerProfileQuery(customer_id="9101", fields=("real_name",)), context("9101")
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_query_model_rejects_non_numeric_customer_id() -> None:
|
|||
|
|
from pydantic import ValidationError
|
|||
|
|
|
|||
|
|
with pytest.raises(ValidationError):
|
|||
|
|
CustomerProfileQuery(customer_id="abc")
|
|||
|
|
|
|||
|
|
|
|||
|
|
@pytest.mark.asyncio
|
|||
|
|
async def test_service_is_injectable_for_tests(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
"""`session_factory` 可注入(单测不连库)。"""
|
|||
|
|
wire(monkeypatch, [row(version=3)])
|
|||
|
|
view = await CustomerProfileService().load(
|
|||
|
|
CustomerProfileQuery(customer_id="9101"), context("9101"), now=NOW
|
|||
|
|
)
|
|||
|
|
assert view.version == "3"
|