新功能:个性化推荐算法
This commit is contained in:
@@ -0,0 +1,8 @@
|
||||
"""
|
||||
Content Repository(候选查询与数据访问层)。
|
||||
|
||||
说明:
|
||||
- 本模块为推荐引擎提供可注入的数据访问接口(与 ORM/SQL 解耦)。
|
||||
- 负责将 DB 存储形态规范化为上层稳定的 ContentProfile 结构。
|
||||
"""
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,36 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.features.personalized_reco.content_repository.types import ContentProfileDTO
|
||||
|
||||
|
||||
class ContentRepository(Protocol):
|
||||
"""
|
||||
推荐引擎依赖的内容数据访问抽象接口(用于解耦 ORM/SQL)。
|
||||
"""
|
||||
|
||||
async def fetch_candidates(
|
||||
self,
|
||||
*,
|
||||
scene: str,
|
||||
user_profile: object,
|
||||
fallback_level: int,
|
||||
limit: int,
|
||||
locale: str,
|
||||
exclude_content_ids: list[int] | None = None,
|
||||
) -> list[ContentProfileDTO]:
|
||||
"""
|
||||
按场景与用户画像拉取候选内容画像(用于候选池)。
|
||||
"""
|
||||
|
||||
async def fetch_contents_by_ids(
|
||||
self,
|
||||
*,
|
||||
content_ids: list[int],
|
||||
locale: str,
|
||||
) -> list[ContentProfileDTO]:
|
||||
"""
|
||||
按 content_id 批量获取内容画像(去重、按输入顺序返回;缺语言/缺记录的 id 跳过)。
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from app.features.personalized_reco.content_repository.types import Locale, normalize_locale
|
||||
|
||||
|
||||
CONTEXT_KEYS: tuple[str, ...] = ("family", "work", "relationship", "friends", "health")
|
||||
NEED_KEYS: tuple[str, ...] = (
|
||||
"emotional_support",
|
||||
"parenting_pressure",
|
||||
"self_worth",
|
||||
"anxiety_relief",
|
||||
"rest_balance",
|
||||
)
|
||||
|
||||
|
||||
def _normalize_discrete_score(v: Any, *, default: float = 0.5) -> float:
|
||||
"""
|
||||
将 suitability 的离散值规范化为 0/0.5/1。
|
||||
|
||||
非法值一律兜底 default(默认 0.5)。
|
||||
"""
|
||||
|
||||
try:
|
||||
if v in (0, 0.0):
|
||||
return 0.0
|
||||
if v in (0.5,):
|
||||
return 0.5
|
||||
if v in (1, 1.0):
|
||||
return 1.0
|
||||
# 允许字符串形式的 "0"/"0.5"/"1"
|
||||
if isinstance(v, str):
|
||||
s = v.strip()
|
||||
if s == "0":
|
||||
return 0.0
|
||||
if s == "0.5":
|
||||
return 0.5
|
||||
if s == "1":
|
||||
return 1.0
|
||||
except Exception:
|
||||
return default
|
||||
return default
|
||||
|
||||
|
||||
def normalize_suitability(raw: Any, *, keys: tuple[str, ...]) -> dict[str, float]:
|
||||
"""
|
||||
解析 suitability JSON,缺失时补齐全 0.5。
|
||||
|
||||
raw 期望为 dict;否则视为缺失。
|
||||
"""
|
||||
|
||||
data: dict[str, Any] = raw if isinstance(raw, dict) else {}
|
||||
return {k: _normalize_discrete_score(data.get(k), default=0.5) for k in keys}
|
||||
|
||||
|
||||
def normalize_review_confidence(raw: Any) -> float:
|
||||
"""
|
||||
review_confidence 缺失/NULL 时兜底 0.7。
|
||||
"""
|
||||
|
||||
try:
|
||||
if raw is None:
|
||||
return 0.7
|
||||
v = float(raw)
|
||||
if 0.0 <= v <= 1.0:
|
||||
return v
|
||||
except Exception:
|
||||
pass
|
||||
return 0.7
|
||||
|
||||
|
||||
def normalize_personalization_power(raw: Any) -> float:
|
||||
"""
|
||||
DB 约定存 0/5/10,读取层输出 0/0.5/1。
|
||||
"""
|
||||
|
||||
try:
|
||||
if raw is None:
|
||||
return 0.0
|
||||
v = int(raw)
|
||||
if v == 0:
|
||||
return 0.0
|
||||
if v == 5:
|
||||
return 0.5
|
||||
if v == 10:
|
||||
return 1.0
|
||||
except Exception:
|
||||
pass
|
||||
return 0.0
|
||||
|
||||
|
||||
_RISK_FLAG_MAP: dict[str, str] = {
|
||||
"block_stage_unknown": "unsafe_for_stage_unknown",
|
||||
"block_stage_parenting": "unsafe_for_stage_parenting",
|
||||
"block_emotion_low": "unsafe_for_emotion_low",
|
||||
"block_health_sensitive": "block_health_medical",
|
||||
}
|
||||
|
||||
|
||||
def normalize_risk_flags(raw_flags: list[str] | None) -> list[str]:
|
||||
"""
|
||||
risk_flags 旧→新映射、去重、稳定排序(字典序)。
|
||||
"""
|
||||
|
||||
flags = raw_flags or []
|
||||
mapped: set[str] = set()
|
||||
for f in flags:
|
||||
if not f:
|
||||
continue
|
||||
name = _RISK_FLAG_MAP.get(f, f)
|
||||
mapped.add(name)
|
||||
return sorted(mapped)
|
||||
|
||||
|
||||
def pick_text(*, text_en: str | None, text_tc: str | None, locale: str) -> str | None:
|
||||
"""
|
||||
按 locale 选择输出文案文本。
|
||||
|
||||
当前仅支持 EN/TC,且不允许语言回退:
|
||||
- locale=en*:必须使用 text_en
|
||||
- locale=tc/zh-TW/zh-HK:必须使用 text_tc
|
||||
"""
|
||||
|
||||
loc: Locale = normalize_locale(locale)
|
||||
if loc == "en":
|
||||
return text_en if text_en else None
|
||||
# loc == "tc"
|
||||
return text_tc if text_tc else None
|
||||
|
||||
@@ -0,0 +1,265 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Iterable
|
||||
|
||||
from sqlalchemy import Select, and_, desc, not_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
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.interface import ContentRepository
|
||||
from app.features.personalized_reco.content_repository.normalization import (
|
||||
CONTEXT_KEYS,
|
||||
NEED_KEYS,
|
||||
normalize_personalization_power,
|
||||
normalize_review_confidence,
|
||||
normalize_risk_flags,
|
||||
normalize_suitability,
|
||||
pick_text,
|
||||
)
|
||||
from app.features.personalized_reco.content_repository.types import ContentProfileDTO
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _UserSignals:
|
||||
"""
|
||||
从 user_profile 中提取 repository 级别需要的最小信号。
|
||||
|
||||
注意:更复杂的规则(Hard Filter/Scoring/Rerank)不在本层处理。
|
||||
"""
|
||||
|
||||
missing_need: bool
|
||||
missing_context: bool
|
||||
missing_emotion: bool
|
||||
stage: str | None # expecting/parenting/unknown/general/None
|
||||
|
||||
|
||||
def _bool(v: Any) -> bool:
|
||||
return bool(v)
|
||||
|
||||
|
||||
def _extract_user_signals(user_profile: object) -> _UserSignals:
|
||||
"""
|
||||
兼容 pydantic model / dict / 其他对象的最小字段读取。
|
||||
"""
|
||||
|
||||
def _get(obj: Any, key: str, default: Any = None) -> Any:
|
||||
if obj is None:
|
||||
return default
|
||||
if isinstance(obj, dict):
|
||||
return obj.get(key, default)
|
||||
return getattr(obj, key, default)
|
||||
|
||||
need = _get(user_profile, "need", {}) or {}
|
||||
context = _get(user_profile, "context", {}) or {}
|
||||
emotion_score = _get(user_profile, "emotion_score", None)
|
||||
|
||||
missing_need = len(need) == 0
|
||||
missing_context = len(context) == 0
|
||||
missing_emotion = emotion_score is None
|
||||
|
||||
# stage: from user_profile.stage (one-hot)
|
||||
stage_obj = _get(user_profile, "stage", None)
|
||||
stage: str | None = None
|
||||
if stage_obj is not None:
|
||||
expecting = _get(stage_obj, "expecting", None)
|
||||
parenting = _get(stage_obj, "parenting", None)
|
||||
unknown = _get(stage_obj, "unknown", None)
|
||||
if _bool(expecting):
|
||||
stage = "expecting"
|
||||
elif _bool(parenting):
|
||||
stage = "parenting"
|
||||
elif _bool(unknown):
|
||||
stage = "unknown"
|
||||
|
||||
return _UserSignals(
|
||||
missing_need=missing_need,
|
||||
missing_context=missing_context,
|
||||
missing_emotion=missing_emotion,
|
||||
stage=stage,
|
||||
)
|
||||
|
||||
|
||||
def _dedupe_preserve_order(ids: Iterable[int]) -> list[int]:
|
||||
seen: set[int] = set()
|
||||
out: list[int] = []
|
||||
for i in ids:
|
||||
if i in seen:
|
||||
continue
|
||||
seen.add(i)
|
||||
out.append(i)
|
||||
return out
|
||||
|
||||
|
||||
class SqlAlchemyContentRepository(ContentRepository):
|
||||
"""
|
||||
基于 SQLAlchemy AsyncSession 的 ContentRepository 实现。
|
||||
"""
|
||||
|
||||
def __init__(self, session: AsyncSession):
|
||||
self._session = session
|
||||
|
||||
async def fetch_contents_by_ids(self, *, content_ids: list[int], locale: str) -> list[ContentProfileDTO]:
|
||||
"""
|
||||
- 输入去重
|
||||
- 输出顺序与输入一致(按首次出现顺序)
|
||||
- 缺记录或缺目标语言文本:跳过
|
||||
- 不产生 N+1(主体+画像一次,flags 一次)
|
||||
"""
|
||||
|
||||
unique_ids = _dedupe_preserve_order(content_ids)
|
||||
if not unique_ids:
|
||||
return []
|
||||
|
||||
# locale 文本存在性过滤(不允许语言回退)
|
||||
# en -> 必须 text_en;tc -> 必须 text_tc
|
||||
# 过滤在 DB 层做,避免后续组装无意义
|
||||
from app.features.personalized_reco.content_repository.types import normalize_locale
|
||||
|
||||
loc = normalize_locale(locale)
|
||||
text_filter = Content.text_en.is_not(None) if loc == "en" else Content.text_tc.is_not(None)
|
||||
|
||||
stmt: Select = (
|
||||
select(Content, ContentProfile)
|
||||
.join(ContentProfile, Content.content_id == ContentProfile.content_id)
|
||||
.where(and_(Content.content_id.in_(unique_ids), text_filter))
|
||||
)
|
||||
|
||||
rows = (await self._session.execute(stmt)).all()
|
||||
if not rows:
|
||||
return []
|
||||
|
||||
# 先组装主体+画像,后续再补 risk_flags
|
||||
by_id: dict[int, dict[str, Any]] = {}
|
||||
valid_ids: list[int] = []
|
||||
for content, profile in rows:
|
||||
cid = int(content.content_id)
|
||||
text = pick_text(text_en=content.text_en, text_tc=content.text_tc, locale=locale)
|
||||
if not text:
|
||||
continue
|
||||
by_id[cid] = {
|
||||
"content": content,
|
||||
"profile": profile,
|
||||
"text": text,
|
||||
}
|
||||
valid_ids.append(cid)
|
||||
|
||||
if not by_id:
|
||||
return []
|
||||
|
||||
# 批量取 flags(避免 join 行膨胀)
|
||||
flags_stmt = select(ContentRiskFlag.content_id, ContentRiskFlag.flag).where(
|
||||
ContentRiskFlag.content_id.in_(list(by_id.keys()))
|
||||
)
|
||||
flags_rows = (await self._session.execute(flags_stmt)).all()
|
||||
flags_map: dict[int, list[str]] = defaultdict(list)
|
||||
for cid, flag in flags_rows:
|
||||
flags_map[int(cid)].append(str(flag))
|
||||
|
||||
result_by_id: dict[int, ContentProfileDTO] = {}
|
||||
for cid, payload in by_id.items():
|
||||
content: Content = payload["content"]
|
||||
profile: ContentProfile = payload["profile"]
|
||||
text: str = payload["text"]
|
||||
|
||||
dto = ContentProfileDTO(
|
||||
content_id=cid,
|
||||
text=text,
|
||||
stage=profile.stage, # type: ignore[arg-type]
|
||||
emotion_score=float(profile.emotion_score) if profile.emotion_score is not None else None,
|
||||
context_suitability=normalize_suitability(profile.context_suitability_json, keys=CONTEXT_KEYS),
|
||||
need_suitability=normalize_suitability(profile.need_suitability_json, keys=NEED_KEYS),
|
||||
personalization_power=normalize_personalization_power(profile.personalization_power),
|
||||
risk_flags=normalize_risk_flags(flags_map.get(cid)),
|
||||
author_id=content.author_id,
|
||||
template_id=content.template_id,
|
||||
review_confidence=normalize_review_confidence(profile.review_confidence),
|
||||
)
|
||||
result_by_id[cid] = dto
|
||||
|
||||
# 按输入顺序返回(跳过缺失/被过滤的)
|
||||
out: list[ContentProfileDTO] = []
|
||||
for cid in unique_ids:
|
||||
dto = result_by_id.get(cid)
|
||||
if dto is not None:
|
||||
out.append(dto)
|
||||
return out
|
||||
|
||||
async def fetch_candidates(
|
||||
self,
|
||||
*,
|
||||
scene: str,
|
||||
user_profile: object,
|
||||
fallback_level: int,
|
||||
limit: int,
|
||||
locale: str,
|
||||
exclude_content_ids: list[int] | None = None,
|
||||
) -> list[ContentProfileDTO]:
|
||||
"""
|
||||
两段式候选召回:
|
||||
1) 先查候选 content_id 列表(含粗过滤、locale 过滤、limit*multiplier)
|
||||
2) 再批量补全字段(复用 fetch_contents_by_ids)
|
||||
"""
|
||||
|
||||
if limit <= 0:
|
||||
return []
|
||||
|
||||
signals = _extract_user_signals(user_profile)
|
||||
effective_fallback = int(fallback_level)
|
||||
if signals.missing_need or signals.missing_context or signals.missing_emotion:
|
||||
effective_fallback = max(effective_fallback, 1)
|
||||
|
||||
# locale 文本存在性过滤(不允许语言回退)
|
||||
from app.features.personalized_reco.content_repository.types import normalize_locale
|
||||
|
||||
loc = normalize_locale(locale)
|
||||
text_filter = Content.text_en.is_not(None) if loc == "en" else Content.text_tc.is_not(None)
|
||||
|
||||
filters: list[Any] = [text_filter]
|
||||
if exclude_content_ids:
|
||||
filters.append(not_(Content.content_id.in_(exclude_content_ids)))
|
||||
|
||||
# fallback 约束(repository 只做“降级约束”,不做 hard filter)
|
||||
if effective_fallback >= 1:
|
||||
# personalization_power <= 5 代表 <= 0.5
|
||||
filters.append(ContentProfile.personalization_power <= 5)
|
||||
if effective_fallback >= 2:
|
||||
filters.append(ContentProfile.personalization_power == 0)
|
||||
filters.append(ContentProfile.stage == "general")
|
||||
if effective_fallback >= 3:
|
||||
filters.append(ContentProfile.is_safe_pool.is_(True))
|
||||
filters.append(ContentProfile.personalization_power == 0)
|
||||
filters.append(ContentProfile.stage == "general")
|
||||
|
||||
# stage 粗过滤(仅 L0/L1 才做“用户阶段 + general”;L2/L3 已强制 general)
|
||||
if effective_fallback < 2:
|
||||
user_stage = signals.stage
|
||||
if user_stage in {"expecting", "parenting"}:
|
||||
filters.append(ContentProfile.stage.in_([user_stage, "general"]))
|
||||
else:
|
||||
# unknown 或无法判定:仅取 general,避免误推
|
||||
filters.append(ContentProfile.stage == "general")
|
||||
|
||||
multiplier = 5
|
||||
raw_limit = max(limit * multiplier, limit)
|
||||
|
||||
stmt_ids = (
|
||||
select(Content.content_id)
|
||||
.join(ContentProfile, Content.content_id == ContentProfile.content_id)
|
||||
.where(and_(*filters))
|
||||
.order_by(desc(ContentProfile.updated_at))
|
||||
.limit(raw_limit)
|
||||
)
|
||||
|
||||
candidate_ids_rows = (await self._session.execute(stmt_ids)).scalars().all()
|
||||
candidate_ids = [int(x) for x in candidate_ids_rows]
|
||||
if not candidate_ids:
|
||||
return []
|
||||
|
||||
# 复用按 ID 批量补全(会再次做 locale 过滤,但成本可接受,且可保证一致行为)
|
||||
items = await self.fetch_contents_by_ids(content_ids=candidate_ids, locale=locale)
|
||||
return items[:limit]
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
# 当前阶段仅支持 EN / TC(繁体中文)
|
||||
Locale = Literal["en", "tc"]
|
||||
|
||||
|
||||
def normalize_locale(locale: str) -> Locale:
|
||||
"""
|
||||
将客户端传入的 locale 归一化为内部枚举(仅 EN / TC)。
|
||||
|
||||
约定:
|
||||
- 任何以 "en" 开头的 locale 归一化为 "en"(例如 en、en-US)
|
||||
- "tc"/"zh-TW"/"zh-HK" 归一化为 "tc"
|
||||
- 其他 locale 视为不支持
|
||||
"""
|
||||
|
||||
raw = (locale or "").strip()
|
||||
if not raw:
|
||||
raise ValueError("locale 不能为空(当前仅支持 en/tc)")
|
||||
|
||||
low = raw.lower()
|
||||
if low.startswith("en"):
|
||||
return "en"
|
||||
if low in {"tc", "zh-tw", "zh-hk", "zh_tw", "zh_hk"}:
|
||||
return "tc"
|
||||
|
||||
raise ValueError(f"不支持的 locale:{locale!r}(当前仅支持 en/tc)")
|
||||
|
||||
|
||||
ContentStage = Literal["general", "expecting", "parenting", "unknown"]
|
||||
|
||||
|
||||
class ContentProfileDTO(BaseModel):
|
||||
"""
|
||||
推荐模块消费的内容画像(稳定字段契约)。
|
||||
|
||||
注意:
|
||||
- text 已按 locale 选择,不允许语言回退(缺语言文本的内容不返回)
|
||||
- emotion_score 为 None 表示 general
|
||||
- personalization_power 对上统一为 0/0.5/1
|
||||
- review_confidence 缺失时兜底 0.7
|
||||
"""
|
||||
|
||||
content_id: int
|
||||
text: str
|
||||
stage: ContentStage
|
||||
emotion_score: Optional[float] = None
|
||||
|
||||
context_suitability: dict[str, float] = Field(default_factory=dict)
|
||||
need_suitability: dict[str, float] = Field(default_factory=dict)
|
||||
|
||||
personalization_power: float
|
||||
risk_flags: list[str] = Field(default_factory=list)
|
||||
|
||||
# 可选字段
|
||||
author_id: Optional[str] = None
|
||||
template_id: Optional[str] = None
|
||||
review_confidence: float = 0.7
|
||||
|
||||
Reference in New Issue
Block a user