diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 01dc5f4..0737f0b 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -39,16 +39,13 @@ jobs: run: uv python install ${{ env.PYTHON_VERSION }} - name: Install dependencies - run: uv sync --frozen --no-dev - - - name: Install dev dependencies - run: uv sync --frozen + run: uv sync --frozen --extra dev - name: Run ruff check - run: uv run ruff check src/ tests/ + run: uv run --frozen --extra dev ruff check src/ tests/ - name: Run ruff format check - run: uv run ruff format --check src/ tests/ + run: uv run --frozen --extra dev ruff format --check src/ tests/ # ===== 2. Test(单元测试 + 集成测试)===== test: @@ -84,19 +81,19 @@ jobs: run: uv python install ${{ env.PYTHON_VERSION }} - name: Install dependencies - run: uv sync --frozen + run: uv sync --frozen --extra dev - name: Generate demo data run: | # 创建 .env 用于 generate_data.py echo "OPENAI_API_KEY=ci-test-key" > .env echo "OPENAI_BASE_URL=https://api.siliconflow.cn/v1" >> .env - uv run python generate_data.py + uv run --frozen --extra dev python generate_data.py continue-on-error: true - name: Run tests run: | - uv run pytest tests/ \ + uv run --frozen --extra dev pytest tests/ \ -v \ --tb=short \ --timeout=60 \ diff --git a/Dockerfile b/Dockerfile index f8c9f7f..2b95527 100644 --- a/Dockerfile +++ b/Dockerfile @@ -54,8 +54,8 @@ EXPOSE 8000 # 健康检查 HEALTHCHECK --interval=30s --timeout=5s --start-period=30s --retries=3 \ - CMD curl -f http://localhost:8000/api/health || exit 1 + CMD [".venv/bin/python", "-c", "from urllib.request import urlopen; urlopen('http://localhost:8000/api/health', timeout=5)"] # 启动命令 -CMD [".venv/bin/uv", "run", "uvicorn", "src.api.main:app", \ +CMD [".venv/bin/python", "-m", "uvicorn", "src.api.main:app", \ "--host", "0.0.0.0", "--port", "8000", "--log-level", "info"] diff --git a/pyproject.toml b/pyproject.toml index be7abf9..b9a7691 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -47,6 +47,7 @@ dev = [ "pytest>=9.1.0", "pytest-asyncio>=0.24.0", "pytest-cov>=6.0.0", + "pytest-timeout>=2.3.1", "httpx>=0.27.0", "sentence-transformers>=3.0", "mcp>=1.27.2", diff --git a/src/agents/comment_agent.py b/src/agents/comment_agent.py index 29432da..f71aad1 100644 --- a/src/agents/comment_agent.py +++ b/src/agents/comment_agent.py @@ -8,20 +8,21 @@ 2. 解析 YAML 格式的评论分析数据 3. 输出结构化的分析结果 """ + import re -import yaml -from typing import List, Optional from pathlib import Path + +import yaml from langchain_core.documents import Document class CommentAnalyzer: """分析笔记评论区,提取用户真实需求信号""" - def __init__(self, raw_dir: Optional[str] = None): + def __init__(self, raw_dir: str | None = None): self.raw_dir = raw_dir - def analyze(self, documents: List[Document]) -> List[dict]: + def analyze(self, documents: list[Document]) -> list[dict]: """ 从检索到的文档中提取评论分析数据。 @@ -98,7 +99,7 @@ def analyze(self, documents: List[Document]) -> List[dict]: # ---- 内部方法 ---- - def _resolve_source(self, doc: Document) -> Optional[Path]: + def _resolve_source(self, doc: Document) -> Path | None: """从 Document metadata 中解析原始文件路径""" source = doc.metadata.get("source", "") if not source: diff --git a/src/agents/creator_agent.py b/src/agents/creator_agent.py index cfdfba2..d8e99df 100644 --- a/src/agents/creator_agent.py +++ b/src/agents/creator_agent.py @@ -4,9 +4,9 @@ 与 InsightGenerator 并行:同一份 DemandAggregator 输出,不同的 prompt 模板。 把用户评论数据变成选题 + 脚本大纲 + 封面方案。 """ -from langchain_core.messages import HumanMessage -from langchain_core.prompts import ChatPromptTemplate + from langchain_openai import ChatOpenAI + from src.config import LLM_CONFIG from src.core.prompt_loader import get_prompt_loader @@ -22,19 +22,25 @@ def _get_prompt(self): return self.prompt_loader.load("creator_report", "v1") def _build_msg(self, aggregated: dict, category: str = "") -> list: - complaints_str = "\n".join( - f" {i+1}. 「{c}」出现 {f} 次" - for i, (c, f) in enumerate(aggregated["top_complaints"][:10]) - ) or " 暂无" + complaints_str = ( + "\n".join( + f" {i + 1}. 「{c}」出现 {f} 次" + for i, (c, f) in enumerate(aggregated["top_complaints"][:10]) + ) + or " 暂无" + ) - intents_str = "\n".join( - f" {i+1}. 「{t}」出现 {f} 次" - for i, (t, f) in enumerate(aggregated["top_purchase_intents"][:10]) - ) or " 暂无" + intents_str = ( + "\n".join( + f" {i + 1}. 「{t}」出现 {f} 次" + for i, (t, f) in enumerate(aggregated["top_purchase_intents"][:10]) + ) + or " 暂无" + ) - comparisons_str = "\n".join( - f" - {c}" for c in aggregated["comparison_patterns"][:10] - ) or " 暂无" + comparisons_str = ( + "\n".join(f" - {c}" for c in aggregated["comparison_patterns"][:10]) or " 暂无" + ) brands_str = ", ".join(aggregated["related_brands"]) or "暂无" differentiations_str = ", ".join(aggregated.get("differentiation_directions", [])) or "暂无" @@ -88,7 +94,11 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") lines.append("") lines.append("【数据亮点】") - lines.append(f" 📊 {aggregated['note_count']}篇笔记 → {len(aggregated['top_complaints'])}个痛点 + {len(aggregated['top_purchase_intents'])}个需求信号") + lines.append( + f" 📊 {aggregated['note_count']}篇笔记 → " + f"{len(aggregated['top_complaints'])}个痛点 + " + f"{len(aggregated['top_purchase_intents'])}个需求信号" + ) lines.append("") # 提取数据 @@ -96,11 +106,10 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: intents = aggregated["top_purchase_intents"] brands = aggregated.get("related_brands", []) avg_price = aggregated.get("avg_price", 0) - avg_cost = aggregated.get("avg_cost", 0) # 方案1: 避坑/测评向 lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") - lines.append(f"【方案一】🔍 避坑测评") + lines.append("【方案一】🔍 避坑测评") lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") if pains: pain = pains[0][0] @@ -110,7 +119,10 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(f" 标题:《{category}选购避坑指南,买前必看》") lines.append("") lines.append(" 【脚本大纲】") - lines.append(f" 前5秒:展示热门{category}产品,抛出问题「{pains[0][0] if pains else '买错等于浪费钱'}」") + lines.append( + f" 前5秒:展示热门{category}产品," + f"抛出问题「{pains[0][0] if pains else '买错等于浪费钱'}」" + ) lines.append(" 5-15秒:引用真实评论引发共鸣") if len(pains) >= 2: lines.append(f" 「{pains[0][0]}」「{pains[1][0]}」") @@ -121,7 +133,7 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: # 方案2: 推荐/种草向 lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") - lines.append(f"【方案二】🌟 好物种草") + lines.append("【方案二】🌟 好物种草") lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") if intents: intent = intents[0][0] @@ -133,18 +145,23 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(" 【脚本大纲】") lines.append(f" 前5秒:直接展示{category}使用效果「这也太好用了吧」") if intents: - lines.append(f" 5-15秒:抛出用户最大需求「{intents[0][0] if intents else '想要高性价比'}」") - lines.append(f" 核心段:第1件→第2件→第3件,每件15秒展示+口播") + lines.append( + f" 5-15秒:抛出用户最大需求「{intents[0][0] if intents else '想要高性价比'}」" + ) + lines.append(" 核心段:第1件→第2件→第3件,每件15秒展示+口播") lines.append(" 结尾:「评论区告诉我你最想试哪款」") lines.append(" 互动引导:收藏+关注,下期继续挖宝") lines.append("") # 方案3: 对比/横评向 lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") - lines.append(f"【方案三】⚔️ 品牌横评") + lines.append("【方案三】⚔️ 品牌横评") lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") if len(brands) >= 2: - lines.append(f" 标题:《{brands[0]} vs {brands[1]} vs {brands[2] if len(brands)>=3 else brands[1]},{category}谁更强?》") + lines.append( + f" 标题:《{brands[0]} vs {brands[1]} vs " + f"{brands[2] if len(brands) >= 3 else brands[1]},{category}谁更强?》" + ) else: lines.append(f" 标题:《{category}热门品牌横评,到底选哪个?》") lines.append(" 类型:对比评测 | 平台:B站+小红书 | 预计互动:⭐⭐⭐⭐") @@ -159,15 +176,17 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: # 发布建议 lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") - lines.append(f"【发布策略】(适用于3个方案)") + lines.append("【发布策略】(适用于3个方案)") lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") lines.append(" 🕐 推荐发布:工作日 19:00-21:00 或周末 10:00-12:00") - lines.append(f" 🏷️ 核心标签:") - lines.append(f" 方案一:#避坑 #真实测评 #购物踩雷") - lines.append(f" 方案二:#好物推荐 #种草 #年度爱用") - lines.append(f" 方案三:#对比评测 #理性消费 #品牌测评") + lines.append(" 🏷️ 核心标签:") + lines.append(" 方案一:#避坑 #真实测评 #购物踩雷") + lines.append(" 方案二:#好物推荐 #种草 #年度爱用") + lines.append(" 方案三:#对比评测 #理性消费 #品牌测评") if avg_price > 0: - lines.append(f" 💰 客单价 ¥{avg_price},适合新品牌用方案一打口碑,老品牌用方案二冲销量") + lines.append( + f" 💰 客单价 ¥{avg_price},适合新品牌用方案一打口碑,老品牌用方案二冲销量" + ) lines.append("") return "\n".join(lines) diff --git a/src/agents/demand_agent.py b/src/agents/demand_agent.py index 7a18f2a..992aaf4 100644 --- a/src/agents/demand_agent.py +++ b/src/agents/demand_agent.py @@ -3,14 +3,14 @@ 将 CommentAnalyzer 输出的多条评论分析结果进行聚合统计, 找出高频投诉、热门需求、品牌竞争格局等可行动的信号。 """ + from collections import Counter -from typing import List, Dict, Any class DemandAggregator: """聚合多个文档的评论数据,提取高价值需求信号""" - def aggregate(self, analyses: List[dict]) -> dict: + def aggregate(self, analyses: list[dict]) -> dict: """ 输入:CommentAnalyzer 输出的分析结果列表 输出:聚合后的需求洞察 + 电商选品评分 @@ -56,7 +56,6 @@ def aggregate(self, analyses: List[dict]) -> dict: total_cost = 0 total_weight = 0.0 profit_margins = [] - logistics_scores = [] competition_levels = [] entry_difficulties = [] differentiation_ops = [] @@ -119,37 +118,46 @@ def aggregate(self, analyses: List[dict]) -> dict: price_cost_ratio = avg_price / avg_cost else: price_cost_ratio = 3.0 - profit_score = min(100, int( - min(price_cost_ratio / 5, 1.0) * 40 + # 定价倍率 满分40 - min(avg_margin / 0.7, 1.0) * 40 + # 利润率 满分40 - 20 # 基础分 - )) + profit_score = min( + 100, + int( + min(price_cost_ratio / 5, 1.0) * 40 # 定价倍率 满分40 + + min(avg_margin / 0.7, 1.0) * 40 # 利润率 满分40 + + 20 # 基础分 + ), + ) # 2️⃣ 物流友好度评分 - logistics_score = min(100, int( - (avg_weight < 0.3) * 30 + # 轻量 - (avg_weight < 1.0) * 15 + # 不太重 - 35 + # 基础分 - (competition_levels.count("高") < n * 0.5) * 10 # 偏低竞争加分 - )) + logistics_score = min( + 100, + int( + (avg_weight < 0.3) * 30 # 轻量 + + (avg_weight < 1.0) * 15 # 不太重 + + 35 # 基础分 + + (competition_levels.count("高") < n * 0.5) * 10 # 偏低竞争加分 + ), + ) # 3️⃣ 竞争强度评分(分越高越推荐进入) low_comp_ratio = competition_levels.count("低") / max(n, 1) easy_entry_ratio = entry_difficulties.count("低") / max(n, 1) brand_count = len(all_brands) - competition_score = min(100, int( - low_comp_ratio * 40 + # 低竞争占比 - easy_entry_ratio * 30 + # 低进入门槛 - min((5 - brand_count) * 5, 20) + # 品牌少=机会大 - 10 # 基础分 - )) + competition_score = min( + 100, + int( + low_comp_ratio * 40 # 低竞争占比 + + easy_entry_ratio * 30 # 低进入门槛 + + min((5 - brand_count) * 5, 20) # 品牌少=机会大 + + 10 # 基础分 + ), + ) # 4️⃣ 综合选品评分 (加权) selection_score = int( - profit_score * 0.30 + # 利润权重 30% - logistics_score * 0.20 + # 物流权重 20% - competition_score * 0.20 + # 竞争权重 20% - demand_score * 0.30 # 需求权重 30% + profit_score * 0.30 # 利润权重 30% + + logistics_score * 0.20 # 物流权重 20% + + competition_score * 0.20 # 竞争权重 20% + + demand_score * 0.30 # 需求权重 30% ) # 5️⃣ 差异化方向(去重取前5) @@ -167,7 +175,6 @@ def aggregate(self, analyses: List[dict]) -> dict: "total_ask_link": total_ask, "avg_likes": round(avg_likes, 1), "demand_score": demand_score, - # 电商选品评分 "avg_price": round(avg_price, 1), "avg_cost": round(avg_cost, 1), @@ -194,7 +201,6 @@ def _empty_result(self) -> dict: "total_ask_link": 0, "avg_likes": 0, "demand_score": 0, - "avg_price": 0, "avg_cost": 0, "avg_profit_margin": 0, diff --git a/src/agents/insight_agent.py b/src/agents/insight_agent.py index 990bebc..e513d11 100644 --- a/src/agents/insight_agent.py +++ b/src/agents/insight_agent.py @@ -3,13 +3,13 @@ 基于 DemandAggregator 的聚合结果,用 LLM 生成可执行的市场洞察报告。 是 Phase 3 的最终输出环节。 """ + from langchain_core.messages import HumanMessage -from langchain_core.prompts import ChatPromptTemplate from langchain_openai import ChatOpenAI + from src.config import LLM_CONFIG from src.core.prompt_loader import get_prompt_loader - THREE_TIER_HINT = HumanMessage( content="⚠️ 重要:在【选品综合评分】之后,必须输出【三档价位选品】章节!" "按低价/中价/高价三档展开,每档包含:价格带、产品方向、功能亮点、目标人群、预估利润。" @@ -30,19 +30,25 @@ def _get_prompt(self): def _build_msg(self, aggregated: dict, category: str = "") -> list: """构建消息列表(含三档价位强制指令)""" - complaints_str = "\n".join( - f" {i+1}. 「{c}」出现 {f} 次" - for i, (c, f) in enumerate(aggregated["top_complaints"][:10]) - ) or " 暂无" + complaints_str = ( + "\n".join( + f" {i + 1}. 「{c}」出现 {f} 次" + for i, (c, f) in enumerate(aggregated["top_complaints"][:10]) + ) + or " 暂无" + ) - intents_str = "\n".join( - f" {i+1}. 「{t}」出现 {f} 次" - for i, (t, f) in enumerate(aggregated["top_purchase_intents"][:10]) - ) or " 暂无" + intents_str = ( + "\n".join( + f" {i + 1}. 「{t}」出现 {f} 次" + for i, (t, f) in enumerate(aggregated["top_purchase_intents"][:10]) + ) + or " 暂无" + ) - comparisons_str = "\n".join( - f" - {c}" for c in aggregated["comparison_patterns"][:10] - ) or " 暂无" + comparisons_str = ( + "\n".join(f" - {c}" for c in aggregated["comparison_patterns"][:10]) or " 暂无" + ) brands_str = ", ".join(aggregated["related_brands"]) or "暂无" differentiations_str = ", ".join(aggregated.get("differentiation_directions", [])) or "暂无" @@ -121,7 +127,9 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append("【市场概况】") lines.append(f"品类:{category or '未分类'}") lines.append(f"分析笔记数:{aggregated['note_count']} 篇") - lines.append(f"平均点赞:{aggregated['avg_likes']} | 总求链接:{aggregated['total_ask_link']}") + lines.append( + f"平均点赞:{aggregated['avg_likes']} | 总求链接:{aggregated['total_ask_link']}" + ) evergreen = aggregated.get("evergreen_ratio", 0.8) lines.append(f"季节特性:{'✅ 常青款为主' if evergreen > 0.5 else '⚠️ 偏季节性'}") lines.append("") @@ -135,7 +143,7 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: if avg_price > 0: lines.append(f"平均售价:¥{avg_price} | 平均成本:¥{avg_cost}") lines.append(f"定价倍率:{ratio}x {'✅ 达标(≥3x)' if ratio >= 3 else '⚠️ 偏低(<3x)'}") - lines.append(f"预估利润率:{margin*100:.0f}%") + lines.append(f"预估利润率:{margin * 100:.0f}%") lines.append(f"利润评分:{profit_score}/100") else: lines.append("(暂无售价数据,建议参考1688/拼多多比价)") @@ -145,7 +153,9 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: avg_weight = aggregated.get("avg_weight", 0) logistics_score = aggregated.get("logistics_score", 0) if avg_weight > 0: - weight_level = "轻量级" if avg_weight < 0.3 else ("中量级" if avg_weight < 1.0 else "重量级") + weight_level = ( + "轻量级" if avg_weight < 0.3 else ("中量级" if avg_weight < 1.0 else "重量级") + ) lines.append(f"平均重量:{avg_weight}kg({weight_level})") lines.append(f"物流评分:{logistics_score}/100") if logistics_score >= 70: @@ -162,7 +172,10 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: comp_score = aggregated.get("competition_score", 50) if aggregated["related_brands"]: lines.append(f"涉及品牌:{', '.join(aggregated['related_brands'])}") - lines.append(f"竞争评分:{comp_score}/100({'✅ 蓝海' if comp_score >= 60 else '⚠️ 中等' if comp_score >= 35 else '❌ 红海'})") + competition_label = ( + "✅ 蓝海" if comp_score >= 60 else "⚠️ 中等" if comp_score >= 35 else "❌ 红海" + ) + lines.append(f"竞争评分:{comp_score}/100({competition_label})") if aggregated["comparison_patterns"]: lines.append("用户对比:") for c in aggregated["comparison_patterns"][:5]: @@ -172,7 +185,7 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(f"【用户痛点 TOP {min(len(aggregated['top_complaints']), 5)}】") if aggregated["top_complaints"]: for i, (c, f) in enumerate(aggregated["top_complaints"][:5]): - lines.append(f" {i+1}. {c}(出现 {f} 次)") + lines.append(f" {i + 1}. {c}(出现 {f} 次)") else: lines.append(" 暂无明显投诉") lines.append("") @@ -180,7 +193,7 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append("【需求信号】") if aggregated["top_purchase_intents"]: for i, (t, f) in enumerate(aggregated["top_purchase_intents"][:5]): - lines.append(f" {i+1}. {t}(出现 {f} 次)") + lines.append(f" {i + 1}. {t}(出现 {f} 次)") else: lines.append(" 暂无明确信号") lines.append("") @@ -201,27 +214,27 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: low_price = max(30, int(avg_price * 0.4)) mid_price = int(avg_price * 0.8) high_price = int(avg_price * 1.5) - lines.append(f" 💰 低价位(走量引流款):¥{low_price}-{int(low_price*1.5)}") - lines.append(f" - 基础功能款,锁定价格敏感用户,利润率约{int(margin*100*0.7)}%") - lines.append(f" 💰 中价位(利润主力款):¥{mid_price}-{int(mid_price*1.4)}") - lines.append(f" - 主流功能+品质升级,利润率约{int(margin*100)}%") - lines.append(f" 💰 高价位(品牌形象款):¥{high_price}-{int(high_price*1.6)}") - lines.append(f" - 高端材质/设计,利润率约{int(margin*100*1.2)}%") + lines.append(f" 💰 低价位(走量引流款):¥{low_price}-{int(low_price * 1.5)}") + lines.append(f" - 基础功能款,锁定价格敏感用户,利润率约{int(margin * 100 * 0.7)}%") + lines.append(f" 💰 中价位(利润主力款):¥{mid_price}-{int(mid_price * 1.4)}") + lines.append(f" - 主流功能+品质升级,利润率约{int(margin * 100)}%") + lines.append(f" 💰 高价位(品牌形象款):¥{high_price}-{int(high_price * 1.6)}") + lines.append(f" - 高端材质/设计,利润率约{int(margin * 100 * 1.2)}%") else: lines.append(" (需补充价格数据后生成)") lines.append("") lines.append("【选品综合评分】") - lines.append(f"┌─────────────────────┬──────┐") - lines.append(f"│ 维度 │ 评分 │") - lines.append(f"├─────────────────────┼──────┤") + lines.append("┌─────────────────────┬──────┐") + lines.append("│ 维度 │ 评分 │") + lines.append("├─────────────────────┼──────┤") lines.append(f"│ 利润空间 │ {aggregated.get('profit_score', 0):>3} │") lines.append(f"│ 物流友好 │ {aggregated.get('logistics_score', 0):>3} │") lines.append(f"│ 竞争强度(分越高越好) │ {aggregated.get('competition_score', 0):>3} │") lines.append(f"│ 市场需求 │ {aggregated.get('demand_score', 0):>3} │") - lines.append(f"├─────────────────────┼──────┤") + lines.append("├─────────────────────┼──────┤") lines.append(f"│ 选品综合评分 │ {sel_score:>3} │") - lines.append(f"└─────────────────────┴──────┘") + lines.append("└─────────────────────┴──────┘") if sel_score >= 70: lines.append("✅ 推荐进入,综合条件良好") elif sel_score >= 45: @@ -236,7 +249,7 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: if sel_score < 45: warnings.append("综合评分偏低,建议寻找替代品类") if avg_return > 0.08: - warnings.append(f"退货率偏高({avg_return*100:.0f}%),注意控制品质") + warnings.append(f"退货率偏高({avg_return * 100:.0f}%),注意控制品质") if comp_score < 35: warnings.append("竞争激烈,可能需要大量广告投入") if evergreen < 0.5: diff --git a/src/api/dependencies.py b/src/api/dependencies.py index 4a3696d..4118d6d 100644 --- a/src/api/dependencies.py +++ b/src/api/dependencies.py @@ -3,7 +3,9 @@ ==================================== 所有 API 端点通过 Depends(get_app_state) 获取 AppState。 """ -from fastapi import Request, HTTPException + +from fastapi import HTTPException, Request + from src.core.state import AppState @@ -19,4 +21,3 @@ async def get_app_state(request: Request) -> AppState: async def get_app_state_or_none(request: Request) -> AppState: """不抛 503 的版本,供内部调用方自行处理""" return request.app.state.app_state - diff --git a/src/api/main.py b/src/api/main.py index c193093..725b5eb 100644 --- a/src/api/main.py +++ b/src/api/main.py @@ -5,29 +5,40 @@ 启动: uv run uvicorn src.api.main:app --port 8000 """ -import sys + import os -import uuid +import sys import traceback -from pathlib import Path +import uuid from contextlib import asynccontextmanager +from pathlib import Path import structlog from fastapi import FastAPI, Request -from fastapi.staticfiles import StaticFiles -from fastapi.responses import FileResponse, JSONResponse from fastapi.middleware.cors import CORSMiddleware +from fastapi.responses import FileResponse, JSONResponse +from fastapi.staticfiles import StaticFiles from starlette.middleware.base import BaseHTTPMiddleware sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))) +from src.api.routes import ( + crawl, + health, + insight, + insight_stream, + inspiration, + opportunities, + qa, + qa_stream, + trending, +) from src.config import settings from src.core.state import init_app_state -from src.api.routes import health, qa, insight, crawl, qa_stream, insight_stream, opportunities, trending, inspiration - # ===== 中间件 ===== + class RequestIDMiddleware(BaseHTTPMiddleware): """为每个请求注入 X-Request-ID,绑定到 structlog""" @@ -75,6 +86,7 @@ async def global_exception_handler(request: Request, exc: Exception): # ===== 生命周期 ===== + @asynccontextmanager async def lifespan(app: FastAPI): """应用启动时初始化 AppState,关闭时清理""" @@ -103,16 +115,18 @@ async def lifespan(app: FastAPI): if settings.rate_limit_enabled: try: from slowapi import Limiter, _rate_limit_exceeded_handler - from slowapi.util import get_remote_address from slowapi.errors import RateLimitExceeded + from slowapi.util import get_remote_address limiter = Limiter(key_func=get_remote_address, default_limits=["200/minute"]) app.state.limiter = limiter app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler) import structlog + structlog.get_logger().info("rate_limit_enabled") except ImportError: import structlog + structlog.get_logger().warning("rate_limit_skipped", reason="slowapi_not_installed") # ---- 注册中间件(顺序重要)---- @@ -146,6 +160,7 @@ async def lifespan(app: FastAPI): if react_assets.is_dir(): app.mount("/assets", StaticFiles(directory=str(react_assets)), name="react_assets") + # 统一前端入口:dist 优先 > static 回退 @app.get("/") async def serve_frontend_root(): @@ -154,10 +169,14 @@ async def serve_frontend_root(): return FileResponse(str(react_dist / "index.html")) if static_dir.exists(): return FileResponse(str(static_dir / "index.html")) - return JSONResponse({"message": "Frontend not built. Run: cd frontend && npm run build"}, status_code=404) + return JSONResponse( + {"message": "Frontend not built. Run: cd frontend && npm run build"}, status_code=404 + ) + # SPA fallback: 非 API 路径回退到 index.html(仅生产模式) if react_dist.is_dir(): + @app.get("/{full_path:path}") async def serve_react_spa(full_path: str): file_path = react_dist / full_path @@ -165,6 +184,7 @@ async def serve_react_spa(full_path: str): return FileResponse(str(file_path)) return FileResponse(str(react_dist / "index.html")) + # 开发回退:托管旧 static/ 目录 if static_dir.exists() and not react_dist.is_dir(): app.mount("/static", StaticFiles(directory=str(static_dir)), name="static") @@ -172,6 +192,7 @@ async def serve_react_spa(full_path: str): if __name__ == "__main__": import uvicorn + print("RedNote Insight API starting...") print(" API: http://localhost:8000") print(" Front: http://localhost:8000") diff --git a/src/api/routes/crawl.py b/src/api/routes/crawl.py index 22860a5..431622d 100644 --- a/src/api/routes/crawl.py +++ b/src/api/routes/crawl.py @@ -1,7 +1,9 @@ """crawl.py — 数据抓取与登录端点""" -import os -import json + import asyncio +import json +import os + from fastapi import APIRouter, Depends from fastapi.responses import JSONResponse, StreamingResponse from pydantic import BaseModel @@ -26,12 +28,15 @@ def _sse_event(event: str, data: dict | str) -> str: async def crawler_status(state: AppState = Depends(get_app_state)): """查询爬虫状态""" from src.crawler import CrawlerInterface + crawler = CrawlerInterface(raw_dir=state.raw_dir) - return JSONResponse(content={ - "available": crawler.is_available, - "needs_login": crawler.needs_login, - "is_cloud": crawler.is_cloud, - }) + return JSONResponse( + content={ + "available": crawler.is_available, + "needs_login": crawler.needs_login, + "is_cloud": crawler.is_cloud, + } + ) @router.post("/api/crawler/login") @@ -40,6 +45,7 @@ async def crawler_login(state: AppState = Depends(get_app_state)): SSE 端点:交互式小红书登录。 打开浏览器→显示二维码→等待扫码→保存 cookie。 """ + async def event_stream(): from src.crawler import CrawlerInterface @@ -50,10 +56,12 @@ async def event_stream(): return if not crawler.needs_login: - yield _sse_event("error", {"message": f"爬虫不可用,请检查配置"}) + yield _sse_event("error", {"message": "爬虫不可用,请检查配置"}) return - yield _sse_event("stage", {"stage": "login", "message": "正在打开小红书登录页,请在浏览器中扫码登录..."}) + yield _sse_event( + "stage", {"stage": "login", "message": "正在打开小红书登录页,请在浏览器中扫码登录..."} + ) await asyncio.sleep(0) login_ok = await asyncio.to_thread(crawler.login, 5) @@ -77,6 +85,7 @@ async def event_stream(): async def trigger_crawl(req: CrawlRequest, state: AppState = Depends(get_app_state)): """触发数据抓取""" from src.crawler import CrawlerInterface + crawler = CrawlerInterface(raw_dir=state.raw_dir) result = await asyncio.to_thread(crawler.crawl, req.category, req.count) @@ -85,9 +94,11 @@ async def trigger_crawl(req: CrawlRequest, state: AppState = Depends(get_app_sta await state.rebuild_indexes() state.stats["total_notes"] = len(os.listdir(state.raw_dir)) - return JSONResponse(content={ - "success": result["count"] > 0, - "method": result["method"], - "count": result["count"], - "message": f"抓取完成: {result['count']} 篇" if result["count"] > 0 else "抓取失败", - }) + return JSONResponse( + content={ + "success": result["count"] > 0, + "method": result["method"], + "count": result["count"], + "message": f"抓取完成: {result['count']} 篇" if result["count"] > 0 else "抓取失败", + } + ) diff --git a/src/api/routes/health.py b/src/api/routes/health.py index 52edcc7..a93cf4d 100644 --- a/src/api/routes/health.py +++ b/src/api/routes/health.py @@ -1,5 +1,7 @@ """health.py — 健康检查 + 统计端点""" + from fastapi import APIRouter, Depends + from src.api.dependencies import get_app_state from src.core.state import AppState diff --git a/src/api/routes/insight.py b/src/api/routes/insight.py index 778659d..d90d6bf 100644 --- a/src/api/routes/insight.py +++ b/src/api/routes/insight.py @@ -1,6 +1,8 @@ """insight.py — 选品洞察端点""" -import time + import asyncio +import time + from fastapi import APIRouter, Depends from pydantic import BaseModel @@ -31,23 +33,27 @@ async def run_insight(req: InsightRequest, state: AppState = Depends(get_app_sta result = await _run_insight_async(req.category, state, mode=req.mode) elapsed = round(time.time() - t0, 2) return InsightResponse( - success=True, category=req.category, mode=req.mode, - report=result["report"], notes_count=result["notes_count"], - generated_count=result["generated_count"], elapsed=elapsed, + success=True, + category=req.category, + mode=req.mode, + report=result["report"], + notes_count=result["notes_count"], + generated_count=result["generated_count"], + elapsed=elapsed, ) async def _run_insight_async(query: str, state: AppState, mode: str = "selection") -> dict: """执行洞察管道(全异步)""" from src.agents.comment_agent import CommentAnalyzer + from src.agents.creator_agent import CreatorGenerator from src.agents.demand_agent import DemandAggregator from src.agents.insight_agent import InsightGenerator - from src.agents.creator_agent import CreatorGenerator from src.config import RERANKER_THRESHOLD from src.crawler import CrawlerInterface - MIN_NOTES = 10 - CRAWL_COUNT = 30 + min_notes = 10 + crawl_count = 30 async def _do_insight(docs, category): analyzer = CommentAnalyzer(raw_dir=state.raw_dir) @@ -67,7 +73,9 @@ async def _do_insight(docs, category): report += f"\n\n(注:LLM 生成失败,使用模板兜底。错误:{e})" return report - docs = await state.hybrid_retriever.ahybrid_search(query, k=MIN_NOTES, bm25_k=40, final_k=MIN_NOTES) + docs = await state.hybrid_retriever.ahybrid_search( + query, k=min_notes, bm25_k=40, final_k=min_notes + ) if not docs: docs = [] @@ -89,36 +97,53 @@ async def _do_insight(docs, category): else: return { "report": f"知识库无「{query}」数据,且爬虫未登录。\n\n" - f"请使用流式接口(POST /api/insight/stream)触发交互式登录,\n" - f"或先在命令行运行:\n" - f" uv run python src/real_crawler.py \"{query}\"\n" - f"完成登录后再试。", - "notes_count": 0, "generated_count": 0, + f"请使用流式接口(POST /api/insight/stream)触发交互式登录,\n" + f"或先在命令行运行:\n" + f' uv run python src/real_crawler.py "{query}"\n' + f"完成登录后再试。", + "notes_count": 0, + "generated_count": 0, } else: return { "report": f"知识库无「{query}」数据,且爬虫不可用。\n\n" - f"请先在命令行运行 `uv run python src/real_crawler.py \"{query}\"` 登录并抓取数据。", - "notes_count": 0, "generated_count": 0, + "请先在命令行运行 `uv run python " + f'src/real_crawler.py "{query}"` 登录并抓取数据。', + "notes_count": 0, + "generated_count": 0, } - result = await asyncio.to_thread(crawler.crawl, query, CRAWL_COUNT) + result = await asyncio.to_thread(crawler.crawl, query, crawl_count) crawled_count = result["count"] if crawled_count == 0: - return {"report": f"抱歉,无法从小红书获取「{query}」的数据。", "notes_count": 0, "generated_count": 0} + return { + "report": f"抱歉,无法从小红书获取「{query}」的数据。", + "notes_count": 0, + "generated_count": 0, + } await state.rebuild_indexes() await asyncio.sleep(0.5) - fresh_docs = await state.hybrid_retriever.ahybrid_search(query, k=MIN_NOTES, bm25_k=40, final_k=MIN_NOTES) + fresh_docs = await state.hybrid_retriever.ahybrid_search( + query, k=min_notes, bm25_k=40, final_k=min_notes + ) fresh_scores = await state.reranker.arerank(query, fresh_docs) if fresh_docs else [] - fresh_relevant = [doc for doc, s in zip(fresh_docs, fresh_scores) if s >= RERANKER_THRESHOLD] + fresh_relevant = [ + doc for doc, s in zip(fresh_docs, fresh_scores) if s >= RERANKER_THRESHOLD + ] if not fresh_relevant: - return {"report": f"已从小红书抓取 {crawled_count} 篇笔记,但检索仍未匹配。", "notes_count": 0, "generated_count": crawled_count} + return { + "report": f"已从小红书抓取 {crawled_count} 篇笔记,但检索仍未匹配。", + "notes_count": 0, + "generated_count": crawled_count, + } report = await _do_insight(fresh_relevant, query) report = f"(📥 已从小红书实时抓取「{query}」{crawled_count} 篇真实笔记)\n\n{report}" - return {"report": report, - "notes_count": len(relevant) if crawled_count == 0 else len(fresh_relevant), - "generated_count": crawled_count} + return { + "report": report, + "notes_count": len(relevant) if crawled_count == 0 else len(fresh_relevant), + "generated_count": crawled_count, + } diff --git a/src/api/routes/insight_stream.py b/src/api/routes/insight_stream.py index 4c20e3c..553432a 100644 --- a/src/api/routes/insight_stream.py +++ b/src/api/routes/insight_stream.py @@ -10,16 +10,17 @@ -d '{"category":"磁吸感应灯"}' """ +import asyncio import json import time -import asyncio + from fastapi import APIRouter, Depends from fastapi.responses import StreamingResponse from pydantic import BaseModel from src.api.dependencies import get_app_state -from src.core.state import AppState from src.config import RERANKER_THRESHOLD +from src.core.state import AppState from src.logger import logger router = APIRouter(tags=["insight-stream"]) @@ -46,14 +47,14 @@ async def _stream_generator(gen, aggregated: dict, category: str, event_type: st yield _sse_event(event_type, {"token": report}) -async def _run_analysis_pipeline(category: str, state: AppState, MIN_NOTES=10, CRAWL_COUNT=30): +async def _run_analysis_pipeline(category: str, state: AppState, min_notes=10, crawl_count=30): """运行检索→分析→聚合管道,返回 (aggregated, notes_count)""" from src.agents.comment_agent import CommentAnalyzer from src.agents.demand_agent import DemandAggregator from src.crawler import CrawlerInterface docs = await state.hybrid_retriever.ahybrid_search( - category, k=MIN_NOTES, bm25_k=40, final_k=MIN_NOTES + category, k=min_notes, bm25_k=40, final_k=min_notes ) if docs: scores = await state.reranker.arerank(category, docs) @@ -79,7 +80,7 @@ async def _run_analysis_pipeline(category: str, state: AppState, MIN_NOTES=10, C if not crawler.is_available: return None, 0, False - result = await asyncio.to_thread(crawler.crawl, category, CRAWL_COUNT) + result = await asyncio.to_thread(crawler.crawl, category, crawl_count) if result["count"] == 0: return None, 0, False @@ -87,10 +88,12 @@ async def _run_analysis_pipeline(category: str, state: AppState, MIN_NOTES=10, C await asyncio.sleep(0.5) fresh_docs = await state.hybrid_retriever.ahybrid_search( - category, k=MIN_NOTES, bm25_k=40, final_k=MIN_NOTES + category, k=min_notes, bm25_k=40, final_k=min_notes ) fresh_scores = await state.reranker.arerank(category, fresh_docs) if fresh_docs else [] - fresh_relevant = [doc for doc, s in zip(fresh_docs, fresh_scores) if s >= RERANKER_THRESHOLD] + fresh_relevant = [ + doc for doc, s in zip(fresh_docs, fresh_scores) if s >= RERANKER_THRESHOLD + ] if not fresh_relevant: return None, 0, False @@ -110,42 +113,61 @@ async def event_stream(): category = req.category try: - from src.agents.insight_agent import InsightGenerator from src.agents.creator_agent import CreatorGenerator + from src.agents.insight_agent import InsightGenerator # ── 阶段 1: 检索 ── - yield _sse_event("stage", {"stage": "retrieve", "message": f"正在检索「{category}」相关笔记..."}) + yield _sse_event( + "stage", {"stage": "retrieve", "message": f"正在检索「{category}」相关笔记..."} + ) await asyncio.sleep(0) - docs = await state.hybrid_retriever.ahybrid_search(category, k=10, bm25_k=40, final_k=10) + docs = await state.hybrid_retriever.ahybrid_search( + category, k=10, bm25_k=40, final_k=10 + ) if docs: scores = await state.reranker.arerank(category, docs) relevant = [doc for doc, s in zip(docs, scores) if s >= RERANKER_THRESHOLD] else: relevant = [] - yield _sse_event("stage", { - "stage": "retrieved", - "message": f"检索到 {len(relevant)} 篇相关笔记", - "note_count": len(relevant), - }) + yield _sse_event( + "stage", + { + "stage": "retrieved", + "message": f"检索到 {len(relevant)} 篇相关笔记", + "note_count": len(relevant), + }, + ) await asyncio.sleep(0) # ── 阶段 2: 分析 + 聚合 ── aggregated, notes_count, ok = await _run_analysis_pipeline(category, state) if not ok or aggregated is None: - yield _sse_event("error", {"message": f"无法获取「{category}」的分析数据,请确认品类名称或尝试其他关键词"}) + yield _sse_event( + "error", + { + "message": f"无法获取「{category}」的分析数据," + "请确认品类名称或尝试其他关键词" + }, + ) return - yield _sse_event("stage", { - "stage": "aggregated", - "message": f"识别到 {len(aggregated.get('top_complaints', []))} 个痛点,{len(aggregated.get('top_purchase_intents', []))} 个需求信号", - }) + yield _sse_event( + "stage", + { + "stage": "aggregated", + "message": f"识别到 {len(aggregated.get('top_complaints', []))} 个痛点," + f"{len(aggregated.get('top_purchase_intents', []))} 个需求信号", + }, + ) await asyncio.sleep(0) # ── 阶段 3: 生成选品报告 ── - yield _sse_event("stage", {"stage": "generate_selection", "message": "正在生成选品洞察报告..."}) + yield _sse_event( + "stage", {"stage": "generate_selection", "message": "正在生成选品洞察报告..."} + ) ins_gen = InsightGenerator() async for event in _stream_generator(ins_gen, aggregated, category, "token:selection"): yield event @@ -153,7 +175,9 @@ async def event_stream(): yield _sse_event("stage", {"stage": "selection_done", "message": "选品报告完成"}) # ── 阶段 4: 生成选题方案 ── - yield _sse_event("stage", {"stage": "generate_creator", "message": "正在生成选题方案..."}) + yield _sse_event( + "stage", {"stage": "generate_creator", "message": "正在生成选题方案..."} + ) cr_gen = CreatorGenerator() async for event in _stream_generator(cr_gen, aggregated, category, "token:creator"): yield event @@ -162,10 +186,13 @@ async def event_stream(): # ── 完成 ── elapsed = round(time.time() - t0, 2) - yield _sse_event("done", { - "elapsed": elapsed, - "note_count": notes_count, - }) + yield _sse_event( + "done", + { + "elapsed": elapsed, + "note_count": notes_count, + }, + ) except Exception as e: logger.error(f"insight_stream_error: {e}") diff --git a/src/api/routes/inspiration.py b/src/api/routes/inspiration.py index b5685b2..041073c 100644 --- a/src/api/routes/inspiration.py +++ b/src/api/routes/inspiration.py @@ -5,9 +5,10 @@ GET /api/inspiration?category=美妆 → 按品类筛选 GET /api/inspiration/categories → 列出所有品类 """ + from fastapi import APIRouter, Query -from src.data.inspiration import get_inspiration, get_categories +from src.data.inspiration import get_categories, get_inspiration router = APIRouter(prefix="/api/inspiration", tags=["inspiration"]) diff --git a/src/api/routes/opportunities.py b/src/api/routes/opportunities.py index ff2ad00..a1b411f 100644 --- a/src/api/routes/opportunities.py +++ b/src/api/routes/opportunities.py @@ -8,12 +8,11 @@ GET /api/opportunities → 全部品类排行列表 GET /api/opportunities/{cat} → 单个品类详细信息(未知品类返回启发式估算) """ -import os + import re -import yaml as pyyaml from pathlib import Path -from typing import Optional +import yaml as pyyaml from fastapi import APIRouter, HTTPException from src.config import RAW_DIR @@ -74,22 +73,87 @@ def _safe(val, default=0): def _classify_category(name: str) -> dict: """根据品类名关键词推断品类属性""" if any(k in name for k in DIGITAL_KEYS): - return dict(base_price=60, base_cost=18, base_weight=0.2, base_margin=0.65, base_sales=4500, comp_level="高", cat_type="常青款", tags=["需求旺", "更新快"]) + return dict( + base_price=60, + base_cost=18, + base_weight=0.2, + base_margin=0.65, + base_sales=4500, + comp_level="高", + cat_type="常青款", + tags=["需求旺", "更新快"], + ) if any(k in name for k in CLOTHING_KEYS): - return dict(base_price=188, base_cost=45, base_weight=0.5, base_margin=0.65, base_sales=3500, comp_level="高", cat_type="季节款" if "衣" in name else "常青款", tags=["需求旺", "竞争大"]) + return dict( + base_price=188, + base_cost=45, + base_weight=0.5, + base_margin=0.65, + base_sales=3500, + comp_level="高", + cat_type="季节款" if "衣" in name else "常青款", + tags=["需求旺", "竞争大"], + ) if any(k in name for k in FOOD_KEYS): - return dict(base_price=45, base_cost=15, base_weight=0.4, base_margin=0.60, base_sales=6000, comp_level="高", cat_type="常青款", tags=["复购高", "利润中等"]) + return dict( + base_price=45, + base_cost=15, + base_weight=0.4, + base_margin=0.60, + base_sales=6000, + comp_level="高", + cat_type="常青款", + tags=["复购高", "利润中等"], + ) if any(k in name for k in HOME_KEYS): - return dict(base_price=80, base_cost=25, base_weight=0.6, base_margin=0.65, base_sales=4000, comp_level="中", cat_type="常青款", tags=["刚需品", "利润一般"]) + return dict( + base_price=80, + base_cost=25, + base_weight=0.6, + base_margin=0.65, + base_sales=4000, + comp_level="中", + cat_type="常青款", + tags=["刚需品", "利润一般"], + ) if any(k in name for k in BEAUTY_KEYS): - return dict(base_price=120, base_cost=35, base_weight=0.3, base_margin=0.70, base_sales=5000, comp_level="高", cat_type="常青款", tags=["利润高", "品牌多"]) - return dict(base_price=80, base_cost=25, base_weight=0.5, base_margin=0.60, base_sales=3000, comp_level="中", cat_type="常青款", tags=["需验证", "数据采集中"]) + return dict( + base_price=120, + base_cost=35, + base_weight=0.3, + base_margin=0.70, + base_sales=5000, + comp_level="高", + cat_type="常青款", + tags=["利润高", "品牌多"], + ) + return dict( + base_price=80, + base_cost=25, + base_weight=0.5, + base_margin=0.60, + base_sales=3000, + comp_level="中", + cat_type="常青款", + tags=["需验证", "数据采集中"], + ) -def _compute_scores(avg_price, avg_cost, avg_weight, avg_margin, avg_sales, - avg_likes=0, avg_comments=0, brand_count=0, - competitions=None, difficulties=None, - differentiations=None, cat_types=None, n=0): +def _compute_scores( + avg_price, + avg_cost, + avg_weight, + avg_margin, + avg_sales, + avg_likes=0, + avg_comments=0, + brand_count=0, + competitions=None, + difficulties=None, + differentiations=None, + cat_types=None, + n=0, +): """通用评分计算,可传入估算值或实际聚合值""" if competitions is None: competitions = [] @@ -102,64 +166,90 @@ def _compute_scores(avg_price, avg_cost, avg_weight, avg_margin, avg_sales, price_cost_ratio = avg_price / avg_cost if avg_cost > 0 else 3.0 - profit_score = min(100, int( - min(price_cost_ratio / 5, 1.0) * 40 + - min(avg_margin / 0.7, 1.0) * 40 + - 20 - )) - - logistics_score = min(100, int(35 + - (30 if avg_weight > 0 and avg_weight < 0.3 else 0) + - (15 if avg_weight > 0 and avg_weight < 1.0 else 0) + - (10 if competitions.count("高") < max(n, 1) * 0.5 else 0) - )) - - demand_score = min(100, int( - min(avg_likes / 200, 1.0) * 25 + - min(avg_comments / 50, 1.0) * 20 + - min(avg_sales / 5000, 1.0) * 30 + - min(brand_count * 3, 15) + - 10 - )) + profit_score = min( + 100, int(min(price_cost_ratio / 5, 1.0) * 40 + min(avg_margin / 0.7, 1.0) * 40 + 20) + ) + + logistics_score = min( + 100, + int( + 35 + + (30 if avg_weight > 0 and avg_weight < 0.3 else 0) + + (15 if avg_weight > 0 and avg_weight < 1.0 else 0) + + (10 if competitions.count("高") < max(n, 1) * 0.5 else 0) + ), + ) + + demand_score = min( + 100, + int( + min(avg_likes / 200, 1.0) * 25 + + min(avg_comments / 50, 1.0) * 20 + + min(avg_sales / 5000, 1.0) * 30 + + min(brand_count * 3, 15) + + 10 + ), + ) low_comp_ratio = competitions.count("低") / max(n, 1) easy_entry_ratio = difficulties.count("低") / max(n, 1) - competition_score = min(100, max(0, int( - low_comp_ratio * 40 + - easy_entry_ratio * 30 + - max(0, 5 - brand_count) * 5 + - 10 - ))) - - overall = max(0, min(100, int( - profit_score * 0.30 + - logistics_score * 0.20 + - competition_score * 0.20 + - demand_score * 0.30 - ))) - - rec = "强烈推荐" if overall >= 80 else "可尝试" if overall >= 65 else "谨慎进入" if overall >= 50 else "不建议" + competition_score = min( + 100, + max(0, int(low_comp_ratio * 40 + easy_entry_ratio * 30 + max(0, 5 - brand_count) * 5 + 10)), + ) + + overall = max( + 0, + min( + 100, + int( + profit_score * 0.30 + + logistics_score * 0.20 + + competition_score * 0.20 + + demand_score * 0.30 + ), + ), + ) + + rec = ( + "强烈推荐" + if overall >= 80 + else "可尝试" + if overall >= 65 + else "谨慎进入" + if overall >= 50 + else "不建议" + ) unique_diffs = list(dict.fromkeys([d for d in differentiations if d]))[:5] evergreen = cat_types.count("常青款") / max(n, 1) > 0.6 if n > 0 else True return dict( - scores=dict(profit=profit_score, logistics=logistics_score, - demand=demand_score, competition=competition_score, - overall=overall), - metrics=dict(avg_price=round(avg_price, 1), avg_cost=round(avg_cost, 1), - avg_profit_margin=round(avg_margin, 2), - price_cost_ratio=round(price_cost_ratio, 1), - avg_weight=round(avg_weight, 2), - avg_likes=round(avg_likes, 1), avg_comments=round(avg_comments, 1), - avg_return_rate=0.05, avg_monthly_sales=avg_sales, - brand_count=brand_count), + scores=dict( + profit=profit_score, + logistics=logistics_score, + demand=demand_score, + competition=competition_score, + overall=overall, + ), + metrics=dict( + avg_price=round(avg_price, 1), + avg_cost=round(avg_cost, 1), + avg_profit_margin=round(avg_margin, 2), + price_cost_ratio=round(price_cost_ratio, 1), + avg_weight=round(avg_weight, 2), + avg_likes=round(avg_likes, 1), + avg_comments=round(avg_comments, 1), + avg_return_rate=0.05, + avg_monthly_sales=avg_sales, + brand_count=brand_count, + ), recommendation=rec, differentiation_directions=unique_diffs, evergreen=evergreen, ) -def _calc_category_scores(cat_prefix: str) -> Optional[dict]: +def _calc_category_scores(cat_prefix: str) -> dict | None: """对一个品类下的所有文件做聚合评分""" raw_dir = Path(RAW_DIR) files = sorted(raw_dir.glob(f"{cat_prefix}_*.md")) @@ -209,9 +299,21 @@ def _calc_category_scores(cat_prefix: str) -> Optional[dict]: avg_weight, avg_margin = cls["base_weight"], cls["base_margin"] avg_sales = cls["base_sales"] - scores = _compute_scores(avg_price, avg_cost, avg_weight, avg_margin, avg_sales, - avg_likes, avg_comments, brand_count, - competitions, difficulties, differentiations, cat_types, n) + scores = _compute_scores( + avg_price, + avg_cost, + avg_weight, + avg_margin, + avg_sales, + avg_likes, + avg_comments, + brand_count, + competitions, + difficulties, + differentiations, + cat_types, + n, + ) return { "category": CATEGORY_NAMES.get(cat_prefix, cat_prefix), @@ -224,6 +326,7 @@ def _calc_category_scores(cat_prefix: str) -> Optional[dict]: # ===== 路由 ===== + @router.get("") async def list_opportunities(): """返回全部品类的机会评分排行""" @@ -269,9 +372,13 @@ async def get_opportunity_detail(category_name: str): # ===== 未知品类:启发式估算 ===== cls = _classify_category(category_name) - scores = _compute_scores(cls["base_price"], cls["base_cost"], - cls["base_weight"], cls["base_margin"], - cls["base_sales"]) + scores = _compute_scores( + cls["base_price"], + cls["base_cost"], + cls["base_weight"], + cls["base_margin"], + cls["base_sales"], + ) return { "category": category_name, @@ -281,4 +388,4 @@ async def get_opportunity_detail(category_name: str): **scores, "brands": [], "tags": cls["tags"], - } \ No newline at end of file + } diff --git a/src/api/routes/qa.py b/src/api/routes/qa.py index 047309d..75ab3be 100644 --- a/src/api/routes/qa.py +++ b/src/api/routes/qa.py @@ -1,14 +1,17 @@ """qa.py — QA 问答端点(快速通道)""" + import time + +import jieba from fastapi import APIRouter, Depends -from pydantic import BaseModel from langchain_openai import ChatOpenAI -import jieba +from pydantic import BaseModel from src.api.dependencies import get_app_state -from src.core.state import AppState from src.config import LLM_CONFIG from src.core.prompt_loader import get_prompt_loader +from src.core.query_utils import clean_query, is_brand_comparison +from src.core.state import AppState router = APIRouter(tags=["qa"]) @@ -25,20 +28,40 @@ class QAResponse(BaseModel): elapsed: float -from src.core.query_utils import clean_query, is_brand_comparison - _NO_DATA = "知识库中暂无该品类的数据。请换一个品类试试,或通过选品洞察触发实时抓取。" -_QUESTION_STOP = {"哪个", "哪款", "哪家", "品牌", "推荐", "好", "什么", "怎么", "如何", "多少", "对比", "测评", "排行", "有没有", "值得", "建议", "选择", "区别"} +_QUESTION_STOP = { + "哪个", + "哪款", + "哪家", + "品牌", + "推荐", + "好", + "什么", + "怎么", + "如何", + "多少", + "对比", + "测评", + "排行", + "有没有", + "值得", + "建议", + "选择", + "区别", +} + def _any_relevant(query: str, docs: list, threshold: int = 1) -> bool: q_words = set(w for w in jieba.cut(query) if len(w) > 1 and w not in _QUESTION_STOP) - if not q_words: return True + if not q_words: + return True hits = 0 for d in docs: - content = d.page_content if hasattr(d, 'page_content') else str(d) + content = d.page_content if hasattr(d, "page_content") else str(d) if any(w in content for w in q_words): hits += 1 - if hits >= threshold: return True + if hits >= threshold: + return True return False @@ -59,16 +82,20 @@ async def run_qa(req: QARequest, state: AppState = Depends(get_app_state)): # Rerank 重排序 from src.config import RERANKER_THRESHOLD + scores = await state.reranker.arerank(cleaned, docs) scored = sorted( [(d, s) for d, s in zip(docs, scores) if s >= RERANKER_THRESHOLD], - key=lambda x: x[1], reverse=True + key=lambda x: x[1], + reverse=True, ) docs = [d for d, _ in scored] if scored else docs[:5] - context = "\n---\n".join( - f"[文档{i+1}] {d.page_content}" for i, d in enumerate(docs) - ) if docs else "暂无相关文档" + context = ( + "\n---\n".join(f"[文档{i + 1}] {d.page_content}" for i, d in enumerate(docs)) + if docs + else "暂无相关文档" + ) prompt = get_prompt_loader().load("gen_answer", "v2") msg = prompt.format_messages(context=context, question=req.question) @@ -76,4 +103,6 @@ async def run_qa(req: QARequest, state: AppState = Depends(get_app_state)): resp = await llm.ainvoke(msg) elapsed = round(time.time() - t0, 2) - return QAResponse(success=True, question=req.question, answer=resp.content.strip(), elapsed=elapsed) + return QAResponse( + success=True, question=req.question, answer=resp.content.strip(), elapsed=elapsed + ) diff --git a/src/api/routes/qa_stream.py b/src/api/routes/qa_stream.py index 0fa8cfa..49901b8 100644 --- a/src/api/routes/qa_stream.py +++ b/src/api/routes/qa_stream.py @@ -12,22 +12,22 @@ -d '{"question":"磁吸感应灯哪个品牌好"}' """ +import asyncio import json import time -import re -import asyncio + +import jieba from fastapi import APIRouter, Depends from fastapi.responses import StreamingResponse -from pydantic import BaseModel from langchain_openai import ChatOpenAI +from pydantic import BaseModel from src.api.dependencies import get_app_state -from src.core.state import AppState from src.config import LLM_CONFIG -from src.logger import logger from src.core.prompt_loader import get_prompt_loader -import jieba -from src.core.query_utils import clean_query, is_brand_comparison, resolve_k, resolve_bm25_k +from src.core.query_utils import clean_query, is_brand_comparison +from src.core.state import AppState +from src.logger import logger router = APIRouter(tags=["qa-stream"]) @@ -47,21 +47,37 @@ def _sse_event(event: str, data: dict | str) -> str: # ── 通用疑问词(不参与相关性判断)── -_QUESTION_STOP = {"哪个", "哪款", "哪家", "品牌", "推荐", "好", "什么", "怎么", "如何", "多少", "对比", "测评", "排行", "有没有", "值得", "建议", "选择", "区别"} +_QUESTION_STOP = { + "哪个", + "哪款", + "哪家", + "品牌", + "推荐", + "好", + "什么", + "怎么", + "如何", + "多少", + "对比", + "测评", + "排行", + "有没有", + "值得", + "建议", + "选择", + "区别", +} def _any_doc_relevant(query: str, docs: list, threshold: int = 1) -> bool: """快速相关性检查:至少 threshold 篇文档包含查询中的品类关键词""" # 只取有意义的品类词(排除通用疑问词和短词) - q_words = set( - w for w in jieba.cut(query) - if len(w) > 1 and w not in _QUESTION_STOP - ) + q_words = set(w for w in jieba.cut(query) if len(w) > 1 and w not in _QUESTION_STOP) if not q_words: return True # 全是疑问词时放行 hits = 0 for d in docs: - content = d.page_content if hasattr(d, 'page_content') else str(d) + content = d.page_content if hasattr(d, "page_content") else str(d) if any(w in content for w in q_words): hits += 1 if hits >= threshold: @@ -86,7 +102,7 @@ async def event_stream(): is_brand = is_brand_comparison(cleaned) k = BRAND_K if is_brand else BASE_K - yield _sse_event("stage", {"stage": "retrieve", "message": f"正在检索..."}) + yield _sse_event("stage", {"stage": "retrieve", "message": "正在检索..."}) await asyncio.sleep(0) # ── 阶段 2: 直连检索 ── @@ -94,18 +110,23 @@ async def event_stream(): cleaned, k=k, bm25_k=max(40, k * 5), final_k=k ) - yield _sse_event("stage", { - "stage": "retrieved", - "message": f"检索到 {len(docs)} 篇文档", - "doc_count": len(docs), - }) + yield _sse_event( + "stage", + { + "stage": "retrieved", + "message": f"检索到 {len(docs)} 篇文档", + "doc_count": len(docs), + }, + ) await asyncio.sleep(0) # ── 相关性门禁:无匹配时直接返回 ── if docs and not _any_doc_relevant(cleaned, docs, threshold=1): yield _sse_event("token", {"token": _NO_DATA_MSG}) elapsed = round(time.time() - t0, 2) - yield _sse_event("done", {"answer": _NO_DATA_MSG, "elapsed": elapsed, "doc_count": 0}) + yield _sse_event( + "done", {"answer": _NO_DATA_MSG, "elapsed": elapsed, "doc_count": 0} + ) return # ── 阶段 3: Rerank 重排序 ── @@ -113,19 +134,24 @@ async def event_stream(): await asyncio.sleep(0) from src.config import RERANKER_THRESHOLD + scores = await state.reranker.arerank(cleaned, docs) scored = sorted( [(d, s) for d, s in zip(docs, scores) if s >= RERANKER_THRESHOLD], - key=lambda x: x[1], reverse=True + key=lambda x: x[1], + reverse=True, ) docs = [d for d, _ in scored] if scored else docs[:5] logger.info(f"rerank: {len(scored)}/{len(scores)} docs above threshold") - yield _sse_event("stage", { - "stage": "reranked", - "message": f"重排序完成,保留 {len(docs)} 篇高相关文档", - "doc_count": len(docs), - }) + yield _sse_event( + "stage", + { + "stage": "reranked", + "message": f"重排序完成,保留 {len(docs)} 篇高相关文档", + "doc_count": len(docs), + }, + ) await asyncio.sleep(0) # ── 阶段 4: 流式生成 ── @@ -133,7 +159,7 @@ async def event_stream(): if docs: context = "\n---\n".join( - f"[文档{i+1}] {d.page_content}" for i, d in enumerate(docs) + f"[文档{i + 1}] {d.page_content}" for i, d in enumerate(docs) ) else: context = "暂无相关文档" @@ -153,11 +179,14 @@ async def event_stream(): # ── 阶段 4: 完成 ── elapsed = round(time.time() - t0, 2) - yield _sse_event("done", { - "answer": full_answer, - "elapsed": elapsed, - "doc_count": len(docs), - }) + yield _sse_event( + "done", + { + "answer": full_answer, + "elapsed": elapsed, + "doc_count": len(docs), + }, + ) except Exception as e: logger.error(f"qa_stream_error: {e}") diff --git a/src/api/routes/trending.py b/src/api/routes/trending.py index 3be677a..2cd03ee 100644 --- a/src/api/routes/trending.py +++ b/src/api/routes/trending.py @@ -8,17 +8,15 @@ GET /api/trending → 热门搜索词排行 GET /api/trending/refresh → 触发爬虫刷新热词数据 """ -import os -import re + +import asyncio import json +import os import random -import asyncio -from pathlib import Path -from typing import Optional from datetime import datetime, timedelta +from pathlib import Path -from fastapi import APIRouter, HTTPException -from pydantic import BaseModel +from fastapi import APIRouter from src.config import RAW_DIR @@ -26,8 +24,11 @@ # ===== 热词缓存文件 ===== -TRENDING_CACHE = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))), - "data", "trending_cache.json") +TRENDING_CACHE = os.path.join( + os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))), + "data", + "trending_cache.json", +) # ===== 热门选品词库(按品类分类) ===== @@ -41,38 +42,33 @@ {"keyword": "盲盒", "category": "潮玩", "trend": "up", "hots": 0}, {"keyword": "手机壳", "category": "数码", "trend": "stable", "hots": 0}, {"keyword": "蓝牙耳机", "category": "数码", "trend": "stable", "hots": 0}, - # 👗 服饰 {"keyword": "健身服", "category": "服饰", "trend": "up", "hots": 0}, {"keyword": "风衣", "category": "服饰", "trend": "seasonal", "hots": 0}, {"keyword": "瑜伽裤", "category": "服饰", "trend": "up", "hots": 0}, {"keyword": "冲锋衣", "category": "服饰", "trend": "up", "hots": 0}, - # 🍜 食品 {"keyword": "辣条", "category": "食品", "trend": "stable", "hots": 0}, {"keyword": "养生茶", "category": "食品", "trend": "up", "hots": 0}, {"keyword": "即食早餐", "category": "食品", "trend": "up", "hots": 0}, - # 💄 美妆个护 {"keyword": "素颜霜", "category": "美妆", "trend": "up", "hots": 0}, {"keyword": "护发精油", "category": "个护", "trend": "up", "hots": 0}, {"keyword": "补水面膜", "category": "美妆", "trend": "stable", "hots": 0}, - # 🐱 宠物 {"keyword": "猫粮", "category": "宠物", "trend": "up", "hots": 0}, {"keyword": "宠物玩具", "category": "宠物", "trend": "up", "hots": 0}, - # 🔧 其他热门 {"keyword": "健身器材", "category": "运动", "trend": "stable", "hots": 0}, {"keyword": "茶杯", "category": "家居", "trend": "stable", "hots": 0}, ] -def _load_cache() -> Optional[dict]: +def _load_cache() -> dict | None: """加载热词缓存""" if os.path.exists(TRENDING_CACHE): try: - with open(TRENDING_CACHE, "r", encoding="utf-8") as f: + with open(TRENDING_CACHE, encoding="utf-8") as f: return json.load(f) except Exception: return None @@ -94,10 +90,17 @@ def _estimate_hots_from_notes(category: str) -> int: # 找品类前缀 prefix_map = { - "磁吸感应灯": "cixi", "健身服": "健身", "风衣": "风衣", - "辣条": "辣条", "茶杯": "茶杯", "桌面收纳": "dorm", - "盲盒": "box", "装饰画": "deco", "香薰": "scent", - "收纳盒": "store", "健身器材": "选健", + "磁吸感应灯": "cixi", + "健身服": "健身", + "风衣": "风衣", + "辣条": "辣条", + "茶杯": "茶杯", + "桌面收纳": "dorm", + "盲盒": "box", + "装饰画": "deco", + "香薰": "scent", + "收纳盒": "store", + "健身器材": "选健", } prefix = prefix_map.get(category) if not prefix: @@ -128,22 +131,26 @@ def _generate_trending(refresh: bool = False) -> list: items = [] for kw in HOT_KEYWORDS: hots = _estimate_hots_from_notes(kw["keyword"]) - items.append({ - "keyword": kw["keyword"], - "category": kw["category"], - "trend": kw["trend"], - "hots": hots, - "has_data": hots > 0, - }) + items.append( + { + "keyword": kw["keyword"], + "category": kw["category"], + "trend": kw["trend"], + "hots": hots, + "has_data": hots > 0, + } + ) # 按热度排序 items.sort(key=lambda x: x["hots"], reverse=True) # 缓存 - _save_cache({ - "updated_at": now.isoformat(), - "items": items, - }) + _save_cache( + { + "updated_at": now.isoformat(), + "items": items, + } + ) return items @@ -152,6 +159,7 @@ async def _trigger_crawl_for_keyword(keyword: str) -> bool: """触发爬虫抓取关键词数据""" try: from src.crawler import CrawlerInterface + crawler = CrawlerInterface(raw_dir=str(RAW_DIR)) if not crawler.is_available: return False @@ -195,4 +203,4 @@ async def batch_crawl(): "updated_at": datetime.now().isoformat(), "crawling": True, "message": "后台正在采集前 10 个热词数据,1-2 分钟后刷新查看结果", - } \ No newline at end of file + } diff --git a/src/core/database.py b/src/core/database.py index 3f67d70..f15ee63 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -12,18 +12,18 @@ """ import uuid +from collections.abc import AsyncGenerator from datetime import datetime -from typing import AsyncGenerator, Optional -from sqlalchemy import Column, String, Text, DateTime, Index -from sqlalchemy.dialects.postgresql import UUID, JSONB +from pgvector.sqlalchemy import Vector +from sqlalchemy import Column, DateTime, Index, Text +from sqlalchemy.dialects.postgresql import JSONB, UUID from sqlalchemy.ext.asyncio import ( AsyncSession, async_sessionmaker, create_async_engine, ) from sqlalchemy.orm import declarative_base -from pgvector.sqlalchemy import Vector from src.config import settings from src.logger import logger @@ -33,7 +33,7 @@ # ===== 引擎 & 会话工厂 ===== _engine = None -_session_factory: Optional[async_sessionmaker] = None +_session_factory: async_sessionmaker | None = None def _get_database_url() -> str: @@ -85,6 +85,7 @@ async def get_db() -> AsyncGenerator[AsyncSession, None]: # ===== 数据表定义 ===== + class DocumentTable(Base): """文档向量表 — 替代 ChromaDB collection""" @@ -117,14 +118,15 @@ def __repr__(self): # ===== 数据库初始化 ===== + async def init_db() -> None: """创建所有表和索引(幂等:已存在的表不会重建)""" engine = get_engine() async with engine.begin() as conn: # 确保 pgvector 扩展已启用 - await conn.run_sync(lambda sync_conn: sync_conn.execute( - "CREATE EXTENSION IF NOT EXISTS vector" - )) + await conn.run_sync( + lambda sync_conn: sync_conn.execute("CREATE EXTENSION IF NOT EXISTS vector") + ) # 创建所有表 await conn.run_sync(Base.metadata.create_all) logger.info("database_initialized", tables=["documents"]) @@ -140,6 +142,7 @@ async def drop_db() -> None: # ===== 向量操作辅助 ===== + async def insert_documents( session: AsyncSession, chunks: list, @@ -156,15 +159,17 @@ async def insert_documents( """ rows = [] for chunk, emb in zip(chunks, embeddings): - rows.append(DocumentTable( - content=chunk.page_content, - metadata_=chunk.metadata, - embedding=emb, - )) + rows.append( + DocumentTable( + content=chunk.page_content, + metadata_=chunk.metadata, + embedding=emb, + ) + ) session.add_all(rows) await session.flush() - logger.info(f"inserted_documents", count=len(rows)) + logger.info("inserted_documents", count=len(rows)) return len(rows) @@ -197,14 +202,12 @@ async def search_by_vector( ) rows = result.fetchall() - return [ - {"content": row[0], "metadata": row[1], "score": float(row[2])} - for row in rows - ] + return [{"content": row[0], "metadata": row[1], "score": float(row[2])} for row in rows] async def get_document_count(session: AsyncSession) -> int: """获取文档总数""" from sqlalchemy import text + result = await session.execute(text("SELECT COUNT(*) FROM documents")) return result.scalar() or 0 diff --git a/src/core/prompt_loader.py b/src/core/prompt_loader.py index d5fdbdc..8e00f50 100644 --- a/src/core/prompt_loader.py +++ b/src/core/prompt_loader.py @@ -13,10 +13,7 @@ """ import os -import re from pathlib import Path -from functools import lru_cache -from typing import Optional import yaml from langchain_core.prompts import ChatPromptTemplate @@ -27,7 +24,7 @@ class PromptLoader: """从 YAML 文件加载 Prompt 模板""" - def __init__(self, prompts_dir: Optional[str] = None): + def __init__(self, prompts_dir: str | None = None): if prompts_dir is None: # 默认路径:src/prompts/ prompts_dir = os.path.join( @@ -69,20 +66,22 @@ def load(self, name: str, version: str = "v1") -> ChatPromptTemplate: path = self._get_prompt_path(name, version) - with open(path, "r", encoding="utf-8") as f: + with open(path, encoding="utf-8") as f: data = yaml.safe_load(f) system = data.get("system", "") human = data.get("human", "") - template = ChatPromptTemplate.from_messages([ - ("system", system.strip()), - ("human", human.strip()), - ]) + template = ChatPromptTemplate.from_messages( + [ + ("system", system.strip()), + ("human", human.strip()), + ] + ) self._cache[cache_key] = template logger.info( - f"prompt_loaded", + "prompt_loaded", name=name, version=version, path=str(path), @@ -92,7 +91,7 @@ def load(self, name: str, version: str = "v1") -> ChatPromptTemplate: def load_raw(self, name: str, version: str = "v1") -> dict: """加载原始 YAML 数据(不包装为 ChatPromptTemplate)""" path = self._get_prompt_path(name, version) - with open(path, "r", encoding="utf-8") as f: + with open(path, encoding="utf-8") as f: return yaml.safe_load(f) def list_prompts(self) -> list[dict]: @@ -103,20 +102,22 @@ def list_prompts(self) -> list[dict]: for path in sorted(self.prompts_dir.glob("*.yaml")): try: - with open(path, "r", encoding="utf-8") as f: + with open(path, encoding="utf-8") as f: data = yaml.safe_load(f) - prompts.append({ - "name": data.get("name", path.stem), - "version": data.get("version", "unknown"), - "description": data.get("description", ""), - "file": path.name, - }) + prompts.append( + { + "name": data.get("name", path.stem), + "version": data.get("version", "unknown"), + "description": data.get("description", ""), + "file": path.name, + } + ) except Exception: continue return prompts - def invalidate_cache(self, name: Optional[str] = None): + def invalidate_cache(self, name: str | None = None): """清除缓存(用于热重载)""" if name is None: self._cache.clear() @@ -133,7 +134,7 @@ def reload(self, name: str, version: str = "v1") -> ChatPromptTemplate: # ── 单例 ────────────────────────────────────────── -_global_loader: Optional[PromptLoader] = None +_global_loader: PromptLoader | None = None def get_prompt_loader() -> PromptLoader: diff --git a/src/core/query_utils.py b/src/core/query_utils.py index 1b089f4..a74379a 100644 --- a/src/core/query_utils.py +++ b/src/core/query_utils.py @@ -4,19 +4,20 @@ 项目内 graph.py / qa.py / qa_stream.py 共用。 统一去除噪音 token、检测品牌对比意图。 """ + import re # 噪声:4位以内纯数字 token -_NOISE_RE = re.compile(r'(? str: """清洗查询:去纯数字噪音 + 压缩空白""" - cleaned = _NOISE_RE.sub('', query) - cleaned = re.sub(r'\s{2,}', ' ', cleaned).strip() + cleaned = _NOISE_RE.sub("", query) + cleaned = re.sub(r"\s{2,}", " ", cleaned).strip() return cleaned if cleaned else query diff --git a/src/core/state.py b/src/core/state.py index 6e125c6..9a6e4a6 100644 --- a/src/core/state.py +++ b/src/core/state.py @@ -14,10 +14,11 @@ await state.initialize() await state.rebuild_indexes() """ -import os + import asyncio +import os +from collections.abc import Callable from dataclasses import dataclass, field -from typing import Optional, Callable from src.config import settings from src.logger import logger @@ -35,7 +36,7 @@ class AppState: chunks: list = field(default_factory=list) bm25: any = None hybrid_retriever: any = None - bm25_search: Optional[Callable] = None + bm25_search: Callable | None = None reranker: any = None # --- 路径 --- @@ -45,22 +46,24 @@ class AppState: # --- 生命周期 --- _lock: asyncio.Lock = field(default_factory=asyncio.Lock) _initialized: bool = False - _error: Optional[str] = None + _error: str | None = None _use_pg: bool = False # True=PG+pgvector, False=ChromaDB # --- 统计 --- - stats: dict = field(default_factory=lambda: { - "categories": [], - "total_notes": 0, - "total_chunks": 0, - }) + stats: dict = field( + default_factory=lambda: { + "categories": [], + "total_notes": 0, + "total_chunks": 0, + } + ) @property def is_ready(self) -> bool: return self._initialized and not self._error @property - def error(self) -> Optional[str]: + def error(self) -> str | None: return self._error async def initialize(self) -> None: @@ -81,16 +84,18 @@ async def initialize(self) -> None: async def rebuild_indexes(self) -> None: """增量入库 + 重建 BM25/Hybrid/Graph 索引""" async with self._lock: + import jieba + from rank_bm25 import BM25Okapi + from src.ingestion import incremental_ingest, rebuild_all_chunks from src.retrievers import HybridRetriever, PgHybridRetriever - from rank_bm25 import BM25Okapi - import jieba logger.info("[AppState] 重建全量索引...") if self._use_pg: # PG 增量入库 from src.ingestion import incremental_ingest_to_pg + await incremental_ingest_to_pg(self.raw_dir) else: # ChromaDB 增量入库 @@ -108,7 +113,10 @@ async def rebuild_indexes(self) -> None: def bms(q, k=3): scores = bm25.get_scores(list(jieba.cut(q))) - return [chunks[i] for i in sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)[:k]] + return [ + chunks[i] + for i in sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)[:k] + ] self.chunks = chunks self.bm25 = bm25 @@ -119,10 +127,16 @@ def bms(q, k=3): async def _do_initialize(self) -> None: """实际初始化逻辑:优先 PG,回退 ChromaDB""" - from src.ingestion import rebuild_all_chunks, load_vectorstore - from src.retrievers import HybridRetriever, APIReranker, create_pg_vectorstore, PgHybridRetriever - from rank_bm25 import BM25Okapi import jieba + from rank_bm25 import BM25Okapi + + from src.ingestion import load_vectorstore, rebuild_all_chunks + from src.retrievers import ( + APIReranker, + HybridRetriever, + PgHybridRetriever, + create_pg_vectorstore, + ) project_root = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) self.raw_dir = os.path.join(project_root, "data", "raw") @@ -130,8 +144,11 @@ async def _do_initialize(self) -> None: loop = asyncio.get_running_loop() - raw_files = [f for f in os.listdir(self.raw_dir) if f.endswith((".txt", ".md"))] \ - if os.path.exists(self.raw_dir) else [] + raw_files = ( + [f for f in os.listdir(self.raw_dir) if f.endswith((".txt", ".md"))] + if os.path.exists(self.raw_dir) + else [] + ) if not raw_files: self._error = "暂无数据,请用 generate_data.py 生成数据后刷新" return @@ -153,6 +170,7 @@ async def _do_initialize(self) -> None: # PG 可用但为空,写入数据 logger.info("pg_empty, ingesting...") from src.ingestion import ingest_to_pg + await ingest_to_pg(self.raw_dir) pg_store = await create_pg_vectorstore() if pg_store: @@ -170,10 +188,13 @@ async def _do_initialize(self) -> None: if os.path.exists(chroma_db_file): vectorstore = await loop.run_in_executor(None, load_vectorstore) else: - from src.ingestion import load_raw_documents, chunk_documents, build_vectorstore + from src.ingestion import build_vectorstore, chunk_documents, load_raw_documents + docs = load_raw_documents() chunks_for_build = chunk_documents(docs) - vectorstore = await loop.run_in_executor(None, lambda: build_vectorstore(chunks_for_build)) + vectorstore = await loop.run_in_executor( + None, lambda: build_vectorstore(chunks_for_build) + ) logger.info("using_chromadb_vectorstore") # ===== 构建 BM25 + Hybrid + Graph ===== @@ -183,7 +204,7 @@ async def _do_initialize(self) -> None: if self._use_pg: hr = PgHybridRetriever(vectorstore, chunks) # PG 版 else: - hr = HybridRetriever(vectorstore, chunks) # ChromaDB 版 + hr = HybridRetriever(vectorstore, chunks) # ChromaDB 版 def bm25_search(query: str, k: int = 3): tokenized_query = list(jieba.cut(query)) @@ -200,8 +221,11 @@ def bm25_search(query: str, k: int = 3): def _refresh_stats(self) -> None: categories = list(set(d.metadata.get("category", "未分类") for d in self.chunks)) - raw_files = [f for f in os.listdir(self.raw_dir) if f.endswith((".txt", ".md"))] \ - if os.path.exists(self.raw_dir) else [] + raw_files = ( + [f for f in os.listdir(self.raw_dir) if f.endswith((".txt", ".md"))] + if os.path.exists(self.raw_dir) + else [] + ) self.stats = { "categories": categories, "total_notes": len(raw_files), @@ -238,4 +262,3 @@ async def init_app_state() -> AppState: state = AppState() await state.initialize() return state - diff --git a/src/crawler.py b/src/crawler.py index c5e0b3c..25c7b03 100644 --- a/src/crawler.py +++ b/src/crawler.py @@ -4,14 +4,12 @@ 基于 DrissionPage 的真实小红书爬虫封装。 当知识库中没有用户查询的品类数据时,自动打开浏览器抓取真实笔记和评论。 """ -import os -import json -import random # ============================================================ # 统一入口 # ============================================================ + class CrawlerInterface: """ 爬虫统一接口,支持本地和云端(Streamlit Cloud)两种模式。 @@ -27,6 +25,7 @@ def __init__(self, raw_dir: str, cookies_json: str = ""): try: from src.real_crawler import XHSCrawler + self._crawler = XHSCrawler(cookies_json=cookies_json) self._is_cloud = self._crawler.is_cloud_mode if self._crawler.is_logged_in: @@ -36,7 +35,10 @@ def __init__(self, raw_dir: str, cookies_json: str = ""): if self._is_cloud: print("[Crawler] [Cloud] 云端未登录。请在 Streamlit Secrets 中配置 XHS_COOKIES") else: - print("[Crawler] [WARN] 本地未登录。运行: uv run python src/real_crawler.py \"品类名\" 登录") + print( + "[Crawler] [WARN] 本地未登录。运行: " + 'uv run python src/real_crawler.py "品类名" 登录' + ) except Exception as e: self._init_error = str(e) print(f"[Crawler] [FAIL] 爬虫初始化失败: {e}") @@ -91,15 +93,19 @@ def crawl(self, category: str, count: int = 30, comment_level: str = "all") -> d return { "method": "error", "count": 0, - "error": "未登录。请在本地运行 `uv run python scripts/export_cookies.py` 导出 cookie,然后粘贴到 Streamlit Secrets → XHS_COOKIES" + "error": "未登录。请在本地运行 `uv run python scripts/export_cookies.py` " + "导出 cookie,然后粘贴到 Streamlit Secrets → XHS_COOKIES", } return { "method": "error", "count": 0, - "error": "未登录小红书。请先在命令行运行: uv run python src/real_crawler.py \"品类名\" 完成登录" + "error": "未登录小红书。请先在命令行运行: " + 'uv run python src/real_crawler.py "品类名" 完成登录', } - saved = self._crawler.crawl(category, count=count, with_comments=True, comment_level=comment_level) + saved = self._crawler.crawl( + category, count=count, with_comments=True, comment_level=comment_level + ) return { "method": "real", "count": saved, @@ -116,10 +122,12 @@ def fetch_hot_list(self, max_items: int = 30) -> list[dict]: if not self._crawler: print(f"[Crawler] 爬虫不可用,使用兜底热榜: {self._init_error}") from src.real_crawler import XHSCrawler + return XHSCrawler._fallback_hot_list() try: return self._crawler.fetch_hot_search(max_items=max_items) except Exception as e: print(f"[Crawler] 热榜抓取失败: {e}") from src.real_crawler import XHSCrawler + return XHSCrawler._fallback_hot_list() diff --git a/src/data/inspiration.py b/src/data/inspiration.py index 94ff053..47ff1a7 100644 --- a/src/data/inspiration.py +++ b/src/data/inspiration.py @@ -43,7 +43,11 @@ {"keyword": "速食面", "type": "both", "tip": "囤货必选,选题:泡面天花板是哪款"}, {"keyword": "牛肉干", "type": "selection", "tip": "高蛋白零食溢价高,选品:内蒙古产地直发"}, {"keyword": "果酒", "type": "both", "tip": "女性微醺经济,选题:10元以内好喝果酒"}, - {"keyword": "代餐奶昔", "type": "both", "tip": "减脂赛道火爆,选题:喝了一周代餐的真实感受"}, + { + "keyword": "代餐奶昔", + "type": "both", + "tip": "减脂赛道火爆,选题:喝了一周代餐的真实感受", + }, {"keyword": "火锅底料", "type": "both", "tip": "居家火锅趋势,选题:自制火锅比店里香"}, {"keyword": "蛋黄酥", "type": "both", "tip": "烘焙零食爆品,选品方向:手工vs工厂代工"}, {"keyword": "黑芝麻丸", "type": "both", "tip": "养生零食化趋势,选题:黑芝麻丸真的养发吗"}, @@ -80,13 +84,21 @@ {"keyword": "瑜伽裤", "type": "both", "tip": "lululemon平替方向,选题:10条平价瑜伽裤横评"}, {"keyword": "冲锋衣", "type": "both", "tip": "功能性溢价高,选品方向:城市户外风"}, {"keyword": "风衣", "type": "topic", "tip": "经典款永不过时,选题:不同身材风衣穿搭公式"}, - {"keyword": "显瘦穿搭", "type": "topic", "tip": "搜索量巨大,选题:苹果型/梨型身材穿搭避雷"}, + { + "keyword": "显瘦穿搭", + "type": "topic", + "tip": "搜索量巨大,选题:苹果型/梨型身材穿搭避雷", + }, {"keyword": "羽绒服", "type": "both", "tip": "冬季大单品高客单,选题:500元以下羽绒服推荐"}, {"keyword": "卫衣", "type": "both", "tip": "春秋必备基础款,选题:一周卫衣不重样穿搭"}, {"keyword": "衬衫", "type": "both", "tip": "通勤刚需,选题:白衬衫的100种穿法"}, {"keyword": "牛仔裤", "type": "both", "tip": "万能单品,选题:不同腿型牛仔裤选购指南"}, {"keyword": "连衣裙", "type": "both", "tip": "夏季主力品类,选品:法式碎花裙利润分析"}, - {"keyword": "西装外套", "type": "both", "tip": "职场穿搭升级,选题:平价西装也能穿出高级感"}, + { + "keyword": "西装外套", + "type": "both", + "tip": "职场穿搭升级,选题:平价西装也能穿出高级感", + }, {"keyword": "针织衫", "type": "both", "tip": "秋冬百搭单品,选品:不起球针织衫怎么挑"}, {"keyword": "半身裙", "type": "topic", "tip": "穿搭包容性强,选题:微胖女生半身裙推荐"}, {"keyword": "打底衫", "type": "both", "tip": "基础款高频复购,选品方向:自发热打底衫"}, @@ -99,7 +111,11 @@ {"keyword": "情侣装", "type": "both", "tip": "情感消费溢价高,选题:不土的情侣穿搭"}, ], "数码": [ - {"keyword": "蓝牙耳机", "type": "both", "tip": "竞争激烈但利润稳,选题:百元以下耳机谁最强"}, + { + "keyword": "蓝牙耳机", + "type": "both", + "tip": "竞争激烈但利润稳,选题:百元以下耳机谁最强", + }, {"keyword": "手机壳", "type": "both", "tip": "门槛低利润高,选品方向:设计师联名+季节性"}, {"keyword": "充电宝", "type": "topic", "tip": "实用测评类高收藏,选题:10款充电宝实测"}, {"keyword": "无线鼠标", "type": "both", "tip": "办公刚需,选品方向:静音+人体工学溢价"}, @@ -107,7 +123,11 @@ {"keyword": "键盘", "type": "both", "tip": "客制化赛道火热,选题:机械键盘入门指南"}, {"keyword": "数据线", "type": "selection", "tip": "高频消耗品,选品:快充数据线走量策略"}, {"keyword": "平板支架", "type": "both", "tip": "无纸化学习趋势,选题:iPad配件好物推荐"}, - {"keyword": "手机膜", "type": "selection", "tip": "极致走量品类,选品方向:防窥/蓝光功能膜"}, + { + "keyword": "手机膜", + "type": "selection", + "tip": "极致走量品类,选品方向:防窥/蓝光功能膜", + }, {"keyword": "蓝牙音箱", "type": "both", "tip": "氛围感好物,选题:百元蓝牙音箱音质横评"}, {"keyword": "智能手表", "type": "both", "tip": "健康监测刚需,选题:手环VS手表怎么选"}, {"keyword": "耳机壳", "type": "selection", "tip": "AirPods配件刚需,选品:硅胶耳机壳利润"}, @@ -116,16 +136,36 @@ {"keyword": "屏幕挂灯", "type": "both", "tip": "程序员/设计师刚需,选题:屏幕挂灯值得吗"}, {"keyword": "充电头", "type": "selection", "tip": "快充刚需配件,选品:氮化镓充电器利润"}, {"keyword": "拓展坞", "type": "both", "tip": "轻薄本必备配件,选题:Type-C拓展坞推荐"}, - {"keyword": "平板保护套", "type": "selection", "tip": "iPad必备配件,选品:带笔槽保护套溢价"}, + { + "keyword": "平板保护套", + "type": "selection", + "tip": "iPad必备配件,选品:带笔槽保护套溢价", + }, {"keyword": "车载充电器", "type": "selection", "tip": "汽车配件刚需,选品:双口快充车充"}, {"keyword": "桌面理线器", "type": "both", "tip": "桌面美学趋势,选题:告别桌面乱线的神器"}, - {"keyword": "读卡器", "type": "selection", "tip": "摄影/Vlog必备,选品方向:多合一高速读卡器"}, + { + "keyword": "读卡器", + "type": "selection", + "tip": "摄影/Vlog必备,选品方向:多合一高速读卡器", + }, ], "运动": [ - {"keyword": "健身器材", "type": "both", "tip": "客单价高、利润厚,选品方向:居家小型化趋势"}, + { + "keyword": "健身器材", + "type": "both", + "tip": "客单价高、利润厚,选品方向:居家小型化趋势", + }, {"keyword": "瑜伽垫", "type": "both", "tip": "入门必备,选题:用了一年最不后悔的瑜伽垫"}, - {"keyword": "筋膜枪", "type": "both", "tip": "恢复类热门单品,选题:百元筋膜枪VS千元款实测"}, - {"keyword": "运动内衣", "type": "both", "tip": "女性刚需高频,选题:不同胸型运动内衣怎么选"}, + { + "keyword": "筋膜枪", + "type": "both", + "tip": "恢复类热门单品,选题:百元筋膜枪VS千元款实测", + }, + { + "keyword": "运动内衣", + "type": "both", + "tip": "女性刚需高频,选题:不同胸型运动内衣怎么选", + }, {"keyword": "哑铃", "type": "both", "tip": "居家健身标配,选品方向:可调节款溢价空间大"}, {"keyword": "跑步鞋", "type": "topic", "tip": "永不过时,选题:新手第一双跑鞋怎么选"}, {"keyword": "泡沫轴", "type": "both", "tip": "运动恢复赛道,选题:泡沫轴是智商税吗"}, @@ -136,7 +176,11 @@ {"keyword": "运动水壶", "type": "both", "tip": "健身人群刚需,选题:大容量运动水壶测评"}, {"keyword": "护膝", "type": "both", "tip": "运动防护品类,选题:跑步护膝有没有用"}, {"keyword": "运动毛巾", "type": "both", "tip": "冷感毛巾夏季爆品,选品方向:速干面料"}, - {"keyword": "登山杖", "type": "both", "tip": "户外登山装备,选题:登山杖选碳纤维还是铝合金"}, + { + "keyword": "登山杖", + "type": "both", + "tip": "户外登山装备,选题:登山杖选碳纤维还是铝合金", + }, {"keyword": "滑板", "type": "both", "tip": "潮流运动品类,选题:新手滑板选购指南"}, {"keyword": "泳衣", "type": "both", "tip": "季节性爆品,选题:不同身材泳衣怎么选"}, {"keyword": "运动背包", "type": "both", "tip": "健身通勤双场景,选品:干湿分离运动包"}, @@ -146,22 +190,46 @@ ], "宠物": [ {"keyword": "猫粮", "type": "both", "tip": "刚需高频复购,选题:成分党教你一眼看懂配料表"}, - {"keyword": "宠物玩具", "type": "both", "tip": "冲动消费多、利润高,选品:互动型玩具溢价空间"}, + { + "keyword": "宠物玩具", + "type": "both", + "tip": "冲动消费多、利润高,选品:互动型玩具溢价空间", + }, {"keyword": "猫砂", "type": "topic", "tip": "消耗品测评易火,选题:10款猫砂粉尘对比"}, {"keyword": "狗粮", "type": "both", "tip": "最大刚需品类,选题:国产狗粮真的不如进口吗"}, {"keyword": "宠物零食", "type": "both", "tip": "高频复购高毛利,选品:冻干零食工厂代工"}, {"keyword": "猫爬架", "type": "both", "tip": "大件宠物用品,选品:实木vs剑麻材质利润对比"}, {"keyword": "宠物窝", "type": "both", "tip": "颜值+功能双需求,选题:四季通用宠物窝推荐"}, - {"keyword": "宠物自动喂食器", "type": "both", "tip": "智能养宠趋势,选题:自动喂食器出粮实测"}, + { + "keyword": "宠物自动喂食器", + "type": "both", + "tip": "智能养宠趋势,选题:自动喂食器出粮实测", + }, {"keyword": "宠物梳子", "type": "both", "tip": "换毛季刚需爆品,选题:去浮毛梳子测评"}, - {"keyword": "宠物衣服", "type": "both", "tip": "宠物拟人化消费,选品:秋冬季宠物服装利润率"}, + { + "keyword": "宠物衣服", + "type": "both", + "tip": "宠物拟人化消费,选品:秋冬季宠物服装利润率", + }, {"keyword": "猫罐头", "type": "both", "tip": "主食罐VS零食罐,选题:猫咪最爱的罐头排名"}, - {"keyword": "宠物尿垫", "type": "selection", "tip": "高频消耗品走量,选品:加厚竹炭除臭尿垫"}, + { + "keyword": "宠物尿垫", + "type": "selection", + "tip": "高频消耗品走量,选品:加厚竹炭除臭尿垫", + }, {"keyword": "宠物航空箱", "type": "both", "tip": "宠物出行必备,选题:猫咪外出包怎么选"}, - {"keyword": "宠物饮水机", "type": "both", "tip": "活水概念升级,选题:宠物饮水机过滤效果对比"}, + { + "keyword": "宠物饮水机", + "type": "both", + "tip": "活水概念升级,选题:宠物饮水机过滤效果对比", + }, {"keyword": "宠物牵引绳", "type": "both", "tip": "遛狗刚需,选题:爆冲狗该用什么牵引绳"}, {"keyword": "猫抓板", "type": "both", "tip": "高频消耗品类,选品:瓦楞纸猫抓板走量"}, - {"keyword": "宠物湿巾", "type": "both", "tip": "清洁刚需消耗品,选题:宠物专用vs婴儿湿巾对比"}, + { + "keyword": "宠物湿巾", + "type": "both", + "tip": "清洁刚需消耗品,选题:宠物专用vs婴儿湿巾对比", + }, {"keyword": "鱼粮", "type": "both", "tip": "水族垂直品类,选题:金鱼热带鱼饲料怎么选"}, {"keyword": "仓鼠笼", "type": "both", "tip": "小宠赛道蓝海,选品:亚克力仓鼠笼溢价空间"}, {"keyword": "宠物钙片", "type": "both", "tip": "宠物保健品市场,选题:狗狗真的需要补钙吗"}, @@ -169,7 +237,11 @@ ], "母婴": [ {"keyword": "婴儿湿巾", "type": "both", "tip": "消耗品高复购,选品方向:天然成分溢价"}, - {"keyword": "儿童水杯", "type": "both", "tip": "安全诉求溢价高,选题:不同材质儿童水杯对比"}, + { + "keyword": "儿童水杯", + "type": "both", + "tip": "安全诉求溢价高,选题:不同材质儿童水杯对比", + }, {"keyword": "早教玩具", "type": "topic", "tip": "知识类高收藏,选题:分月龄早教玩具推荐"}, {"keyword": "纸尿裤", "type": "both", "tip": "母婴第一刚需品,选题:10款纸尿裤吸水性实测"}, {"keyword": "婴儿推车", "type": "both", "tip": "高客单价大单品,选题:千元推车测评"}, @@ -229,15 +301,17 @@ def get_inspiration(category: str = None) -> list[dict]: for cat in categories: items = INSPIRATION.get(cat, []) for idx, item in enumerate(items): - result.append({ - "keyword": item["keyword"], - "type": item["type"], - "tip": item["tip"], - "category": cat, - "rank": idx + 1, - "hots": max(95 - idx * 4, 30), - "trend": "up" if idx < 3 else "stable", - }) + result.append( + { + "keyword": item["keyword"], + "type": item["type"], + "tip": item["tip"], + "category": cat, + "rank": idx + 1, + "hots": max(95 - idx * 4, 30), + "trend": "up" if idx < 3 else "stable", + } + ) result.sort(key=lambda x: -x["hots"]) for i, r in enumerate(result): diff --git a/src/ingestion.py b/src/ingestion.py index d617a13..8082086 100644 --- a/src/ingestion.py +++ b/src/ingestion.py @@ -2,23 +2,27 @@ ingestion.py - 文档加载 + 向量库构建 合并自原 step01_ingestion + step02_vectorstore """ + import os -from langchain_community.document_loaders import TextLoader, DirectoryLoader -from langchain_text_splitters import RecursiveCharacterTextSplitter -from langchain_openai import OpenAIEmbeddings + from langchain_chroma import Chroma +from langchain_community.document_loaders import DirectoryLoader, TextLoader from langchain_core.documents import Document +from langchain_openai import OpenAIEmbeddings +from langchain_text_splitters import RecursiveCharacterTextSplitter -from src.config import RAW_DIR, CHROMA_DIR, EMBEDDING_CONFIG +from src.config import CHROMA_DIR, EMBEDDING_CONFIG, RAW_DIR from src.logger import logger - # ===== 文档清洗 ===== + def _extract_fm_prefix(text: str) -> str: """从 frontmatter 中提取品牌/价格关键信息,返回元数据前缀文本""" import re + import yaml as _yaml + fm_match = re.match(r"^---\s*\n(.*?)\n---", text, re.DOTALL) if not fm_match: return "" @@ -43,6 +47,7 @@ def _extract_fm_prefix(text: str) -> str: def _clean_frontmatter(text: str) -> str: """统一的 frontmatter + HTML 注释清理,保留关键元数据到文本前缀""" import re + prefix = _extract_fm_prefix(text) text = re.sub(r"^---\s*\n.*?\n---\s*\n", "", text, count=1, flags=re.DOTALL) text = re.sub(r"", "", text, flags=re.DOTALL) @@ -132,17 +137,19 @@ def incremental_ingest(raw_dir: str, vectorstore: Chroma) -> list: 用法: new_chunks = incremental_ingest(str(RAW_DIR), vectorstore) """ - from langchain_community.document_loaders import TextLoader, DirectoryLoader + from langchain_community.document_loaders import DirectoryLoader, TextLoader # 1. 加载所有 .md 文件(包括已有的,Chromadb 内置去重) loader = DirectoryLoader( - str(raw_dir), glob="**/*.md", + str(raw_dir), + glob="**/*.md", loader_cls=TextLoader, loader_kwargs={"encoding": "utf-8"}, show_progress=False, ) loader_txt = DirectoryLoader( - str(raw_dir), glob="**/*.txt", + str(raw_dir), + glob="**/*.txt", loader_cls=TextLoader, loader_kwargs={"encoding": "utf-8"}, show_progress=False, @@ -170,16 +177,18 @@ def rebuild_all_chunks(raw_dir: str) -> list: 重新加载全部文档并 chunk,用于重建 BM25 索引。 返回完整的 chunks 列表。 """ - from langchain_community.document_loaders import TextLoader, DirectoryLoader + from langchain_community.document_loaders import DirectoryLoader, TextLoader loader = DirectoryLoader( - str(raw_dir), glob="**/*.md", + str(raw_dir), + glob="**/*.md", loader_cls=TextLoader, loader_kwargs={"encoding": "utf-8"}, show_progress=False, ) loader_txt = DirectoryLoader( - str(raw_dir), glob="**/*.txt", + str(raw_dir), + glob="**/*.txt", loader_cls=TextLoader, loader_kwargs={"encoding": "utf-8"}, show_progress=False, @@ -196,14 +205,16 @@ def rebuild_all_chunks(raw_dir: str) -> list: # PostgreSQL + pgvector 异步 Ingestion # ================================================================ + async def ingest_to_pg(raw_dir: str = None) -> int: """将所有文档 embedding 后写入 PostgreSQL + pgvector 异步版本,直接写入 PG 替代 ChromaDB。 返回写入的文档数。 """ - from src.core.database import get_db, insert_documents, init_db - from langchain_community.document_loaders import TextLoader, DirectoryLoader + from langchain_community.document_loaders import DirectoryLoader, TextLoader + + from src.core.database import get_db, init_db, insert_documents if raw_dir is None: raw_dir = str(RAW_DIR) @@ -213,13 +224,15 @@ async def ingest_to_pg(raw_dir: str = None) -> int: # 2. 加载文档 loader = DirectoryLoader( - raw_dir, glob="**/*.md", + raw_dir, + glob="**/*.md", loader_cls=TextLoader, loader_kwargs={"encoding": "utf-8"}, show_progress=False, ) loader_txt = DirectoryLoader( - raw_dir, glob="**/*.txt", + raw_dir, + glob="**/*.txt", loader_cls=TextLoader, loader_kwargs={"encoding": "utf-8"}, show_progress=False, @@ -238,16 +251,17 @@ async def ingest_to_pg(raw_dir: str = None) -> int: logger.info(f"embedding {len(texts)} chunks...") # 批量 embedding(每次最多 100 条,避免 API 超限) - BATCH_SIZE = 100 + batch_size = 100 total_inserted = 0 async for session in get_db(): - for i in range(0, len(texts), BATCH_SIZE): - batch_texts = texts[i:i + BATCH_SIZE] - batch_chunks = chunks[i:i + BATCH_SIZE] + for i in range(0, len(texts), batch_size): + batch_texts = texts[i : i + batch_size] + batch_chunks = chunks[i : i + batch_size] # embed 在 event loop 外执行(OAI embedding 是同步的) import asyncio + loop = asyncio.get_running_loop() batch_embeddings = await loop.run_in_executor( None, lambda: embeddings_model.embed_documents(batch_texts) @@ -266,9 +280,10 @@ async def incremental_ingest_to_pg(raw_dir: str = None) -> int: 返回新插入的文档数。 """ - from src.core.database import get_db, insert_documents, init_db, DocumentTable + from langchain_community.document_loaders import DirectoryLoader, TextLoader from sqlalchemy import select - from langchain_community.document_loaders import TextLoader, DirectoryLoader + + from src.core.database import DocumentTable, get_db, init_db, insert_documents if raw_dir is None: raw_dir = str(RAW_DIR) @@ -285,13 +300,15 @@ async def incremental_ingest_to_pg(raw_dir: str = None) -> int: # 2. 加载新文档 loader = DirectoryLoader( - raw_dir, glob="**/*.md", + raw_dir, + glob="**/*.md", loader_cls=TextLoader, loader_kwargs={"encoding": "utf-8"}, show_progress=False, ) loader_txt = DirectoryLoader( - raw_dir, glob="**/*.txt", + raw_dir, + glob="**/*.txt", loader_cls=TextLoader, loader_kwargs={"encoding": "utf-8"}, show_progress=False, @@ -313,15 +330,16 @@ async def incremental_ingest_to_pg(raw_dir: str = None) -> int: embeddings_model = get_embeddings() texts = [chunk.page_content for chunk in chunks] - BATCH_SIZE = 100 + batch_size = 100 total_inserted = 0 async for session in get_db(): - for i in range(0, len(texts), BATCH_SIZE): - batch_texts = texts[i:i + BATCH_SIZE] - batch_chunks = chunks[i:i + BATCH_SIZE] + for i in range(0, len(texts), batch_size): + batch_texts = texts[i : i + batch_size] + batch_chunks = chunks[i : i + batch_size] import asyncio + loop = asyncio.get_running_loop() batch_embeddings = await loop.run_in_executor( None, lambda: embeddings_model.embed_documents(batch_texts) diff --git a/src/logger.py b/src/logger.py index f54cab6..577cdd6 100644 --- a/src/logger.py +++ b/src/logger.py @@ -1,5 +1,7 @@ -import structlog import logging + +import structlog + from src.config import settings @@ -39,4 +41,4 @@ def setup_logging() -> structlog.BoundLogger: # ── 模块级 logger 实例 ──────────────────────────── -logger = setup_logging() \ No newline at end of file +logger = setup_logging() diff --git a/src/real_crawler.py b/src/real_crawler.py index de0ac00..6d3dace 100644 --- a/src/real_crawler.py +++ b/src/real_crawler.py @@ -9,25 +9,27 @@ 首次运行会打开浏览器窗口,请手动扫码登录。登录后 cookie 会保存, 后续运行自动复用,无需重复登录。 """ -import os -import sys + +import argparse import json -import time +import os import random import re -import argparse -from pathlib import Path +import sys +import time sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from DrissionPage import ChromiumPage, ChromiumOptions -from src.config import RAW_DIR +from DrissionPage import ChromiumOptions, ChromiumPage +from src.config import RAW_DIR # ============================================================ # 配置 # ============================================================ -COOKIE_FILE = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "data", "cookies.json") +COOKIE_FILE = os.path.join( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "data", "cookies.json" +) SEARCH_URL = "https://www.xiaohongshu.com/search_result?keyword={}&source=web_search_result_notes" @@ -44,15 +46,16 @@ def __init__(self, headless: bool = False, cookies_json: str = ""): self.cookies_json = cookies_json # 云端模式:从 secrets 传入的 cookie JSON self._logged_in = False self._is_cloud = self._detect_cloud() - self._need_relogin = False # search() 发现登录弹窗时设为 True + self._need_relogin = False # search() 发现登录弹窗时设为 True self._init_browser() @staticmethod def _detect_cloud() -> bool: """检测是否运行在 Streamlit Cloud 环境""" # Streamlit Cloud 设置了 STREAMLIT_SERVER_ADDRESS 环境变量 - return bool(os.environ.get("STREAMLIT_SERVER_ADDRESS")) or \ - bool(os.environ.get("STREAMLIT_RUNTIME")) + return bool(os.environ.get("STREAMLIT_SERVER_ADDRESS")) or bool( + os.environ.get("STREAMLIT_RUNTIME") + ) def _init_browser(self): """初始化浏览器,尝试复用已保存的登录态""" @@ -64,19 +67,28 @@ def _init_browser(self): co.set_argument("--disable-features=VizDisplayCompositor") co.set_argument("--window-size=1920,1080") # 随机化 User-Agent,避免被识别为爬虫 - ua = random.choice([ - "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/125.0.0.0 Safari/537.36", - "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/126.0.0.0 Safari/537.36", - "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/124.0.0.0 Safari/537.36", - ]) + ua = random.choice( + [ + "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 " + "(KHTML, like Gecko) Chrome/125.0.0.0 Safari/537.36", + "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 " + "(KHTML, like Gecko) Chrome/126.0.0.0 Safari/537.36", + "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 " + "(KHTML, like Gecko) Chrome/124.0.0.0 Safari/537.36", + ] + ) co.set_user_agent(ua) if self._is_cloud or self.headless: co.headless(True) # Streamlit Cloud 上 Chromium 的安装路径 if self._is_cloud: - for browser_path in ["/usr/bin/chromium-browser", "/usr/bin/chromium", - "/usr/bin/google-chrome", "/usr/bin/chrome"]: + for browser_path in [ + "/usr/bin/chromium-browser", + "/usr/bin/chromium", + "/usr/bin/google-chrome", + "/usr/bin/chrome", + ]: if os.path.exists(browser_path): co.set_browser_path(browser_path) print(f"[Crawler] [Cloud] 云端模式,使用: {browser_path}") @@ -140,7 +152,7 @@ def _load_cookies(self): # 2. 本地模式:从文件加载 if not cookies and os.path.exists(COOKIE_FILE): try: - with open(COOKIE_FILE, "r", encoding="utf-8") as f: + with open(COOKIE_FILE, encoding="utf-8") as f: cookies = json.load(f) print(f"[Crawler] [File] 从文件加载了 {len(cookies)} 个 cookie") except Exception as e: @@ -177,9 +189,14 @@ def _load_cookies(self): self._logged_in = False if self._is_cloud: - print("[Crawler] [Cloud] 云端未登录。请在本地导出 cookie 并配置 Streamlit Secrets: XHS_COOKIES") + print( + "[Crawler] [Cloud] 云端未登录。请在本地导出 cookie " + "并配置 Streamlit Secrets: XHS_COOKIES" + ) else: - print("[Crawler] [WARN] 未登录。请运行: uv run python src/real_crawler.py \"品类名\" 来登录") + print( + '[Crawler] [WARN] 未登录。请运行: uv run python src/real_crawler.py "品类名" 来登录' + ) def _is_login_modal_shown(self) -> bool: """检查当前页面是否显示了登录弹窗(CSS 检测 + HTML 文本兜底)""" @@ -262,7 +279,15 @@ def _verify_logged_in(self) -> bool: # 方法1: 检查是否存在小红书认证 cookie xhs_cookies = self.page.cookies(all_domains=True, all_info=False) - auth_cookie_names = {"a1", "web_session", "session", "sid", "authorization", "token", "xhs"} + auth_cookie_names = { + "a1", + "web_session", + "session", + "sid", + "authorization", + "token", + "xhs", + } for cookie in xhs_cookies: name = cookie.get("name", "").lower() if name in auth_cookie_names: @@ -411,7 +436,10 @@ def _navigate_and_dismiss(target_url: str) -> bool: if login_still_shown: print("[Crawler] [Search] [WARN] 登录弹窗持续存在,请重新登录") - print("[Crawler] [Search] 已登录状态可能已过期,建议运行: uv run python src/real_crawler.py") + print( + "[Crawler] [Search] 已登录状态可能已过期,建议运行: " + "uv run python src/real_crawler.py" + ) # 即使弹窗还在,也尝试抓取(可能弹窗背后有内容) # ===== 策略 A (首选): 用浏览器 fetch() 直接调搜索 API ===== @@ -448,13 +476,15 @@ def _navigate_and_dismiss(target_url: str) -> bool: nid = item.get("id", "") if nid and nid not in seen_ids: seen_ids.add(nid) - notes.append({ - "id": nid, - "title": item.get("title", f"{keyword}_{nid[:8]}"), - "url": f"https://www.xiaohongshu.com/explore/{nid}", - "likes": item.get("likes", 0), - "author": item.get("author", ""), - }) + notes.append( + { + "id": nid, + "title": item.get("title", f"{keyword}_{nid[:8]}"), + "url": f"https://www.xiaohongshu.com/explore/{nid}", + "likes": item.get("likes", 0), + "author": item.get("author", ""), + } + ) print(f"[Crawler] [Search] [OK] API 搜索到 {len(notes)} 篇笔记") else: print(f"[Crawler] [Search] [WARN] API 返回异常: {str(api_notes_raw)[:100]}") @@ -465,7 +495,6 @@ def _navigate_and_dismiss(target_url: str) -> bool: if not notes: print("[Crawler] [Search] API 无结果,退回到页面解析...") no_new = 0 - last_count = 0 max_scrolls = max(count // 3, 30) scroll_count = 0 @@ -475,22 +504,22 @@ def _navigate_and_dismiss(target_url: str) -> bool: time.sleep(random.uniform(1.5, 3.0)) page_html = self.page.html or "" - hex_ids = re.findall(r'/explore/([a-f0-9]{24})', page_html) + hex_ids = re.findall(r"/explore/([a-f0-9]{24})", page_html) if hex_ids: seen_ids = {n["id"] for n in notes} for note_id in hex_ids: if note_id not in seen_ids: seen_ids.add(note_id) - notes.append({ - "id": note_id, - "title": f"{keyword}_{note_id[:8]}", - "url": f"https://www.xiaohongshu.com/explore/{note_id}", - }) + notes.append( + { + "id": note_id, + "title": f"{keyword}_{note_id[:8]}", + "url": f"https://www.xiaohongshu.com/explore/{note_id}", + } + ) no_new = 0 - last_count = len(notes) else: no_new += 1 - last_count = len(notes) print(f"\r[Crawler] 已发现 {len(notes)} 篇...", end="") print(f"\n[Crawler] [OK] 搜索完成,共 {len(notes[:count])} 篇笔记") @@ -498,7 +527,9 @@ def _navigate_and_dismiss(target_url: str) -> bool: # 如果一篇都没找到且有登录弹窗 → 需要重新登录 if not notes and self._is_login_modal_shown(): self._need_relogin = True - print("[Crawler] [Search] [WARN] 搜索到 0 篇笔记 + 登录弹窗 -> cookie 过期,需要重新登录") + print( + "[Crawler] [Search] [WARN] 搜索到 0 篇笔记 + 登录弹窗 -> cookie 过期,需要重新登录" + ) else: self._need_relogin = False @@ -509,8 +540,12 @@ def _is_note_page_blocked(self) -> bool: try: page_text = self.page.text or "" blocked_keywords = [ - "暂时无法浏览", "请打开小红书App扫码查看", "扫码查看", - "小红书如何扫码", "问题反馈", "返回首页", + "暂时无法浏览", + "请打开小红书App扫码查看", + "扫码查看", + "小红书如何扫码", + "问题反馈", + "返回首页", ] matches = sum(1 for kw in blocked_keywords if kw in page_text) return matches >= 3 # 匹配到 3 个以上关键词才判定为被拦截 @@ -544,12 +579,16 @@ def _try_fetch_via_api(self, note_id: str) -> dict: """ raw = self.page.run_js(js) import json - if raw and raw.startswith('{'): + + if raw and raw.startswith("{"): data = json.loads(raw) result.update(data) - print(f" [API] [OK] 获取内容成功({len(result['content'])} 字符,{result['likes']} 赞)") - elif raw and 'NO_DATA' in str(raw): - print(f" [API] [WARN] API 返回空数据") + print( + f" [API] [OK] 获取内容成功({len(result['content'])} 字符," + f"{result['likes']} 赞)" + ) + elif raw and "NO_DATA" in str(raw): + print(" [API] [WARN] API 返回空数据") else: print(f" [API] [WARN] 失败: {str(raw)[:80]}") except Exception as e: @@ -559,7 +598,7 @@ def _try_fetch_via_api(self, note_id: str) -> dict: def get_note_detail(self, note: dict) -> dict: """获取单篇笔记的正文内容 — API 优先""" - title_short = note.get('title', '')[:30] + title_short = note.get("title", "")[:30] note_id = note.get("id", "") print(f"[Crawler] [Page] {title_short}...") @@ -580,7 +619,7 @@ def get_note_detail(self, note: dict) -> dict: time.sleep(1) blocked = self._is_note_page_blocked() if blocked: - print(f" [WARN] 被拦截,放弃此篇") + print(" [WARN] 被拦截,放弃此篇") note["content"] = note.get("content", "") note["likes"] = note.get("likes", 0) note["author"] = note.get("author", "") @@ -588,8 +627,13 @@ def get_note_detail(self, note: dict) -> dict: # 获取正文 content = "" - for desc_sel in ["css:#detail-desc", "css:.note-text", "css:.desc", - "css:.note-scroller", "css:.content"]: + for desc_sel in [ + "css:#detail-desc", + "css:.note-text", + "css:.desc", + "css:.note-scroller", + "css:.content", + ]: try: desc_el = self.page.ele(desc_sel) if desc_el: @@ -601,8 +645,11 @@ def get_note_detail(self, note: dict) -> dict: likes = 0 try: - for like_sel in ["css:.like-wrapper .count", "css:.like .count", - "css:.interact-item .count"]: + for like_sel in [ + "css:.like-wrapper .count", + "css:.like .count", + "css:.interact-item .count", + ]: like_el = self.page.ele(like_sel) if like_el: likes = int(re.sub(r"\D", "", like_el.text) or "0") @@ -661,7 +708,7 @@ def _try_fetch_comments_via_api(self, note_id: str) -> list[str]: }})() """ raw = self.page.run_js(js) - if raw and raw.startswith('['): + if raw and raw.startswith("["): comments = json.loads(raw) print(f" [API] [OK] 评论 API 获取到 {len(comments)} 条") return comments @@ -690,7 +737,7 @@ def get_comments(self, note: dict, max_comments: int = 30, level: str = "all") - api_comments = self._try_fetch_comments_via_api(note_id) if api_comments: if level == "hot": - return api_comments[:max(5, max_comments // 3)] + return api_comments[: max(5, max_comments // 3)] elif level == "top": return api_comments[:max_comments] # API 返回的顶级在前 return api_comments[:max_comments] @@ -720,7 +767,7 @@ def get_comments(self, note: dict, max_comments: int = 30, level: str = "all") - "css:.comment-list > .comment-item > .content", "css:.comments-wrapper > .comment-item > .content", "css:.note-scroller > .comment-item > .content", - "css:.comment-item > .content", # 兜底:所有 comment-item 的直接子 content + "css:.comment-item > .content", # 兜底:所有 comment-item 的直接子 content ] elif level == "hot": # 热评:通常排在前几条,只取前 1/3 @@ -754,8 +801,10 @@ def get_comments(self, note: dict, max_comments: int = 30, level: str = "all") - if not comment_els: # 全量抓取 for sel in [ - "css:.comment-item .content", "css:.comment-content", - "css:.comments .content", "css:.note-comment .content", + "css:.comment-item .content", + "css:.comment-content", + "css:.comments .content", + "css:.note-comment .content", ]: try: found = self.page.eles(sel) @@ -776,14 +825,22 @@ def get_comments(self, note: dict, max_comments: int = 30, level: str = "all") - grandparent = parent.parent() if parent else None gp_classes = grandparent.attr("class") or "" if grandparent else "" # 如果父/祖父元素包含 reply/sub/child 关键词,跳过 - if any(k in parent_classes.lower() for k in ["reply", "sub", "child", "二级"]): + if any( + k in parent_classes.lower() + for k in ["reply", "sub", "child", "二级"] + ): continue - if any(k in gp_classes.lower() for k in ["reply", "sub", "child", "二级"]): + if any( + k in gp_classes.lower() for k in ["reply", "sub", "child", "二级"] + ): continue filtered.append(el) except Exception: filtered.append(el) - print(f"[Crawler] [Comments] DOM 过滤后: {len(filtered)}/{len(comment_els)} 条一级评论") + print( + "[Crawler] [Comments] DOM 过滤后: " + f"{len(filtered)}/{len(comment_els)} 条一级评论" + ) comment_els = filtered # 提取文本 @@ -797,9 +854,11 @@ def get_comments(self, note: dict, max_comments: int = 30, level: str = "all") - # 热评模式:只取前 1/3(按页面排序,热评在前) if level == "hot": - comments = comments[:max(5, max_comments // 3)] + comments = comments[: max(5, max_comments // 3)] - print(f"[Crawler] [Comments] 获取到 {len(comments[:max_comments])} 条评论 (level={level})") + print( + f"[Crawler] [Comments] 获取到 {len(comments[:max_comments])} 条评论 (level={level})" + ) except Exception as e: print(f" [WARN] 评论获取失败: {e}") @@ -813,8 +872,9 @@ def save_note(self, note: dict, category: str, comments: list[str] = None): import yaml # 分析评论 - complaints, purchase_intents, comparison_mentions, high_freq_words = \ - self._analyze_comments(comments or []) + complaints, purchase_intents, comparison_mentions, high_freq_words = self._analyze_comments( + comments or [] + ) # 构建 frontmatter(与 generate_data.py 格式一致) frontmatter = { @@ -840,8 +900,11 @@ def save_note(self, note: dict, category: str, comments: list[str] = None): "purchase_intent": purchase_intents[:10], "comparison_mentions": comparison_mentions[:5], "related_brands": [category], - "ask_link_count": sum(1 for c in (comments or []) - if any(kw in c for kw in ["链接", "在哪买", "求", "想要"])), + "ask_link_count": sum( + 1 + for c in (comments or []) + if any(kw in c for kw in ["链接", "在哪买", "求", "想要"]) + ), } ecommerce = { @@ -864,10 +927,14 @@ def save_note(self, note: dict, category: str, comments: list[str] = None): "", "---", "", ] @@ -888,17 +955,66 @@ def _analyze_comments(self, comments: list[str]) -> tuple: all_words = [] complaint_kw = [ - "太贵", "不好", "差", "烂", "后悔", "千万别", "踩雷", "坑", "不行", - "小", "短", "丑", "慢", "掉", "坏", "退", "缺点", "失望", "难用", - "不值", "垃圾", "无语", "鸡肋", "智商税", "浪费", "别买", "慎入", - "不好用", "有问题", "坏了", "掉了", "破了", "不推荐", "避雷", + "太贵", + "不好", + "差", + "烂", + "后悔", + "千万别", + "踩雷", + "坑", + "不行", + "小", + "短", + "丑", + "慢", + "掉", + "坏", + "退", + "缺点", + "失望", + "难用", + "不值", + "垃圾", + "无语", + "鸡肋", + "智商税", + "浪费", + "别买", + "慎入", + "不好用", + "有问题", + "坏了", + "掉了", + "破了", + "不推荐", + "避雷", ] intent_kw = [ - "想买", "求链接", "在哪买", "多少钱", "推荐", "种草", "求", "想要", - "怎么买", "哪里买", "想入", "好想要", "被种草", "下单", "链接", + "想买", + "求链接", + "在哪买", + "多少钱", + "推荐", + "种草", + "求", + "想要", + "怎么买", + "哪里买", + "想入", + "好想要", + "被种草", + "下单", + "链接", ] comparison_kw = [ - "比", "不如", "还是", "更", "不如买", "对比", "选择", + "比", + "不如", + "还是", + "更", + "不如买", + "对比", + "选择", ] for c in comments: @@ -916,12 +1032,15 @@ def _analyze_comments(self, comments: list[str]) -> tuple: # 简单高频词统计 from collections import Counter + word_counts = Counter(all_words) high_freq = [w for w, _ in word_counts.most_common(15)] return complaints, purchase_intents, comparison_mentions, high_freq - def crawl(self, category: str, count: int = 30, with_comments: bool = True, comment_level: str = "all"): + def crawl( + self, category: str, count: int = 30, with_comments: bool = True, comment_level: str = "all" + ): """完整抓取流程 Args: @@ -931,21 +1050,23 @@ def crawl(self, category: str, count: int = 30, with_comments: bool = True, comm comment_level: 评论层级 ("all" / "top" / "hot") """ if not self._logged_in: - msg = ("[Crawler] [FAIL] 未登录,无法抓取\n" - "[Crawler] [Hint] 本地: uv run python src/real_crawler.py \"品类名\" 登录\n" - "[Crawler] [Hint] 云端: 在 Streamlit Secrets 中配置 XHS_COOKIES") + msg = ( + "[Crawler] [FAIL] 未登录,无法抓取\n" + '[Crawler] [Hint] 本地: uv run python src/real_crawler.py "品类名" 登录\n' + "[Crawler] [Hint] 云端: 在 Streamlit Secrets 中配置 XHS_COOKIES" + ) print(msg) return 0 os.makedirs(RAW_DIR, exist_ok=True) - print(f"\n{'='*60}") + print(f"\n{'=' * 60}") print(f"[Crawler] [Crawl] 开始抓取: {category}") - print(f"{'='*60}") + print(f"{'=' * 60}") # 1. 搜索笔记(如发现需重新登录,自动触发一次登录后重试) - MAX_RETRIES = 1 - for attempt in range(MAX_RETRIES + 1): + max_retries = 1 + for attempt in range(max_retries + 1): notes = self.search(category, count) if notes: break # 找到笔记 → 继续 @@ -954,7 +1075,9 @@ def crawl(self, category: str, count: int = 30, with_comments: bool = True, comm # 不是因为登录弹窗导致的空结果 → 不重试 break - print(f"\n[Crawler] [Crawl] [RETRY] cookie 已过期,尝试重新登录(第{attempt+1}次)...") + print( + f"\n[Crawler] [Crawl] [RETRY] cookie 已过期,尝试重新登录(第{attempt + 1}次)..." + ) self._logged_in = False # 强制设为未登录态 if not self._is_cloud: if not self.login_interactive(): @@ -962,7 +1085,9 @@ def crawl(self, category: str, count: int = 30, with_comments: bool = True, comm return 0 print("[Crawler] [OK] 重新登录成功,重新搜索...") else: - print("[Crawler] [FAIL] 云端模式下 cookie 过期,请更新 Streamlit Secrets XHS_COOKIES") + print( + "[Crawler] [FAIL] 云端模式下 cookie 过期,请更新 Streamlit Secrets XHS_COOKIES" + ) return 0 if not notes: @@ -973,7 +1098,7 @@ def crawl(self, category: str, count: int = 30, with_comments: bool = True, comm saved = 0 print(f"\n[Crawler] [Crawl] 搜索到 {len(notes)} 篇笔记,开始逐篇抓取详情...") for i, note in enumerate(notes): - print(f"\n[Crawler] [{i+1}/{len(notes)}] {note.get('title', '')[:40]}") + print(f"\n[Crawler] [{i + 1}/{len(notes)}] {note.get('title', '')[:40]}") try: note = self.get_note_detail(note) @@ -992,6 +1117,7 @@ def crawl(self, category: str, count: int = 30, with_comments: bool = True, comm except Exception as e: print(f" [ERR] 处理笔记失败: {e}") import traceback + traceback.print_exc() # 尝试保存一个基础版本 try: @@ -1014,7 +1140,7 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: 策略 B: 探索页热门笔记标题提取关键词 策略 C: 兜底内置词库 """ - print(f"\n[Crawler] [HotSearch] 开始抓取小红书热榜...") + print("\n[Crawler] [HotSearch] 开始抓取小红书热榜...") items = [] try: @@ -1044,8 +1170,6 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: if search_clicked: time.sleep(3) # 等下拉面板渲染 - page_html = self.page.html or "" - # 方法 1: 面板文本提取 panel_texts = [] for panel_sel in [ @@ -1084,7 +1208,13 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: continue if suggest_els: - combined = " ".join([el.text.strip() for el in suggest_els if el.text and len(el.text.strip()) >= 2]) + combined = " ".join( + [ + el.text.strip() + for el in suggest_els + if el.text and len(el.text.strip()) >= 2 + ] + ) if combined: panel_texts = [combined] @@ -1093,33 +1223,36 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: seen_keywords = set() rank = 0 for pt in panel_texts: - for line in pt.replace('\t', '\n').split('\n'): + for line in pt.replace("\t", "\n").split("\n"): kw = line.strip() # 清理:去掉数字前缀、热度标记等 import re as re_mod - kw = re_mod.sub(r'^\d+[\.\、\)\s]*', '', kw) - kw = re_mod.sub(r'\s*(热|新|荐|HOT|爆)$', '', kw) + + kw = re_mod.sub(r"^\d+[\.\、\)\s]*", "", kw) + kw = re_mod.sub(r"\s*(热|新|荐|HOT|爆)$", "", kw) kw = kw.strip() if not kw or len(kw) < 2 or len(kw) > 25: continue if kw in seen_keywords: continue - if re_mod.match(r'^[\d\.\s\-—,,]+$', kw): + if re_mod.match(r"^[\d\.\s\-—,,]+$", kw): continue - if '小红书' in kw or '登录' in kw or '注册' in kw: + if "小红书" in kw or "登录" in kw or "注册" in kw: continue seen_keywords.add(kw) rank += 1 - items.append({ - "keyword": kw, - "rank": rank, - "tag": "热" if rank <= 5 else ("新" if rank <= 15 else ""), - "category": self._guess_category(kw), - "trend": "up" if rank <= 10 else "stable", - "hots": max(100 - rank * 3, 10), - }) + items.append( + { + "keyword": kw, + "rank": rank, + "tag": "热" if rank <= 5 else ("新" if rank <= 15 else ""), + "category": self._guess_category(kw), + "trend": "up" if rank <= 10 else "stable", + "hots": max(100 - rank * 3, 10), + } + ) if len(items) >= max_items: break if len(items) >= max_items: @@ -1160,8 +1293,14 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: # 从页面链接提取 try: links = self.page.eles("css:a[href*='/explore/']", timeout=3) - note_links = [l for l in links if l and l.text and len(l.text.strip()) > 3 - and 'footer' not in (getattr(l, 'parent', None) or '')] + note_links = [ + link + for link in links + if link + and link.text + and len(link.text.strip()) > 3 + and "footer" not in (getattr(link, "parent", None) or "") + ] title_els = note_links print(f"[Crawler] [HotSearch] 探索页链接提取: {len(title_els)} 条") except Exception: @@ -1169,19 +1308,55 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: # 过滤页面 chrome(导航、版权、菜单等) blacklist = { - "创作中心", "业务合作", "关于我们", "联系我们", "用户协议", - "隐私政策", "举报", "帮助", "反馈", "登录", "注册", - "首页", "发现", "消息", "通知", "我", "搜索", - "下载", "APP", "小程序", "桌面版", "手机版", - "关注", "推荐", "热门", "最新", "商品", "店铺", - "收藏", "点赞", "评论", "分享", "更多", - "小红书", "沪ICP", "ICP备", "备案", "版权所有", - "Cookie", "隐私", "条款", "广告", "推广", + "创作中心", + "业务合作", + "关于我们", + "联系我们", + "用户协议", + "隐私政策", + "举报", + "帮助", + "反馈", + "登录", + "注册", + "首页", + "发现", + "消息", + "通知", + "我", + "搜索", + "下载", + "APP", + "小程序", + "桌面版", + "手机版", + "关注", + "推荐", + "热门", + "最新", + "商品", + "店铺", + "收藏", + "点赞", + "评论", + "分享", + "更多", + "小红书", + "沪ICP", + "ICP备", + "备案", + "版权所有", + "Cookie", + "隐私", + "条款", + "广告", + "推广", } # 从标题提取关键词 title_keywords = {} import re as re_mod + for el in title_els[:60]: try: t = el.text.strip() @@ -1192,9 +1367,11 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: continue # 拆分提取有意义的词组 - for kw in re_mod.split(r'[,。,\.、\s##||【】\[\]()\(\)]+', t): + for kw in re_mod.split(r"[,。,\.、\s##||【】\[\]()\(\)]+", t): kw = kw.strip() - if 2 <= len(kw) <= 15 and not re_mod.match(r'^[\d\.\s\-—,,、/\??!!]+$', kw): + if 2 <= len(kw) <= 15 and not re_mod.match( + r"^[\d\.\s\-—,,、/\??!!]+$", kw + ): if kw not in blacklist: title_keywords[kw] = title_keywords.get(kw, 0) + 1 except Exception: @@ -1208,14 +1385,16 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: if kw in existing or len(kw) < 2 or kw in blacklist: continue rank += 1 - items.append({ - "keyword": kw, - "rank": rank, - "tag": "热" if rank <= 5 else "", - "category": self._guess_category(kw), - "trend": "up" if freq >= 3 else "stable", - "hots": min(freq * 25, 100), - }) + items.append( + { + "keyword": kw, + "rank": rank, + "tag": "热" if rank <= 5 else "", + "category": self._guess_category(kw), + "trend": "up" if freq >= 3 else "stable", + "hots": min(freq * 25, 100), + } + ) existing.add(kw) if len(items) >= max_items: break @@ -1233,6 +1412,7 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: except Exception as e: print(f"[Crawler] [HotSearch] 异常: {e}") import traceback + traceback.print_exc() items = self._fallback_hot_list() @@ -1242,15 +1422,37 @@ def fetch_hot_search(self, max_items: int = 30) -> list[dict]: def _guess_category(text: str) -> str: """根据关键词猜测品类""" cat_map = { - "穿搭": "服饰", "衣服": "服饰", "裙子": "服饰", "鞋": "服饰", - "化妆": "美妆", "护肤": "美妆", "口红": "美妆", "面膜": "美妆", - "零食": "食品", "吃": "食品", "蛋糕": "食品", "奶茶": "食品", - "家居": "家居", "收纳": "家居", "装修": "家居", "灯": "家居", - "手机": "数码", "耳机": "数码", "电脑": "数码", - "健身": "运动", "运动": "运动", "瑜伽": "运动", - "猫": "宠物", "狗": "宠物", "宠物": "宠物", - "旅行": "旅游", "旅游": "旅游", "酒店": "旅游", - "养娃": "母婴", "宝宝": "母婴", "孕": "母婴", + "穿搭": "服饰", + "衣服": "服饰", + "裙子": "服饰", + "鞋": "服饰", + "化妆": "美妆", + "护肤": "美妆", + "口红": "美妆", + "面膜": "美妆", + "零食": "食品", + "吃": "食品", + "蛋糕": "食品", + "奶茶": "食品", + "家居": "家居", + "收纳": "家居", + "装修": "家居", + "灯": "家居", + "手机": "数码", + "耳机": "数码", + "电脑": "数码", + "健身": "运动", + "运动": "运动", + "瑜伽": "运动", + "猫": "宠物", + "狗": "宠物", + "宠物": "宠物", + "旅行": "旅游", + "旅游": "旅游", + "酒店": "旅游", + "养娃": "母婴", + "宝宝": "母婴", + "孕": "母婴", } for kw, cat in cat_map.items(): if kw in text: @@ -1261,22 +1463,44 @@ def _guess_category(text: str) -> str: def _fallback_hot_list() -> list[dict]: """兜底热榜(极简内置词库,避免空榜)""" fallback = [ - "穿搭", "化妆", "护肤", "减肥", "健身", - "收纳", "装修", "零食", "奶茶", "咖啡", - "穿搭灵感", "显瘦穿搭", "平价好物", "家居好物", "数码好物", - "通勤穿搭", "早春穿搭", "夜间护肤", "抗老", "美白", - "宠物用品", "旅行攻略", "本地美食", "周末去哪儿", "读书推荐", + "穿搭", + "化妆", + "护肤", + "减肥", + "健身", + "收纳", + "装修", + "零食", + "奶茶", + "咖啡", + "穿搭灵感", + "显瘦穿搭", + "平价好物", + "家居好物", + "数码好物", + "通勤穿搭", + "早春穿搭", + "夜间护肤", + "抗老", + "美白", + "宠物用品", + "旅行攻略", + "本地美食", + "周末去哪儿", + "读书推荐", ] items = [] for i, kw in enumerate(fallback): - items.append({ - "keyword": kw, - "rank": i + 1, - "tag": "热" if i < 5 else "", - "category": XHSCrawler._guess_category(kw), - "trend": "up" if i < 10 else "stable", - "hots": max(100 - i * 3, 15), - }) + items.append( + { + "keyword": kw, + "rank": i + 1, + "tag": "热" if i < 5 else "", + "category": XHSCrawler._guess_category(kw), + "trend": "up" if i < 10 else "stable", + "hots": max(100 - i * 3, 15), + } + ) return items def close(self): diff --git a/src/retrievers.py b/src/retrievers.py index d789226..e9d5b53 100644 --- a/src/retrievers.py +++ b/src/retrievers.py @@ -3,13 +3,14 @@ 合并自原 step03_hybrid_retriever + step04_reranker 支持 ChromaDB(默认)和 PostgreSQL + pgvector(可选)双模式 """ + import asyncio -from src.logger import logger -from typing import List, Optional -from rank_bm25 import BM25Okapi -from langchain_core.documents import Document import jieba +from langchain_core.documents import Document +from rank_bm25 import BM25Okapi + +from src.logger import logger class HybridRetriever: @@ -131,6 +132,7 @@ class APIReranker: def __init__(self, model: str = None, api_key: str = None, base_url: str = None): from src.config import RERANKER_CONFIG + cfg = RERANKER_CONFIG self.model = model or cfg["model"] self.api_key = api_key or cfg["api_key"] @@ -151,9 +153,7 @@ def rerank(self, query: str, documents: list[Document]) -> list[float]: "Content-Type": "application/json", } - resp = requests.post( - f"{self.base_url}/rerank", headers=headers, json=payload, timeout=30 - ) + resp = requests.post(f"{self.base_url}/rerank", headers=headers, json=payload, timeout=30) resp.raise_for_status() data = resp.json() @@ -184,9 +184,7 @@ async def arerank(self, query: str, documents: list[Document]) -> list[float]: } async with httpx.AsyncClient(timeout=30.0) as client: - resp = await client.post( - f"{self.base_url}/rerank", headers=headers, json=payload - ) + resp = await client.post(f"{self.base_url}/rerank", headers=headers, json=payload) resp.raise_for_status() data = resp.json() @@ -201,6 +199,7 @@ async def arerank(self, query: str, documents: list[Document]) -> list[float]: # PostgreSQL + pgvector 向量检索 # ================================================================ + class PGVectorStore: """PostgreSQL + pgvector 向量存储适配器 @@ -215,18 +214,18 @@ class PGVectorStore: def __init__(self, embedding_model=None): if embedding_model is None: from src.ingestion import get_embeddings + self.embedding_model = get_embeddings() else: self.embedding_model = embedding_model - async def similarity_search( - self, query: str, k: int = 5 - ) -> list[Document]: + async def similarity_search(self, query: str, k: int = 5) -> list[Document]: """向量相似度搜索(返回 Document 对象)""" - from src.core.database import get_db, search_by_vector - # 1. Embed 查询 import asyncio as _asyncio + + from src.core.database import get_db, search_by_vector + loop = _asyncio.get_running_loop() query_embedding = await loop.run_in_executor( None, lambda: self.embedding_model.embed_query(query) @@ -251,9 +250,10 @@ async def similarity_search_with_score( self, query: str, k: int = 5 ) -> list[tuple[Document, float]]: """向量相似度搜索(带分数)""" + import asyncio as _asyncio + from src.core.database import get_db, search_by_vector - import asyncio as _asyncio loop = _asyncio.get_running_loop() query_embedding = await loop.run_in_executor( None, lambda: self.embedding_model.embed_query(query) @@ -275,6 +275,7 @@ async def similarity_search_with_score( async def count(self) -> int: """获取文档总数""" from src.core.database import get_db, get_document_count + async for session in get_db(): return await get_document_count(session) return 0 @@ -348,13 +349,13 @@ def bm25_work(): return [doc_map[rid] for rid in ranked_ids if rid in doc_map] -async def create_pg_vectorstore() -> Optional[PGVectorStore]: +async def create_pg_vectorstore() -> PGVectorStore | None: """尝试创建 PG 向量存储(如果 PG 可用) 返回 None 表示 PG 不可用,应回退到 ChromaDB。 """ try: - from src.core.database import init_db, get_db, get_document_count + from src.core.database import get_db, get_document_count, init_db await init_db() diff --git a/tests/test_agents/test_insight_agent.py b/tests/test_agents/test_insight_agent.py index 15b617e..5051b2d 100644 --- a/tests/test_agents/test_insight_agent.py +++ b/tests/test_agents/test_insight_agent.py @@ -1,10 +1,13 @@ """ test_insight_agent.py — InsightGenerator 单元测试 """ -import pytest -from unittest.mock import MagicMock, AsyncMock, patch -import sys + import os +import sys +from unittest.mock import AsyncMock, MagicMock + +import pytest + sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../..")) @@ -40,12 +43,14 @@ def _make_aggregated(self, **overrides): def test_instantiation(self): from src.agents.insight_agent import InsightGenerator + gen = InsightGenerator() assert gen is not None assert gen.llm is not None def test_generate_fallback_basic(self): from src.agents.insight_agent import InsightGenerator + gen = InsightGenerator() aggregated = self._make_aggregated() report = gen.generate_fallback(aggregated, category="磁吸感应灯") @@ -55,12 +60,14 @@ def test_generate_fallback_basic(self): def test_generate_fallback_empty_data(self): from src.agents.insight_agent import InsightGenerator + gen = InsightGenerator() report = gen.generate_fallback({"note_count": 0}, category="测试") assert "没有足够的评论数据" in report def test_generate_fallback_low_score_warning(self): from src.agents.insight_agent import InsightGenerator + gen = InsightGenerator() aggregated = self._make_aggregated(selection_score=30, competition_score=20) report = gen.generate_fallback(aggregated, category="红海品类") @@ -70,6 +77,7 @@ def test_generate_fallback_low_score_warning(self): async def test_agenerate_returns_report(self): """agenerate() 应返回报告文本""" from src.agents.insight_agent import InsightGenerator + mock_llm = MagicMock() mock_llm.ainvoke = AsyncMock() mock_llm.ainvoke.return_value.content = "测试报告内容" diff --git a/tests/test_agents/test_supervisor.py b/tests/test_agents/test_supervisor.py deleted file mode 100644 index c2eb255..0000000 --- a/tests/test_agents/test_supervisor.py +++ /dev/null @@ -1,63 +0,0 @@ -""" -test_supervisor.py — Supervisor 策略路由测试 -""" -import pytest -from unittest.mock import AsyncMock -import sys -import os -sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../..")) - - -class TestSupervisor: - """Supervisor 策略路由测试""" - - def test_supervisor_instantiation(self): - """Supervisor 可以正常实例化""" - from src.agents.supervisor import Supervisor - sup = Supervisor() - assert sup is not None - assert sup.llm is not None - - @pytest.mark.asyncio - async def test_decide_with_empty_strategies(self): - """adecide() 在策略列表为空时应返回 hybrid""" - from src.agents.supervisor import Supervisor - mock_llm = AsyncMock() - mock_llm.ainvoke.return_value.content = " hybrid " - - sup = Supervisor(llm=mock_llm) - result = await sup.adecide("测试问题", ["vector", "keyword", "hybrid"]) - assert result == "hybrid" - - @pytest.mark.asyncio - async def test_decide_falls_back_to_hybrid(self): - """adecide() 在 LLM 返回无效策略时应退回到 hybrid""" - from src.agents.supervisor import Supervisor - mock_llm = AsyncMock() - mock_llm.ainvoke.return_value.content = "invalid_strategy" - - sup = Supervisor(llm=mock_llm) - result = await sup.adecide("测试问题", ["vector", "keyword", "hybrid"]) - assert result == "hybrid" - - @pytest.mark.asyncio - async def test_decide_returns_valid_strategy(self): - """adecide() 应返回有效策略""" - from src.agents.supervisor import Supervisor - mock_llm = AsyncMock() - mock_llm.ainvoke.return_value.content = "vector" - - sup = Supervisor(llm=mock_llm) - result = await sup.adecide("概念性问题", ["vector", "keyword", "hybrid"]) - assert result == "vector" - - @pytest.mark.asyncio - async def test_adecide_basic(self): - """adecide() 异步版本基本功能""" - from src.agents.supervisor import Supervisor - mock_llm = AsyncMock() - mock_llm.ainvoke.return_value.content = "keyword" - - sup = Supervisor(llm=mock_llm) - result = await sup.adecide("专有名词查询", ["vector", "keyword", "hybrid"]) - assert result == "keyword" diff --git a/tests/test_api/test_health.py b/tests/test_api/test_health.py index 586c765..0754b51 100644 --- a/tests/test_api/test_health.py +++ b/tests/test_api/test_health.py @@ -1,13 +1,14 @@ """ test_health.py — API 健康检查 + 统计端点测试 """ + import pytest -from unittest.mock import MagicMock, AsyncMock, patch -from httpx import AsyncClient, ASGITransport +from httpx import ASGITransport, AsyncClient class MockAppState: """模拟已初始化的 AppState""" + is_ready = True error = None stats = { @@ -20,6 +21,7 @@ class MockAppState: @pytest.fixture def app(): from src.api.main import app + app.state.app_state = MockAppState() return app @@ -60,10 +62,12 @@ async def test_stats_returns_data(app): async def test_stats_503_when_not_ready(): """GET /api/stats — 未初始化应返回 503""" from src.api.main import app + # 用 not-ready 的 mock class NotReady: is_ready = False error = "未初始化" + app.state.app_state = NotReady() transport = ASGITransport(app=app) diff --git a/tests/test_api/test_insight.py b/tests/test_api/test_insight.py index d178d02..49a7fa4 100644 --- a/tests/test_api/test_insight.py +++ b/tests/test_api/test_insight.py @@ -1,8 +1,9 @@ """ test_insight.py — Insight 端点测试 """ + import pytest -from httpx import AsyncClient, ASGITransport +from httpx import ASGITransport, AsyncClient class MockAppState: @@ -14,6 +15,7 @@ class MockAppState: @pytest.fixture def app(): from src.api.main import app + app.state.app_state = MockAppState() return app diff --git a/tests/test_api/test_qa.py b/tests/test_api/test_qa.py index 39b50d9..15d2d1b 100644 --- a/tests/test_api/test_qa.py +++ b/tests/test_api/test_qa.py @@ -1,8 +1,12 @@ """ test_qa.py — QA 端点测试 """ + +from types import SimpleNamespace +from unittest.mock import AsyncMock + import pytest -from httpx import AsyncClient, ASGITransport +from httpx import ASGITransport, AsyncClient class MockAppState: @@ -14,6 +18,7 @@ class MockAppState: @pytest.fixture def app(): from src.api.main import app + app.state.app_state = MockAppState() return app @@ -45,3 +50,32 @@ async def test_qa_validates_input(app): async with AsyncClient(transport=transport, base_url="http://test") as client: resp = await client.post("/api/qa", json={}) assert resp.status_code == 422 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("question", "cleaned", "k"), + [("收纳盒材质 123", "收纳盒材质", 5), ("收纳盒推荐", "收纳盒推荐", 8)], +) +async def test_qa_uses_hybrid_retrieval_without_a_supervisor(monkeypatch, question, cleaned, k): + from langchain_core.documents import Document + + from src.api.routes import qa + + document = Document(page_content="收纳盒材质和收纳盒推荐") + retriever = SimpleNamespace(ahybrid_search=AsyncMock(return_value=[document])) + reranker = SimpleNamespace(arerank=AsyncMock(return_value=[0.9])) + llm = SimpleNamespace(ainvoke=AsyncMock(return_value=SimpleNamespace(content="测试回答"))) + monkeypatch.setattr(qa, "ChatOpenAI", lambda **kwargs: llm) + + response = await qa.run_qa( + qa.QARequest(question=question), + SimpleNamespace(hybrid_retriever=retriever, reranker=reranker), + ) + + assert response.success is True + assert response.question == question + assert response.answer == "测试回答" + retriever.ahybrid_search.assert_awaited_once_with(cleaned, k=k, bm25_k=40, final_k=k) + reranker.arerank.assert_awaited_once_with(cleaned, [document]) + llm.ainvoke.assert_awaited_once() diff --git a/tests/test_demand_agent.py b/tests/test_demand_agent.py index 0e397ba..f471ac1 100644 --- a/tests/test_demand_agent.py +++ b/tests/test_demand_agent.py @@ -1,9 +1,12 @@ """ test_demand_agent.py — DemandAggregator 单元测试 """ -import pytest -import sys + import os +import sys + +import pytest + sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) from src.agents.demand_agent import DemandAggregator @@ -57,15 +60,17 @@ def test_empty_input_returns_empty_result(self): def test_single_analysis_basic_stats(self): agg = DemandAggregator() - result = agg.aggregate([ - _make_analysis( - complaints=["太贵了", "质量差"], - purchase_intent=["想买"], - related_brands=["品牌A"], - ask_link_count=3, - likes=100, - ) - ]) + result = agg.aggregate( + [ + _make_analysis( + complaints=["太贵了", "质量差"], + purchase_intent=["想买"], + related_brands=["品牌A"], + ask_link_count=3, + likes=100, + ) + ] + ) assert result["note_count"] == 1 assert result["avg_likes"] == 100.0 assert result["total_ask_link"] == 3 @@ -98,47 +103,59 @@ def test_complaint_frequency_aggregation(self): def test_profit_score_high_margin(self): """高利润率应得高利润分""" agg = DemandAggregator() - result = agg.aggregate([ - _make_analysis(price=100, cost=10, profit_margin=0.9), - ]) + result = agg.aggregate( + [ + _make_analysis(price=100, cost=10, profit_margin=0.9), + ] + ) assert result["profit_score"] >= 70 def test_profit_score_low_margin(self): """低利润率应得低利润分""" agg = DemandAggregator() - result = agg.aggregate([ - _make_analysis(price=50, cost=45, profit_margin=0.1), - ]) + result = agg.aggregate( + [ + _make_analysis(price=50, cost=45, profit_margin=0.1), + ] + ) assert result["profit_score"] < 60 def test_weight_affects_logistics_score(self): """重量越大物流分越低""" agg_light = DemandAggregator() - r_light = agg_light.aggregate([ - _make_analysis(weight=0.1), - ]) + r_light = agg_light.aggregate( + [ + _make_analysis(weight=0.1), + ] + ) agg_heavy = DemandAggregator() - r_heavy = agg_heavy.aggregate([ - _make_analysis(weight=2.0), - ]) + r_heavy = agg_heavy.aggregate( + [ + _make_analysis(weight=2.0), + ] + ) assert r_light["logistics_score"] >= r_heavy["logistics_score"] def test_low_competition_boosts_score(self): """低竞争品类竞争分更高""" agg = DemandAggregator() - result = agg.aggregate([ - _make_analysis(competition_level="低", entry_difficulty="低"), - _make_analysis(competition_level="低", entry_difficulty="低"), - ]) + result = agg.aggregate( + [ + _make_analysis(competition_level="低", entry_difficulty="低"), + _make_analysis(competition_level="低", entry_difficulty="低"), + ] + ) assert result["competition_score"] >= 60 def test_evergreen_ratio_calculation(self): agg = DemandAggregator() - result = agg.aggregate([ - _make_analysis(category_type="常青款"), - _make_analysis(category_type="常青款"), - _make_analysis(category_type="季节性"), - ]) + result = agg.aggregate( + [ + _make_analysis(category_type="常青款"), + _make_analysis(category_type="常青款"), + _make_analysis(category_type="季节性"), + ] + ) assert result["evergreen_ratio"] == pytest.approx(2 / 3, abs=0.01) def test_empty_analysis_produces_fallback(self): diff --git a/tests/test_query_utils.py b/tests/test_query_utils.py new file mode 100644 index 0000000..7a98019 --- /dev/null +++ b/tests/test_query_utils.py @@ -0,0 +1,42 @@ +"""Cover the query preprocessing and retrieval sizing used by the QA routes.""" + +import pytest + +from src.core.query_utils import ( + clean_query, + is_brand_comparison, + resolve_bm25_k, + resolve_k, +) + + +@pytest.mark.parametrize( + ("query", "expected"), + [ + ("灯 123 推荐", "灯 推荐"), + ("灯 12345", "灯 12345"), + ("123", "123"), + (" 磁吸 感应灯 ", "磁吸 感应灯"), + ], +) +def test_clean_query_preserves_content_and_long_identifiers(query, expected): + assert clean_query(query) == expected + + +@pytest.mark.parametrize( + ("query", "comparison", "k"), + [("品牌A对比品牌B", True, 8), ("收纳盒推荐", True, 8), ("收纳盒材质", False, 5)], +) +def test_comparisons_expand_retrieval(query, comparison, k): + assert is_brand_comparison(query) is comparison + assert resolve_k(query) == k + + +def test_retrieval_sizes_can_be_configured(): + assert resolve_k("收纳盒材质", base=3, brand=7) == 3 + assert resolve_k("收纳盒推荐", base=3, brand=7) == 7 + + +@pytest.mark.parametrize(("k", "expected"), [(5, 40), (9, 45)]) +def test_keyword_retrieval_keeps_a_minimum_candidate_pool(k, expected): + assert resolve_bm25_k("收纳盒材质", base_k=k) == expected diff --git a/uv.lock b/uv.lock index 1460c65..0784f56 100644 --- a/uv.lock +++ b/uv.lock @@ -4030,6 +4030,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/9d/7a/d968e294073affff457b041c2be9868a40c1c71f4a35fcc1e45e5493067b/pytest_cov-7.1.0-py3-none-any.whl", hash = "sha256:a0461110b7865f9a271aa1b51e516c9a95de9d696734a2f71e3e78f46e1d4678", size = 22876, upload-time = "2026-03-21T20:11:14.438Z" }, ] +[[package]] +name = "pytest-timeout" +version = "2.4.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pytest" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ac/82/4c9ecabab13363e72d880f2fb504c5f750433b2b6f16e99f4ec21ada284c/pytest_timeout-2.4.0.tar.gz", hash = "sha256:7e68e90b01f9eff71332b25001f85c75495fc4e3a836701876183c4bcfd0540a", size = 17973, upload-time = "2025-05-05T19:44:34.99Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/fa/b6/3127540ecdf1464a00e5a01ee60a1b09175f6913f0644ac748494d9c4b21/pytest_timeout-2.4.0-py3-none-any.whl", hash = "sha256:c42667e5cdadb151aeb5b26d114aff6bdf5a907f176a007a30b940d3d865b5c2", size = 14382, upload-time = "2025-05-05T19:44:33.502Z" }, +] + [[package]] name = "python-dateutil" version = "2.9.0.post0" @@ -4211,6 +4223,7 @@ dev = [ { name = "pytest" }, { name = "pytest-asyncio" }, { name = "pytest-cov" }, + { name = "pytest-timeout" }, { name = "ruff" }, { name = "sentence-transformers" }, ] @@ -4237,6 +4250,7 @@ requires-dist = [ { name = "pytest", marker = "extra == 'dev'", specifier = ">=9.1.0" }, { name = "pytest-asyncio", marker = "extra == 'dev'", specifier = ">=0.24.0" }, { name = "pytest-cov", marker = "extra == 'dev'", specifier = ">=6.0.0" }, + { name = "pytest-timeout", marker = "extra == 'dev'", specifier = ">=2.3.1" }, { name = "python-dotenv", specifier = ">=1.0" }, { name = "pyyaml", specifier = ">=6.0" }, { name = "rank-bm25", specifier = ">=0.2" },