Files
mindfulness/server/tests/test_content_repository.py
2026-02-02 16:47:37 +08:00

179 lines
6.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
from datetime import datetime
import pytest
from app.db.models.content import Content
from app.db.models.content_profile import ContentProfile
from app.db.models.content_risk_flag import ContentRiskFlag
from app.features.personalized_reco.content_repository.sqlalchemy_repo import SqlAlchemyContentRepository
def _ctx_json() -> dict:
return {"family": 0.5, "work": 0.5, "relationship": 0.5, "friends": 0.5, "health": 0.5}
def _need_json() -> dict:
return {
"emotional_support": 0.5,
"parenting_pressure": 0.5,
"self_worth": 0.5,
"anxiety_relief": 0.5,
"rest_balance": 0.5,
}
@pytest.mark.asyncio
async def test_fetch_contents_by_ids_locale_no_fallback(db_session, query_counter):
# content 1: 只有英文
c1 = Content(text_en="hello", text_tc=None, author_id="a1", template_id="t1")
db_session.add(c1)
await db_session.flush()
db_session.add(
ContentProfile(
content_id=c1.content_id,
stage="general",
emotion_score=None,
context_suitability_json=_ctx_json(),
need_suitability_json=_need_json(),
personalization_power=10,
review_confidence=None,
is_safe_pool=False,
)
)
db_session.add(ContentRiskFlag(content_id=c1.content_id, flag="block_stage_unknown"))
# content 2: 只有繁中
c2 = Content(text_en=None, text_tc="繁體中文", author_id="a2", template_id="t2")
db_session.add(c2)
await db_session.flush()
db_session.add(
ContentProfile(
content_id=c2.content_id,
stage="general",
emotion_score=None,
context_suitability_json=_ctx_json(),
need_suitability_json=_need_json(),
personalization_power=0,
review_confidence=0.9,
is_safe_pool=True,
)
)
db_session.add(ContentRiskFlag(content_id=c2.content_id, flag="block_health_sensitive"))
await db_session.flush()
repo = SqlAlchemyContentRepository(db_session)
start = query_counter()
# en只能拿到有 text_en 的内容(不允许回退到 text_tc
en_items = await repo.fetch_contents_by_ids(content_ids=[c2.content_id, c1.content_id], locale="en")
assert [x.content_id for x in en_items] == [c1.content_id]
assert en_items[0].text == "hello"
assert en_items[0].personalization_power == 1.0
assert en_items[0].review_confidence == 0.7 # NULL -> 0.7
assert "unsafe_for_stage_unknown" in en_items[0].risk_flags
assert "block_stage_unknown" not in en_items[0].risk_flags
# tc只能拿到有 text_tc 的内容(不允许回退到 text_en
tc_items = await repo.fetch_contents_by_ids(content_ids=[c1.content_id, c2.content_id], locale="tc")
assert [x.content_id for x in tc_items] == [c2.content_id]
assert tc_items[0].text == "繁體中文"
assert tc_items[0].review_confidence == 0.9
assert "block_health_medical" in tc_items[0].risk_flags
assert "block_health_sensitive" not in tc_items[0].risk_flags
# 两次调用各自 2 次查询(主体+画像一次flags 一次),总计应为常数级
end = query_counter()
assert (end - start) <= 4
@pytest.mark.asyncio
async def test_fetch_candidates_fallback_and_locale_filter(db_session, query_counter):
# 构造 4 条英文内容power 0/5/10安全池标记不同
contents = []
for i, (power, safe) in enumerate([(0, True), (5, False), (10, False), (0, False)], start=1):
c = Content(text_en=f"en_{i}", text_tc=None, author_id=f"a{i}", template_id=f"t{i}")
db_session.add(c)
await db_session.flush()
db_session.add(
ContentProfile(
content_id=c.content_id,
stage="general",
emotion_score=None,
context_suitability_json=_ctx_json(),
need_suitability_json=_need_json(),
personalization_power=power,
review_confidence=None,
is_safe_pool=safe,
# 让测试数据在候选排序中排到最前,避免依赖“库为空”
updated_at=datetime(2099, 1, 1, 0, 0, 0),
)
)
contents.append(c)
await db_session.flush()
inserted_ids = {int(c.content_id) for c in contents}
repo = SqlAlchemyContentRepository(db_session)
class _MinimalUser:
# need/context/emotion_score 全缺失 -> effective_fallback 至少 L1
need = {}
context = {}
emotion_score = None
stage = {"unknown": 1}
start = query_counter()
# 入参 L0但因为缺失字段effective_fallback=1 -> power<=5排除 power=10
items_l0 = await repo.fetch_candidates(
scene="feed",
user_profile=_MinimalUser(),
fallback_level=0,
limit=3,
locale="en-US",
)
powers = [x.personalization_power for x in items_l0]
assert 1.0 not in powers
assert all(x.text.startswith("en_") for x in items_l0)
# L2强制 power=0 且 stage=general这里都 general只剩 power=0 的两条
items_l2 = await repo.fetch_candidates(
scene="feed",
user_profile=_MinimalUser(),
fallback_level=2,
limit=2,
locale="en",
)
assert all(x.personalization_power == 0.0 for x in items_l2)
assert all(x.text.startswith("en_") for x in items_l2)
# L3只安全池is_safe_pool=true且 power=0
items_l3 = await repo.fetch_candidates(
scene="feed",
user_profile=_MinimalUser(),
fallback_level=3,
limit=1,
locale="en",
)
assert len(items_l3) == 1
assert items_l3[0].text.startswith("en_")
assert items_l3[0].personalization_power == 0.0
# locale 过滤tc 请求下这些内容都没有 text_tc -> 返回空
items_tc = await repo.fetch_candidates(
scene="feed",
user_profile=_MinimalUser(),
fallback_level=0,
limit=10,
locale="tc",
)
# 不要求库为空:只断言“不会把本次插入的 en-only 测试数据返回出来”
assert not any(x.content_id in inserted_ids for x in items_tc)
end = query_counter()
# 期望为常数级(每次 fetch_candidates1 次取 ids + 2 次补全),这里 4 次调用 -> <= 12
assert (end - start) <= 12