diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..cb7e83b --- /dev/null +++ b/.dockerignore @@ -0,0 +1,38 @@ +# Python +__pycache__/ +*.py[cod] +*.egg-info/ +.venv/ +dist/ +build/ + +# Git +.git/ +.gitignore + +# IDE +.vscode/ +.idea/ + +# Environment +.env + +# Logs +*.log + +# Test +.pytest_cache/ +htmlcov/ +.coverage + +# Docker +Dockerfile +.dockerignore +docker-compose.yml + +# OS +.DS_Store +Thumbs.db + +# Project specific +data/chroma_db/*.sqlite3 diff --git a/.env.example b/.env.example index 3582972..d110e21 100644 --- a/.env.example +++ b/.env.example @@ -26,6 +26,32 @@ EMBEDDING_MODEL=BAAI/bge-m3 # Cross-Encoder for relevance scoring RERANKER_MODEL=BAAI/bge-reranker-v2-m3 +# ===== PostgreSQL(Docker Compose 内自动配置)===== +# 本地开发(非 Docker):用自己的 PG 实例 +# DATABASE_URL=postgresql+asyncpg://postgres:postgres@localhost:5432/rednote_insight +# +# Docker Compose 部署:自动通过 docker-compose.yml 注入 +# DATABASE_URL=postgresql+asyncpg://postgres:postgres@db:5432/rednote_insight + +# ===== Redis(Docker Compose 内自动配置)===== +# 本地开发(非 Docker): +# REDIS_URL=redis://localhost:6379/0 +# +# Docker Compose 部署:自动通过 docker-compose.yml 注入 +# REDIS_URL=redis://redis:6379/0 + +# ===== LangFuse 可观测性(可选)===== +# 注册 https://cloud.langfuse.com 获取 +# LANGFUSE_PUBLIC_KEY=pk-... +# LANGFUSE_SECRET_KEY=sk-... +# LANGFUSE_HOST=https://cloud.langfuse.com + +# ===== 限流配置(可选)===== +# RATE_LIMIT_ENABLED=true +# RATE_LIMIT_INSIGHT=10/minute +# RATE_LIMIT_QA=20/minute +# RATE_LIMIT_CRAWL=5/minute + # ===== 小红书 Cookie(Streamlit Cloud 部署用)===== # 在本地运行: uv run python scripts/export_cookies.py # 复制输出的 JSON 粘贴到 Streamlit Cloud Secrets → XHS_COOKIES diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..01dc5f4 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,133 @@ +# ============================================================================= +# RedNote Insight — CI Pipeline +# ============================================================================= +# 触发条件:push 到 main 或 PR 到 main +# 自动运行:lint → test → build +# ============================================================================= + +name: CI + +on: + push: + branches: [main, master, refactor/production] + pull_request: + branches: [main, master] + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +env: + PYTHON_VERSION: "3.11" + UV_VERSION: "0.6.x" + +jobs: + # ===== 1. Lint(代码风格检查)===== + lint: + name: Lint (ruff) + runs-on: ubuntu-latest + timeout-minutes: 5 + steps: + - uses: actions/checkout@v4 + + - name: Install uv + uses: astral-sh/setup-uv@v5 + with: + version: ${{ env.UV_VERSION }} + + - name: Set up Python + run: uv python install ${{ env.PYTHON_VERSION }} + + - name: Install dependencies + run: uv sync --frozen --no-dev + + - name: Install dev dependencies + run: uv sync --frozen + + - name: Run ruff check + run: uv run ruff check src/ tests/ + + - name: Run ruff format check + run: uv run ruff format --check src/ tests/ + + # ===== 2. Test(单元测试 + 集成测试)===== + test: + name: Test (pytest) + runs-on: ubuntu-latest + timeout-minutes: 15 + needs: lint + services: + # PostgreSQL + pgvector(用于向量检索测试) + postgres: + image: pgvector/pgvector:pg16 + env: + POSTGRES_USER: postgres + POSTGRES_PASSWORD: postgres + POSTGRES_DB: rednote_insight_test + ports: + - 5432:5432 + options: >- + --health-cmd pg_isready + --health-interval 10s + --health-timeout 5s + --health-retries 5 + + steps: + - uses: actions/checkout@v4 + + - name: Install uv + uses: astral-sh/setup-uv@v5 + with: + version: ${{ env.UV_VERSION }} + + - name: Set up Python + run: uv python install ${{ env.PYTHON_VERSION }} + + - name: Install dependencies + run: uv sync --frozen + + - 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 + continue-on-error: true + + - name: Run tests + run: | + uv run pytest tests/ \ + -v \ + --tb=short \ + --timeout=60 \ + -p no:warnings + env: + OPENAI_API_KEY: ci-test-key + OPENAI_BASE_URL: https://api.siliconflow.cn/v1 + + # ===== 3. Build(Docker 镜像构建验证)===== + build: + name: Build (Docker) + runs-on: ubuntu-latest + timeout-minutes: 15 + needs: lint + steps: + - uses: actions/checkout@v4 + + - name: Set up Docker Buildx + uses: docker/setup-buildx-action@v3 + + - name: Build Docker image + uses: docker/build-push-action@v6 + with: + context: . + push: false + load: true + tags: rednote-insight:ci-test + cache-from: type=gha + cache-to: type=gha,mode=max + + - name: Verify image + run: | + docker images rednote-insight:ci-test + echo "✅ Docker 镜像构建成功" diff --git a/.gitignore b/.gitignore index 8e5e8c1..56f4772 100644 --- a/.gitignore +++ b/.gitignore @@ -4,6 +4,7 @@ __pycache__/ *.egg-info/ dist/ build/ +.pytest_cache/ # 环境变量 .env @@ -14,6 +15,10 @@ sample_data.csv # 运行时日志 app.log +server.log + +# 缓存文件 +data/trending_cache.json # Vector DB index files are small — we COMMIT them so Streamlit Cloud # doesn't have to rebuild from scratch (saves 30-60s cold start). @@ -29,8 +34,6 @@ Thumbs.db # Streamlit # .streamlit/secrets.toml — 提交此文件以在免费版 Streamlit Cloud 使用 -*.pyc -__pycache__/ # 小红书登录 Cookie(敏感信息,不要提交) data/cookies.json diff --git a/.streamlit/config.toml b/.streamlit/config.toml deleted file mode 100644 index 752fd19..0000000 --- a/.streamlit/config.toml +++ /dev/null @@ -1,20 +0,0 @@ -[theme] -primaryColor = "#ff5a5f" -backgroundColor = "#f8f9fa" -secondaryBackgroundColor = "#ffffff" -textColor = "#1a1a2e" -font = "sans serif" - -[server] -headless = true -runOnSave = false -fileWatcherType = "poll" - -[browser] -gatherUsageStats = false -serverAddress = "0.0.0.0" -serverPort = 8501 - -[client] -showErrorDetails = true -toolbarMode = "minimal" diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..f8c9f7f --- /dev/null +++ b/Dockerfile @@ -0,0 +1,61 @@ +# ============================================================================= +# RedNote Insight — 多阶段 Docker 构建 +# ============================================================================= +# 用法: +# docker build -t rednote-insight . +# docker run -p 8000:8000 --env-file .env rednote-insight +# ============================================================================= + +# ===== Stage 1: Builder — 装依赖 ===== +FROM python:3.11-slim AS builder + +WORKDIR /build + +# 系统依赖(ChromaDB 需要 sqlite3, jieba 等) +RUN apt-get update && apt-get install -y --no-install-recommends \ + build-essential \ + curl \ + && rm -rf /var/lib/apt/lists/* + +# 安装 uv +RUN pip install --no-cache-dir uv + +# 先复制依赖文件(利用 Docker 缓存层) +COPY pyproject.toml uv.lock ./ + +# 安装生产依赖到虚拟环境 +RUN uv sync --frozen --no-dev --no-editable + +# ===== Stage 2: Runtime — 最小镜像 ===== +FROM python:3.11-slim AS runtime + +WORKDIR /app + +# 运行时系统依赖 +RUN apt-get update && apt-get install -y --no-install-recommends \ + libsqlite3-0 \ + && rm -rf /var/lib/apt/lists/* + +# 从 builder 复制虚拟环境 +COPY --from=builder /build/.venv /app/.venv + +# 复制项目源码和配置 +COPY pyproject.toml ./ +COPY src/ ./src/ +COPY static/ ./static/ +COPY data/ ./data/ +COPY .env.example ./ + +# 创建数据目录(如果不存在) +RUN mkdir -p /app/data/raw /app/data/chroma_db + +# 暴露端口 +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/uv", "run", "uvicorn", "src.api.main:app", \ + "--host", "0.0.0.0", "--port", "8000", "--log-level", "info"] diff --git a/README.md b/README.md index c6a2a3d..12ca2f9 100644 --- a/README.md +++ b/README.md @@ -1,22 +1,23 @@ -# 🎯 小红书爆款雷达 — AI 选品洞察引擎 +# 🎯 小红书爆款雷达 — AI 选品 + 选题引擎
- 翻评论 · 找痛点 · 定方向 — 让 AI 从小红书评论区挖出下一个爆款 + 翻评论 · 找痛点 · 定方向 — 一个品类名,两套完整方案
-
-
-
+
+
+
-
-
+
+
+
- 为 AI 应用开发者而建。如果对你有帮助,给个 ⭐ Star -
+MIT © 2026 diff --git a/alembic.ini b/alembic.ini new file mode 100644 index 0000000..d575ac1 --- /dev/null +++ b/alembic.ini @@ -0,0 +1,43 @@ +# Alembic 配置文件 +# ===================== + +[alembic] +# 迁移脚本目录 +script_location = alembic + +# SQLAlchemy URL(可被环境变量覆盖) +sqlalchemy.url = postgresql+asyncpg://postgres:postgres@localhost:5432/rednote_insight + +# 日志 +[loggers] +keys = root,sqlalchemy,alembic + +[handlers] +keys = console + +[formatters] +keys = generic + +[logger_root] +level = WARN +handlers = console + +[logger_sqlalchemy] +level = WARN +handlers = +qualname = sqlalchemy.engine + +[logger_alembic] +level = INFO +handlers = +qualname = alembic + +[handler_console] +class = StreamHandler +args = (sys.stderr,) +level = NOTSET +formatter = generic + +[formatter_generic] +format = %(levelname)-5.5s [%(name)s] %(message)s +datefmt = %H:%M:%S diff --git a/data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/link_lists.bin b/alembic/__init__.py similarity index 100% rename from data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/link_lists.bin rename to alembic/__init__.py diff --git a/alembic/env.py b/alembic/env.py new file mode 100644 index 0000000..d889c69 --- /dev/null +++ b/alembic/env.py @@ -0,0 +1,87 @@ +""" +alembic/env.py — Alembic 迁移环境配置 +======================================== +用于 PostgreSQL + pgvector 数据库迁移管理。 + +用法: + uv run alembic revision --autogenerate -m "描述" + uv run alembic upgrade head + uv run alembic downgrade -1 +""" + +import asyncio +from logging.config import fileConfig + +from sqlalchemy import pool +from sqlalchemy.engine import Connection +from sqlalchemy.ext.asyncio import async_engine_from_config + +from alembic import context + +# Alembic Config 对象 +config = context.config + +# 日志配置 +if config.config_file_name is not None: + fileConfig(config.config_file_name) + +# 目标元数据(所有表定义) +from src.core.database import Base +target_metadata = Base.metadata + + +def get_url() -> str: + """从配置文件或环境变量获取 DATABASE_URL""" + return config.get_main_option( + "sqlalchemy.url", + "postgresql+asyncpg://postgres:postgres@localhost:5432/rednote_insight", + ) + + +def run_migrations_offline() -> None: + """离线模式:生成 SQL 脚本(不连接数据库)""" + url = get_url() + context.configure( + url=url, + target_metadata=target_metadata, + literal_binds=True, + dialect_opts={"paramstyle": "named"}, + ) + + with context.begin_transaction(): + context.run_migrations() + + +def do_run_migrations(connection: Connection) -> None: + context.configure(connection=connection, target_metadata=target_metadata) + + with context.begin_transaction(): + context.run_migrations() + + +async def run_async_migrations() -> None: + """在线模式:连接数据库并执行迁移""" + configuration = config.get_section(config.config_ini_section) + configuration["sqlalchemy.url"] = get_url() + + connectable = async_engine_from_config( + configuration, + prefix="sqlalchemy.", + poolclass=pool.NullPool, + ) + + async with connectable.connect() as connection: + await connection.run_sync(do_run_migrations) + + await connectable.dispose() + + +def run_migrations_online() -> None: + """在线模式入口""" + asyncio.run(run_async_migrations()) + + +if context.is_offline_mode(): + run_migrations_offline() +else: + run_migrations_online() diff --git a/alembic/script.py.mako b/alembic/script.py.mako new file mode 100644 index 0000000..1a2f2d9 --- /dev/null +++ b/alembic/script.py.mako @@ -0,0 +1,25 @@ +"""${message} + +Revision ID: ${up_revision} +Revises: ${down_revision | comma,n} +Create Date: ${create_date} +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +${imports if imports else ""} + +# revision identifiers +revision: str = ${repr(up_revision)} +down_revision: Union[str, None] = ${repr(down_revision)} +branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)} +depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)} + + +def upgrade() -> None: + ${upgrades if upgrades else "pass"} + + +def downgrade() -> None: + ${downgrades if downgrades else "pass"} diff --git a/alembic/versions/001_initial.py b/alembic/versions/001_initial.py new file mode 100644 index 0000000..515b5b6 --- /dev/null +++ b/alembic/versions/001_initial.py @@ -0,0 +1,64 @@ +"""001_initial — 初始迁移:documents 表 + pgvector 扩展 + +Revision ID: 001 +Revises: None +Create Date: 2026-06-23 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + + +# revision identifiers +revision: str = "001" +down_revision: Union[str, None] = None +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + """创建 documents 表 + pgvector 扩展 + HNSW 索引""" + + # 1. 启用 pgvector 扩展 + op.execute("CREATE EXTENSION IF NOT EXISTS vector") + + # 2. 创建 documents 表 + op.create_table( + "documents", + sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True), + sa.Column("content", sa.Text(), nullable=False), + sa.Column("metadata", postgresql.JSONB(), nullable=False), + sa.Column("embedding", postgresql.ARRAY(sa.Float()), nullable=True), + sa.Column("created_at", sa.DateTime(), nullable=False, + server_default=sa.text("now()")), + sa.Column("updated_at", sa.DateTime(), nullable=False, + server_default=sa.text("now()")), + ) + + # 3. 将 embedding 列转为 pgvector 类型 + # (先创建为 ARRAY(Float),再 cast 为 vector) + op.execute(""" + ALTER TABLE documents + ALTER COLUMN embedding TYPE vector(1024) + USING embedding::vector(1024) + """) + + # 4. 创建索引 + op.create_index("ix_documents_created_at", "documents", ["created_at"]) + + # HNSW 向量索引(余弦相似度) + op.execute(""" + CREATE INDEX IF NOT EXISTS ix_documents_embedding_hnsw + ON documents + USING hnsw (embedding vector_cosine_ops) + WITH (m = 16, ef_construction = 200) + """) + + +def downgrade() -> None: + """回滚:删除 documents 表""" + op.drop_index("ix_documents_embedding_hnsw", table_name="documents") + op.drop_index("ix_documents_created_at", table_name="documents") + op.drop_table("documents") diff --git a/alembic/versions/__init__.py b/alembic/versions/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/api.py b/api.py index 1c270ee..7b28904 100644 --- a/api.py +++ b/api.py @@ -1,497 +1,8 @@ """ -api.py — 小红书爆款雷达 FastAPI 后端 -===================================== -Phase 1: 完整 API + 前端页面托管 -Phase 2: 接入真实爬虫替换假数据 +api.py — 向后兼容入口 +====================== +原来的启动方式 `uv run uvicorn api:app` 仍然可用。 -启动: uv run uvicorn api:app --reload --port 8000 +实际应用定义在 src.api.main,避免代码重复。 """ -import os -import sys -import json -import time -from pathlib import Path -from contextlib import asynccontextmanager - -from fastapi import FastAPI, HTTPException -from fastapi.staticfiles import StaticFiles -from fastapi.responses import FileResponse, JSONResponse -from pydantic import BaseModel - -# 确保能导入项目模块 -sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - -# ============================================================ -# 请求/响应模型 -# ============================================================ - -class InsightRequest(BaseModel): - category: str - -class QARequest(BaseModel): - question: str - strategy: str = "hybrid" # auto / vector / keyword / hybrid - -class EvaluateRequest(BaseModel): - categories: list[str] = [] # 为空则评估全部品类 - -class CrawlRequest(BaseModel): - category: str - count: int = 20 - -class InsightResponse(BaseModel): - success: bool - category: str - report: str - notes_count: int - generated_count: int = 0 - elapsed: float - -class QAResponse(BaseModel): - success: bool - question: str - answer: str - elapsed: float - -class StatsResponse(BaseModel): - success: bool - categories: list[str] - total_notes: int - total_chunks: int - message: str - - -# ============================================================ -# 应用生命周期 -# ============================================================ - -_runtime = None # 全局运行时状态 - - -@asynccontextmanager -async def lifespan(app: FastAPI): - """应用启动时初始化,关闭时清理""" - global _runtime - print("[API] Initializing runtime...") - _runtime = _init_runtime() - if _runtime["error"]: - print(f"[API] WARNING: {_runtime['error']}") - else: - print(f"[API] READY - {_runtime['stats']['total_chunks']} chunks") - yield - print("[API] Shutting down") - - -app = FastAPI( - title="小红书爆款雷达 API", - description="翻评论、找痛点、定方向 — AI 选品洞察引擎", - version="1.0.0", - lifespan=lifespan, -) - - -# ============================================================ -# 初始化逻辑(复用原有代码) -# ============================================================ - -def _init_runtime() -> dict: - """初始化向量库、检索器、LangGraph""" - from src.retrievers import HybridRetriever, APIReranker - from src.graph import build_graph - from rank_bm25 import BM25Okapi - import jieba - from src.ingestion import load_raw_documents, chunk_documents, load_vectorstore, build_vectorstore - - project_root = os.path.dirname(os.path.abspath(__file__)) - raw_dir = os.path.join(project_root, "data", "raw") - chroma_dir = os.path.join(project_root, "data", "chroma_db") - chroma_db_file = os.path.join(chroma_dir, "chroma.sqlite3") - - # 检查数据 - raw_files = [f for f in os.listdir(raw_dir) if f.endswith((".txt", ".md"))] if os.path.exists(raw_dir) else [] - if not raw_files: - return {"error": "暂无数据,请用 generate_data.py 生成数据后刷新"} - - # 加载或构建向量库 - if os.path.exists(chroma_db_file): - vectorstore = load_vectorstore() - else: - docs = load_raw_documents() - chunks = chunk_documents(docs) - vectorstore = build_vectorstore(chunks) - - reranker = APIReranker() - - # 加载全部 chunk - from src.ingestion import rebuild_all_chunks - chunks = rebuild_all_chunks(raw_dir) - - # BM25 索引 - tokenized = [list(jieba.cut(d.page_content)) for d in chunks] - bm25 = BM25Okapi(tokenized) - - hybrid_retriever = HybridRetriever(vectorstore, chunks) - - def bm25_search(query: str, k: int = 3): - tokenized_query = list(jieba.cut(query)) - scores = bm25.get_scores(tokenized_query) - top_idx = sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)[:k] - return [chunks[i] for i in top_idx] - - graph = build_graph(vectorstore, bm25_search, hybrid_retriever, reranker=reranker) - - # 统计 - categories = list(set( - d.metadata.get("category", "未分类") - for d in chunks - )) - - return { - "error": None, - "vectorstore": vectorstore, - "chunks": chunks, - "bm25": bm25, - "hybrid_retriever": hybrid_retriever, - "bm25_search": bm25_search, - "graph": graph, - "reranker": reranker, - "raw_dir": raw_dir, - "chroma_dir": chroma_dir, - "stats": { - "categories": categories, - "total_notes": len(raw_files), - "total_chunks": len(chunks), - }, - } - - -def _run_insight(query: str) -> dict: - """执行洞察管道,返回报告""" - from src.agents.comment_agent import CommentAnalyzer - from src.agents.demand_agent import DemandAggregator - from src.agents.insight_agent import InsightGenerator - from src.config import RERANKER_THRESHOLD - from src.crawler import CrawlerInterface - - MIN_NOTES = 10 - CRAWL_COUNT = 30 # 最少爬取 30 篇 - - runtime = _runtime - hybrid_retriever = runtime["hybrid_retriever"] - reranker = runtime["reranker"] - raw_dir = runtime["raw_dir"] - - def _do_insight(docs, category): - analyzer = CommentAnalyzer(raw_dir=raw_dir) - analyses = analyzer.analyze(docs) - if not analyses: - return "没有找到评论分析数据。" - aggregator = DemandAggregator() - aggregated = aggregator.aggregate(analyses) - generator = InsightGenerator() - try: - report = generator.generate(aggregated, category=category) - except Exception as e: - report = generator.generate_fallback(aggregated, category=category) - report += f"\n\n(注:LLM 生成失败,使用模板兜底。错误:{e})" - return report - - def _rebuild_indexes(): - """增量入库 + 重建检索索引""" - from src.ingestion import incremental_ingest, rebuild_all_chunks - from rank_bm25 import BM25Okapi - import jieba - incremental_ingest(raw_dir, runtime["vectorstore"]) - chunks = rebuild_all_chunks(raw_dir) - tokenized = [list(jieba.cut(d.page_content)) for d in chunks] - bm25_new = BM25Okapi(tokenized) - hr_new = HybridRetriever(runtime["vectorstore"], chunks) - - def bm25_search_new(q, k=3): - scores = bm25_new.get_scores(list(jieba.cut(q))) - return [chunks[i] for i in sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)[:k]] - - runtime["chunks"] = chunks - runtime["bm25"] = bm25_new - runtime["hybrid_retriever"] = hr_new - runtime["bm25_search"] = bm25_search_new - - from src.graph import build_graph - runtime["graph"] = build_graph(runtime["vectorstore"], bm25_search_new, hr_new, reranker=reranker) - - # 检索 - docs = hybrid_retriever.hybrid_search(query, k=MIN_NOTES, bm25_k=40, final_k=MIN_NOTES) - if not docs: - docs = [] - - scores = reranker.rerank(query, docs) if docs else [] - relevant = [doc for doc, s in zip(docs, scores) if s >= RERANKER_THRESHOLD] - - crawled_count = 0 - if len(relevant) >= 3: - report = _do_insight(relevant, query) - else: - # 数据不足 → 自动用真实爬虫抓取小红书数据 - crawler = CrawlerInterface(raw_dir=raw_dir) - if not crawler.is_available: - return { - "report": f"知识库无「{query}」数据,且爬虫不可用。\n\n" - f"💡 请先在命令行运行 `uv run python src/real_crawler.py \"{query}\"` 登录并抓取数据。", - "notes_count": 0, - "generated_count": 0, - } - - result = crawler.crawl(query, count=CRAWL_COUNT) - crawled_count = result["count"] - - if crawled_count == 0: - return { - "report": f"抱歉,无法从小红书获取「{query}」的数据。\n" - f"请检查网络连接,或在命令行手动运行: uv run python src/real_crawler.py \"{query}\"", - "notes_count": 0, - "generated_count": 0, - } - - # 增量入库 + 重建索引 - _rebuild_indexes() - - time.sleep(0.5) - fresh_docs = runtime["hybrid_retriever"].hybrid_search(query, k=MIN_NOTES, bm25_k=40, final_k=MIN_NOTES) - fresh_scores = runtime["reranker"].rerank(query, fresh_docs) if fresh_docs else [] - 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, - } - - report = _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} - - -def _run_qa(question: str, strategy: str = "hybrid") -> str: - """执行 QA 管道""" - runtime = _runtime - graph = runtime["graph"] - - result = graph.invoke({ - "question": question, - "rewritten_question": "", - "strategy": strategy if strategy != "auto" else "", - "documents": [], - "relevant_docs": [], - "generation": "", - "retry_count": 0, - }) - response = result["generation"] - - # 没有答案 → 自动用真实爬虫抓取小红书数据 - if "无法回答" in response or "根据现有资料" in response: - from src.crawler import CrawlerInterface - crawler = CrawlerInterface(raw_dir=runtime["raw_dir"]) - - if not crawler.is_available: - return (f"{response}\n\n" - f"💡 知识库无相关数据。请先在命令行运行:\n" - f" `uv run python src/real_crawler.py \"{question}\"` 登录并抓取数据。") - - result = crawler.crawl(question, count=30) - count = result["count"] - - if count > 0: - # 增量入库 + 重建索引 - from src.ingestion import incremental_ingest, rebuild_all_chunks - from rank_bm25 import BM25Okapi - import jieba - incremental_ingest(runtime["raw_dir"], runtime["vectorstore"]) - chunks = rebuild_all_chunks(runtime["raw_dir"]) - tokenized = [list(jieba.cut(d.page_content)) for d in chunks] - bm25_new = BM25Okapi(tokenized) - hr = HybridRetriever(runtime["vectorstore"], chunks) - def bms(q, k=3): - scores = bm25_new.get_scores(list(jieba.cut(q))) - return [chunks[i] for i in sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)[:k]] - runtime["chunks"] = chunks - runtime["bm25"] = bm25_new - runtime["hybrid_retriever"] = hr - runtime["bm25_search"] = bms - from src.graph import build_graph - runtime["graph"] = build_graph(runtime["vectorstore"], bms, hr, reranker=runtime["reranker"]) - time.sleep(0.5) - fresh_graph = runtime["graph"] - result = fresh_graph.invoke({ - "question": question, - "rewritten_question": "", - "strategy": strategy if strategy != "auto" else "", - "documents": [], - "relevant_docs": [], - "generation": "", - "retry_count": 0, - }) - response = result["generation"] - if "无法回答" in response or "根据现有资料" in response: - response = f"(📥 已从小红书抓取 {count} 篇笔记,但检索仍未匹配)\n\n{response}" - else: - response = f"(📥 已从小红书实时抓取「{question}」{count} 篇真实笔记)\n\n{response}" - else: - response = (f"{response}\n\n" - f"💡 自动抓取失败。请手动运行:\n" - f" `uv run python src/real_crawler.py \"{question}\"`") - - return response - - -# ============================================================ -# API 路由 -# ============================================================ - -@app.get("/api/health") -async def health_check(): - return {"status": "ok", "version": "1.0.0"} - - -@app.get("/api/stats", response_model=StatsResponse) -async def get_stats(): - if _runtime is None or _runtime.get("error"): - return StatsResponse( - success=False, - categories=[], - total_notes=0, - total_chunks=0, - message=_runtime.get("error", "未初始化") if _runtime else "未初始化", - ) - stats = _runtime["stats"] - return StatsResponse( - success=True, - categories=stats["categories"], - total_notes=stats["total_notes"], - total_chunks=stats["total_chunks"], - message=f"知识库就绪,共 {len(stats['categories'])} 个品类", - ) - - -@app.post("/api/insight", response_model=InsightResponse) -async def run_insight(req: InsightRequest): - if _runtime is None or _runtime.get("error"): - raise HTTPException(status_code=503, detail=_runtime.get("error", "服务未就绪") if _runtime else "服务未就绪") - - t0 = time.time() - result = _run_insight(req.category) - elapsed = round(time.time() - t0, 2) - - return InsightResponse( - success=True, - category=req.category, - report=result["report"], - notes_count=result["notes_count"], - generated_count=result["generated_count"], - elapsed=elapsed, - ) - - -@app.post("/api/qa", response_model=QAResponse) -async def run_qa(req: QARequest): - if _runtime is None or _runtime.get("error"): - raise HTTPException(status_code=503, detail=_runtime.get("error", "服务未就绪") if _runtime else "服务未就绪") - - t0 = time.time() - answer = _run_qa(req.question, req.strategy) - elapsed = round(time.time() - t0, 2) - - return QAResponse( - success=True, - question=req.question, - answer=answer, - elapsed=elapsed, - ) - - -@app.post("/api/evaluate") -async def run_evaluation(req: EvaluateRequest): - """运行 RAGAS 评估并返回指标""" - if _runtime is None or _runtime.get("error"): - raise HTTPException(status_code=503, detail=_runtime.get("error", "服务未就绪") if _runtime else "服务未就绪") - - from src.evaluation import RAGEvaluator - evaluator = RAGEvaluator( - qa_func=_run_qa, - hybrid_retriever=_runtime["hybrid_retriever"], - reranker=_runtime["reranker"], - ) - - categories = req.categories or None - results = evaluator.evaluate(categories=categories) - - return JSONResponse(content={ - "success": True, - "evaluated_categories": results["categories"], - "total_questions": results["total_questions"], - "ragas_scores": results["ragas_scores"], - "timing_scores": results["timing_scores"], - "overall_score": results["overall_score"], - "grade": results["grade"], - }) - - -@app.post("/api/crawl") -async def trigger_crawl(req: CrawlRequest): - """触发数据抓取 — 优先使用真实爬虫,不可用时降级为 LLM 生成""" - from src.crawler import CrawlerInterface - crawler = CrawlerInterface(raw_dir=_runtime["raw_dir"]) - - result = crawler.crawl(req.category, req.count) - - if result["count"] > 0: - # 增量入库 + 重建索引 - from src.ingestion import incremental_ingest, rebuild_all_chunks - incremental_ingest(_runtime["raw_dir"], _runtime["vectorstore"]) - chunks = rebuild_all_chunks(_runtime["raw_dir"]) - from rank_bm25 import BM25Okapi - import jieba - tokenized = [list(jieba.cut(d.page_content)) for d in chunks] - _runtime["bm25"] = BM25Okapi(tokenized) - _runtime["chunks"] = chunks - _runtime["stats"]["total_chunks"] = len(chunks) - _runtime["stats"]["total_notes"] = len(os.listdir(_runtime["raw_dir"])) - - return JSONResponse(content={ - "success": result["count"] > 0, - "method": result["method"], - "count": result["count"], - "message": f"抓取完成: {result['count']} 篇" if result["count"] > 0 else "抓取失败", - }) - - -# ============================================================ -# 静态文件托管(前端 SPA) -# ============================================================ - -static_dir = Path(__file__).parent / "static" - - -@app.get("/") -async def serve_frontend(): - """托管前端页面""" - return FileResponse(static_dir / "index.html") - - -# 挂载静态资源 -if static_dir.exists(): - app.mount("/static", StaticFiles(directory=str(static_dir)), name="static") - - -# ============================================================ -# 启动入口 -# ============================================================ - -if __name__ == "__main__": - import uvicorn - print("RedNote Insight API starting...") - print(" API: http://localhost:8000") - print(" Front: http://localhost:8000") - print(" Docs: http://localhost:8000/docs") - uvicorn.run("api:app", host="0.0.0.0", port=8000, reload=True) +from src.api.main import app diff --git a/app.py b/app.py deleted file mode 100644 index 5c90e88..0000000 --- a/app.py +++ /dev/null @@ -1,455 +0,0 @@ -""" -app.py - 小红书爆款雷达(Streamlit 入口) -Phase 1: 基础 RAG 问答 -Phase 3: 评论区需求挖掘洞察模式 -Phase 4: 查询时自动抓取 — 知识库没有就现场生成 -""" -import streamlit as st -import sys -import os - -sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - -st.set_page_config(page_title="小红书爆款雷达", page_icon="🎯", layout="wide") -st.title("🎯 小红书爆款雷达") -st.markdown("---") - - -# ====================================================================== -# 第一部分:一次性初始化(缓存) -# ====================================================================== - -@st.cache_resource -def init_base(): - """缓存:Embedding / Reranker 等不需要随数据变化而变化的对象""" - from src.ingestion import load_raw_documents, chunk_documents, load_vectorstore, build_vectorstore, rebuild_all_chunks - from src.retrievers import APIReranker - - project_root = os.path.dirname(os.path.abspath(__file__)) - raw_dir = os.path.join(project_root, "data", "raw") - chroma_dir = os.path.join(project_root, "data", "chroma_db") - chroma_db_file = os.path.join(chroma_dir, "chroma.sqlite3") - - # 检查数据是否存在 - raw_files = [f for f in os.listdir(raw_dir) if f.endswith((".txt", ".md"))] if os.path.exists(raw_dir) else [] - if not raw_files: - return None, "暂无数据,请用 generate_data.py 生成数据后刷新页面。" - - # 加载或构建向量库 - if os.path.exists(chroma_db_file): - vectorstore = load_vectorstore() - else: - docs = load_raw_documents() - chunks = chunk_documents(docs) - vectorstore = build_vectorstore(chunks) - - # Reranker(CrossEncoder API,不随数据变化) - reranker = APIReranker() - - return { - "vectorstore": vectorstore, - "reranker": reranker, - "raw_dir": raw_dir, - "chroma_dir": chroma_dir, - }, None - - -# ====================================================================== -# 第二部分:可变运行时状态(存储在 session_state,支持动态更新) -# ====================================================================== - -def build_runtime(base: dict): - """从当前磁盘数据构建 BM25 / HybridRetriever / LangGraph""" - from src.ingestion import rebuild_all_chunks - from src.retrievers import HybridRetriever - from src.graph import build_graph - from rank_bm25 import BM25Okapi - import jieba - - vectorstore = base["vectorstore"] - raw_dir = base["raw_dir"] - reranker = base["reranker"] - - # 加载全部文档 + chunk - chunks = rebuild_all_chunks(raw_dir) - - # BM25 索引 - tokenized = [list(jieba.cut(d.page_content)) for d in chunks] - bm25 = BM25Okapi(tokenized) - - # HybridRetriever - hybrid_retriever = HybridRetriever(vectorstore, chunks) - - # BM25 搜索函数 - def bm25_search(query: str, k: int = 3): - tokenized_query = list(jieba.cut(query)) - scores = bm25.get_scores(tokenized_query) - top_idx = sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)[:k] - return [chunks[i] for i in top_idx] - - # LangGraph - graph = build_graph(vectorstore, bm25_search, hybrid_retriever, reranker=reranker) - - return { - "chunks": chunks, - "bm25": bm25, - "hybrid_retriever": hybrid_retriever, - "bm25_search": bm25_search, - "graph": graph, - } - - -def reload_after_fetch(): - """ - 当 fetcher 写入了新数据后调用此函数: - 增量入库 → 重建 chunk → 重建 BM25/Hybrid/Graph - """ - from src.ingestion import incremental_ingest, rebuild_all_chunks - from src.retrievers import HybridRetriever - from src.graph import build_graph - from rank_bm25 import BM25Okapi - import jieba - - base = st.session_state.base - vectorstore = base["vectorstore"] - raw_dir = base["raw_dir"] - reranker = base["reranker"] - - # 增量入库 - incremental_ingest(raw_dir, vectorstore) - - # 重建全部 chunks(新老数据一起) - chunks = rebuild_all_chunks(raw_dir) - - # BM25 - tokenized = [list(jieba.cut(d.page_content)) for d in chunks] - bm25 = BM25Okapi(tokenized) - - # Hybrid - hybrid_retriever = HybridRetriever(vectorstore, chunks) - - def bm25_search(query: str, k: int = 3): - tokenized_query = list(jieba.cut(query)) - scores = bm25.get_scores(tokenized_query) - top_idx = sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)[:k] - return [chunks[i] for i in top_idx] - - # Graph - graph = build_graph(vectorstore, bm25_search, hybrid_retriever, reranker=reranker) - - # 更新 session_state - st.session_state.runtime = { - "chunks": chunks, - "bm25": bm25, - "hybrid_retriever": hybrid_retriever, - "bm25_search": bm25_search, - "graph": graph, - } - st.session_state.data_version += 1 - - -# ====================================================================== -# 初始化入口 -# ====================================================================== - -base, error = init_base() -if error: - st.warning(error) - st.info("提示: 运行 `python generate_data.py` 生成演示数据。") - st.stop() - -# 保持 base 在 session_state(供 reload_after_fetch 使用) -if "base" not in st.session_state: - st.session_state.base = base - -# 初次或刷新时构建运行时 -if "runtime" not in st.session_state: - st.session_state.runtime = build_runtime(base) - st.session_state.data_version = 0 - -runtime = st.session_state.runtime -graph = runtime["graph"] -hybrid_retriever = runtime["hybrid_retriever"] -chunks = runtime["chunks"] -raw_dir = base["raw_dir"] -reranker = base["reranker"] - - -# ====================================================================== -# 第三部分:洞察管道(核心变更:无匹配 → 自动抓取) -# ====================================================================== - -def run_insight_pipeline(query: str, status_placeholder=None, stream: bool = False): - """ - 完整的洞察流程。 - stream=False: 返回完整字符串(API 模式) - stream=True: 返回生成器,在生成阶段逐 token yield(Streamlit 流式输出) - """ - from src.agents.comment_agent import CommentAnalyzer - from src.agents.demand_agent import DemandAggregator - from src.agents.insight_agent import InsightGenerator - from src.config import RERANKER_THRESHOLD - - MIN_NOTES = 20 - generator = InsightGenerator() - - def _do_insight(docs, category): - """文档 → 分析 → 聚合 → 报告(非流式)""" - analyzer = CommentAnalyzer(raw_dir=raw_dir) - analyses = analyzer.analyze(docs) - if not analyses: - return "没有找到评论分析数据。" if not stream else iter(["没有找到评论分析数据。"]) - aggregator = DemandAggregator() - aggregated = aggregator.aggregate(analyses) - if stream: - return generator.generate_stream(aggregated, category=category) - try: - report = generator.generate(aggregated, category=category) - except Exception as e: - report = generator.generate_fallback(aggregated, category=category) - report += f"\n\n(注:LLM 生成失败,使用模板兜底。错误:{e})" - return report - - # 1. 扩大检索范围(从 10 → 20 篇) - docs = hybrid_retriever.hybrid_search(query, k=MIN_NOTES, bm25_k=40, final_k=MIN_NOTES) - if not docs: - return "检索失败,请刷新页面重试。" - - # 2. CrossEncoder 过滤 - scores = reranker.rerank(query, docs) - relevant = [doc for doc, s in zip(docs, scores) if s >= RERANKER_THRESHOLD] - - if len(relevant) >= MIN_NOTES: - # ✅ 有足够数据(≥20 篇),正常走管道 - return _do_insight(relevant, query) - - # ❌ 数据不足(< 20 篇)→ 🔥 触发真实爬虫从小红书抓取 - current_count = len(relevant) - fetch_target = max(MIN_NOTES - current_count + 5, 30) - - if status_placeholder: - status_placeholder.info( - f"🔍 品类「**{query}**」当前只有 {current_count} 篇相关笔记," - f"正在从小红书实时抓取 {fetch_target} 篇笔记..." - ) - - from src.crawler import CrawlerInterface - - cookies_json = "" - try: - cookies_json = st.secrets.get("XHS_COOKIES", "") - except Exception: - pass - - crawler = CrawlerInterface(raw_dir=raw_dir, cookies_json=cookies_json) - if not crawler.is_available: - if crawler.is_cloud: - return (f"知识库无「{query}」数据,且云端爬虫未配置 Cookie。\n\n" - f"💡 本地运行 `uv run python scripts/export_cookies.py` 导出 cookie," - f"粘贴到 Streamlit Secrets → XHS_COOKIES") - return (f"知识库无「{query}」数据,且爬虫未登录。\n\n" - f"💡 请先在命令行运行: `uv run python src/real_crawler.py \"{query}\"` 登录后重试。") - - result = crawler.crawl(query, count=fetch_target) - count = result["count"] - - if count == 0: - return f"抱歉,无法从小红书获取「{query}」的数据。请检查网络连接后重试。" - - if status_placeholder: - status_placeholder.success(f"✅ 已从小红书抓取 {count} 篇「{query}」真实笔记,正在入库并分析...") - - # 增量入库 + 重建索引 - reload_after_fetch() - - # 用更新后的 retriever 重新查询 - import time - time.sleep(0.5) # 等 chromadb 落盘 - fresh_docs = st.session_state.runtime["hybrid_retriever"].hybrid_search( - query, k=MIN_NOTES, bm25_k=40, final_k=MIN_NOTES - ) - fresh_scores = reranker.rerank(query, fresh_docs) - fresh_relevant = [doc for doc, s in zip(fresh_docs, fresh_scores) if s >= RERANKER_THRESHOLD] - - if not fresh_relevant: - return f"已从小红书抓取 {count} 篇「{query}」笔记,但检索仍未匹配。请稍后重试或更换关键词。" - - # 用新数据生成洞察 - report = _do_insight(fresh_relevant, query) - report = ( - f"(📥 已从小红书实时抓取「{query}」{count} 篇真实笔记," - f"当前共 {len(fresh_relevant)} 篇相关笔记)\n\n{report}" - ) - return report - - -# ====================================================================== -# 第四部分:Streamlit UI -# ====================================================================== - -# ---- 侧边栏 ---- -with st.sidebar: - st.subheader("模式选择") - mode = st.radio( - "运行模式", - ["问答模式", "洞察模式"], - index=0, - help="问答模式:基于知识库回答问题。洞察模式:分析评论区挖掘选品机会。", - ) - - st.markdown("---") - st.caption( - f"📊 当前知识库:{len(chunks)} 个 chunk" - + (f" · 🆕 有新数据" if st.session_state.data_version > 0 else "") - ) - - st.markdown("**使用提示**") - if mode == "问答模式": - st.caption( - "输入产品相关的问题,例如:\n" - "- 磁吸感应灯哪个品牌好\n" - "- 学生寝室平价好物推荐\n" - "- 收纳盒怎么选" - ) - else: - st.caption( - "输入品类名称获取市场洞察,例如:\n" - "- 磁吸感应灯\n" - "- 寝室改造\n" - "- 桌面收纳\n" - "- 健身服(知识库没有?自动抓取!)" - ) - - if mode == "问答模式": - st.markdown("---") - st.subheader("检索策略") - strategy = st.radio( - "策略", - ["auto", "vector", "keyword", "hybrid"], - index=0, - help="auto: Supervisor 自动选择", - label_visibility="collapsed", - ) - - -# ---- 主界面 ---- -if mode == "问答模式": - # ========== 问答模式 ========== - st.subheader("💬 问答") - - if "qa_messages" not in st.session_state: - st.session_state.qa_messages = [] - - for msg in st.session_state.qa_messages: - with st.chat_message(msg["role"]): - st.markdown(msg["content"]) - - if prompt := st.chat_input("输入你的问题..."): - st.session_state.qa_messages.append({"role": "user", "content": prompt}) - with st.chat_message("user"): - st.markdown(prompt) - - with st.chat_message("assistant"): - status = st.empty() - with st.spinner("思考中..."): - result = graph.invoke({ - "question": prompt, - "rewritten_question": "", - "strategy": strategy if strategy != "auto" else "", - "documents": [], - "relevant_docs": [], - "generation": "", - "retry_count": 0, - }) - response = result["generation"] - - # 🚀 如果没有答案 → 自动从小红书抓取真实数据 → 重新检索回答 - if "无法回答" in response or "根据现有资料" in response: - from src.crawler import CrawlerInterface - - category = prompt # 直接用问题作为品类名 - status.info(f"🔍 知识库中暂无「{category}」相关信息,正在从小红书实时抓取...") - - cookies_json = "" - try: - cookies_json = st.secrets.get("XHS_COOKIES", "") - except Exception: - pass - - crawler = CrawlerInterface(raw_dir=raw_dir, cookies_json=cookies_json) - if not crawler.is_available: - status.warning("⚠️ 爬虫未登录,请先在命令行运行: uv run python src/real_crawler.py \"品类名\"") - else: - result = crawler.crawl(category, count=30) - count = result["count"] - - if count > 0: - status.success(f"✅ 已从小红书抓取 {count} 篇「{category}」真实笔记,正在重新检索回答...") - reload_after_fetch() - - import time - time.sleep(0.5) - - # 使用更新后的 graph 重新问答 - fresh_graph = st.session_state.runtime["graph"] - result = fresh_graph.invoke({ - "question": prompt, - "rewritten_question": "", - "strategy": strategy if strategy != "auto" else "", - "documents": [], - "relevant_docs": [], - "generation": "", - "retry_count": 0, - }) - response = result["generation"] - - if "无法回答" in response or "根据现有资料" in response: - response = ( - f"(📥 已从小红书抓取 {count} 篇真实笔记," - f"但检索仍未匹配到相关信息)\n\n{response}" - ) - else: - response = ( - f"(📥 已从小红书实时抓取 {count} 篇真实笔记作为知识补充)\n\n{response}" - ) - else: - response = f"抱歉,无法从小红书获取「{category}」的数据。请检查网络连接后重试。" - - st.markdown(response) - - st.session_state.qa_messages.append({"role": "assistant", "content": response}) - -else: - # ========== 洞察模式 ========== - st.subheader("📊 选品洞察") - - if "insight_messages" not in st.session_state: - st.session_state.insight_messages = [] - - for msg in st.session_state.insight_messages: - with st.chat_message(msg["role"]): - st.markdown(msg["content"]) - - if prompt := st.chat_input("输入品类名称,例如:磁吸感应灯、健身服..."): - st.session_state.insight_messages.append({"role": "user", "content": prompt}) - with st.chat_message("user"): - st.markdown(prompt) - - with st.chat_message("assistant"): - status = st.empty() - report_container = st.empty() - with st.spinner("分析评论区数据中..."): - stream_gen = run_insight_pipeline(prompt, status_placeholder=status, stream=True) - # 流式输出 - full_report = "" - for chunk in stream_gen: - full_report += chunk - report_container.markdown(full_report + "▌") - report_container.markdown(full_report) - - st.session_state.insight_messages.append({"role": "assistant", "content": full_report}) - - -# ---- 底部 ---- -st.markdown("---") -st.caption(f"🎯 小红书爆款雷达 v0.3 · 问答 + 洞察 + 自动抓取 · 数据版本 {st.session_state.data_version}") diff --git a/data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/length.bin b/data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/length.bin deleted file mode 100644 index 15b1c0c..0000000 Binary files a/data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/length.bin and /dev/null differ diff --git a/data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/data_level0.bin b/data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/data_level0.bin similarity index 99% rename from data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/data_level0.bin rename to data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/data_level0.bin index 4c42049..b842f97 100644 Binary files a/data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/data_level0.bin and b/data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/data_level0.bin differ diff --git a/data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/header.bin b/data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/header.bin similarity index 100% rename from data/chroma_db/9650ba44-355b-43d2-bf46-374351a47dab/header.bin rename to data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/header.bin diff --git a/data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/length.bin b/data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/length.bin new file mode 100644 index 0000000..4b7540e Binary files /dev/null and b/data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/length.bin differ diff --git a/data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/link_lists.bin b/data/chroma_db/a25555c6-5602-4ba8-ac59-d9c3bbf9af5a/link_lists.bin new file mode 100644 index 0000000..e69de29 diff --git a/data/chroma_db/chroma.sqlite3 b/data/chroma_db/chroma.sqlite3 index 9a80345..86f214f 100644 Binary files a/data/chroma_db/chroma.sqlite3 and b/data/chroma_db/chroma.sqlite3 differ diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_66e806c9_66e806c9.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_66e806c9_66e806c9.md" new file mode 100644 index 0000000..b70fd1a --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_66e806c9_66e806c9.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_66e806c9 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_66f05644_66f05644.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_66f05644_66f05644.md" new file mode 100644 index 0000000..240f0be --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_66f05644_66f05644.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_66f05644 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_66fbcd5e_66fbcd5e.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_66fbcd5e_66fbcd5e.md" new file mode 100644 index 0000000..f34a0be --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_66fbcd5e_66fbcd5e.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_66fbcd5e +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_671953bf_671953bf.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_671953bf_671953bf.md" new file mode 100644 index 0000000..d351959 --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_671953bf_671953bf.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_671953bf +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_6747fe00_6747fe00.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_6747fe00_6747fe00.md" new file mode 100644 index 0000000..af4b48a --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_6747fe00_6747fe00.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_6747fe00 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_6756b34e_6756b34e.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_6756b34e_6756b34e.md" new file mode 100644 index 0000000..2a78f4e --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_6756b34e_6756b34e.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_6756b34e +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_67a1cc3d_67a1cc3d.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_67a1cc3d_67a1cc3d.md" new file mode 100644 index 0000000..1eeef4d --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_67a1cc3d_67a1cc3d.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_67a1cc3d +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_67b332da_67b332da.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_67b332da_67b332da.md" new file mode 100644 index 0000000..4eb76c8 --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_67b332da_67b332da.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_67b332da +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_68d647e2_68d647e2.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_68d647e2_68d647e2.md" new file mode 100644 index 0000000..4a16511 --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_68d647e2_68d647e2.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_68d647e2 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_68e1e1d8_68e1e1d8.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_68e1e1d8_68e1e1d8.md" new file mode 100644 index 0000000..9c48d4e --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_68e1e1d8_68e1e1d8.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_68e1e1d8 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_6927199e_6927199e.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_6927199e_6927199e.md" new file mode 100644 index 0000000..453037a --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_6927199e_6927199e.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_6927199e +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_693cdee3_693cdee3.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_693cdee3_693cdee3.md" new file mode 100644 index 0000000..1367ae7 --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_693cdee3_693cdee3.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_693cdee3 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_6971f12b_6971f12b.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_6971f12b_6971f12b.md" new file mode 100644 index 0000000..69e98ff --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_6971f12b_6971f12b.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_6971f12b +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_6981c746_6981c746.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_6981c746_6981c746.md" new file mode 100644 index 0000000..9fe60cc --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_6981c746_6981c746.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_6981c746 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\346\211\213\346\234\272\345\243\263_69aaa4f4_69aaa4f4.md" "b/data/raw/\346\211\213\346\234\272\345\243\263_69aaa4f4_69aaa4f4.md" new file mode 100644 index 0000000..76768d6 --- /dev/null +++ "b/data/raw/\346\211\213\346\234\272\345\243\263_69aaa4f4_69aaa4f4.md" @@ -0,0 +1,39 @@ +--- +author: '' +brand: 手机壳 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-26' +likes: 0 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 手机壳 +- '' +title: 手机壳_69aaa4f4 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git "a/data/raw/\350\243\205\351\245\260\347\224\273_6a13cb70_6a13cb70.md" "b/data/raw/\350\243\205\351\245\260\347\224\273_6a13cb70_6a13cb70.md" new file mode 100644 index 0000000..d26f2bb --- /dev/null +++ "b/data/raw/\350\243\205\351\245\260\347\224\273_6a13cb70_6a13cb70.md" @@ -0,0 +1,39 @@ +--- +author: 奶芙芙的🏠 +brand: 装饰画 +category_type: 常青款 +comments: 0 +cost: 0 +date: '2026-06-30' +likes: 417 +price: 0 +return_rate: 0.05 +size: '' +tags: +- 装饰画 +- 奶芙芙的🏠 +title: 装饰画_6a13cb70 +weight: 0.5 +--- + + + +--- + \ No newline at end of file diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..6317a88 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,83 @@ +# ============================================================================= +# RedNote Insight — Docker Compose 编排 +# ============================================================================= +# 用法: +# cp .env.example .env && vim .env # 填入 API Key +# docker-compose up -d # 一键启动 +# docker-compose logs -f api # 查看日志 +# docker-compose down # 停止 +# ============================================================================= + +version: "3.8" + +services: + # ===== API 服务 ===== + api: + build: . + container_name: rednote-api + ports: + - "8000:8000" + env_file: + - .env + environment: + - DATABASE_URL=postgresql+asyncpg://postgres:postgres@db:5432/rednote_insight + - REDIS_URL=redis://redis:6379/0 + depends_on: + db: + condition: service_healthy + redis: + condition: service_started + volumes: + - api_data:/app/data + restart: unless-stopped + networks: + - rednote-net + + # ===== PostgreSQL + pgvector ===== + db: + image: pgvector/pgvector:pg16 + container_name: rednote-db + ports: + - "5432:5432" + environment: + POSTGRES_USER: postgres + POSTGRES_PASSWORD: postgres + POSTGRES_DB: rednote_insight + volumes: + - pg_data:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U postgres -d rednote_insight"] + interval: 10s + timeout: 5s + retries: 5 + start_period: 10s + restart: unless-stopped + networks: + - rednote-net + + # ===== Redis(缓存 / 限流)===== + redis: + image: redis:7-alpine + container_name: rednote-redis + ports: + - "6379:6379" + volumes: + - redis_data:/data + command: redis-server --appendonly yes --maxmemory 128mb --maxmemory-policy allkeys-lru + restart: unless-stopped + networks: + - rednote-net + +# ===== 持久化卷 ===== +volumes: + pg_data: + driver: local + redis_data: + driver: local + api_data: + driver: local + +# ===== 网络 ===== +networks: + rednote-net: + driver: bridge diff --git a/import_data.py b/import_data.py deleted file mode 100644 index ff9a2fa..0000000 --- a/import_data.py +++ /dev/null @@ -1,423 +0,0 @@ -""" -import_data.py - 真实数据导入工具 -===================================== -将 CSV/Excel 中的笔记数据导入为 .md 格式, -兼容现有 RAG 问答 + 洞察管道。 - -用法: - # 查看 CSV 格式说明 - python import_data.py --help - - # 导入 CSV(自动生成评论分析数据) - python import_data.py --input my_data.csv - - # 导入 + 用 LLM 丰富评论分析(需 API Key) - python import_data.py --input my_data.csv --enrich - - # 导入后自动重建向量库 - python import_data.py --input my_data.csv --rebuild - python import_data.py --input my_data.xlsx --sheet Sheet1 --rebuild - -输入格式: - title,content,brand,likes,date,tags,comments,author - 标题,笔记正文,品牌名,点赞数,日期,标签|逗号分隔,评论数,作者名 - -只有 title 和 content 是必填,其余缺失会自动填充默认值。 -""" -import os -import sys -import csv -import json -import random -import argparse -import re -from pathlib import Path -from typing import List, Optional, Dict, Any -from datetime import datetime - -# 确保能找到 src 包 -sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - - -# ============================================================ -# 评论分析数据自动生成(基于内容关键词) -# ============================================================ - -# 关键词 → 常见投诉映射 -COMPLAINT_KEYWORDS = { - "续航": ["续航不够持久", "充电太频繁"], - "价格": ["价格偏高", "性价比一般"], - "质量": ["质量一般", "用不久就坏了"], - "材质": ["材质一般", "手感不好"], - "大小": ["尺寸偏小", "比想象中小"], - "颜色": ["颜色和图片有差异", "色差严重"], - "安装": ["安装不方便", "安装说明不清晰"], - "充电": ["充电太慢", "续航不够持久"], - "磁吸": ["磁吸不够牢固", "容易掉"], - "灯": ["亮度不够", "灯光刺眼"], - "声音": ["噪音有点大", "运行声音明显"], - "容量": ["容量太小", "装不了多少东西"], - "设计": ["设计一般", "不够美观"], - "包装": ["包装简陋", "包装破损"], - "售后": ["售后服务差", "退货麻烦"], -} - -# 关键词 → 常见需求信号映射 -INTENT_KEYWORDS = { - "学生": ["适合学生党吗", "性价比怎么样"], - "寝室": ["宿舍能用吗", "查寝会被扣分吗"], - "租房": ["适合租房党吗", "搬家好带走吗"], - "礼物": ["送人合适吗", "包装好看吗"], - "新手": ["新手适合吗", "操作难不难"], - "卧室": ["适合卧室用吗", "什么色温适合卧室"], - "厨房": ["防水吗", "厨房能用吗"], - "礼物": ["送人合适吗", "有礼品包装吗"], - "质量": ["质量怎么样", "耐不耐用"], - "价格": ["有优惠吗", "什么时候降价"], - "尺寸": ["尺寸多大", "能放下吗"], - "颜色": ["有什么颜色", "哪个颜色好看"], -} - -# 通用评论高频词 -DEFAULT_HIGH_FREQ = ["求链接", "好用吗", "收藏了", "什么牌子", "多少钱"] - -# 通用品牌对比 -DEFAULT_COMPARISONS = ["比其他品牌性价比高", "比线下便宜"] - - -def infer_comment_analysis(title: str, content: str, brand: str) -> Dict: - """从笔记标题+内容中推断可能的评论分析数据""" - combined = (title + " " + content).lower() - complaints = [] - intents = [] - - # 匹配关键词 - for keyword, complaint_list in COMPLAINT_KEYWORDS.items(): - if keyword in combined: - complaints.append(random.choice(complaint_list)) - - for keyword, intent_list in INTENT_KEYWORDS.items(): - if keyword in combined: - intents.append(random.choice(intent_list)) - - # 保底:至少一条 - if not complaints: - complaints.append(random.choice(list(COMPLAINT_KEYWORDS.values()))[0]) - if not intents: - intents.append(random.choice(list(INTENT_KEYWORDS.values()))[0]) - - # 去重 - complaints = list(dict.fromkeys(complaints)) - intents = list(dict.fromkeys(intents)) - - return { - "high_freq_words": DEFAULT_HIGH_FREQ.copy(), - "complaints": complaints[:5], - "purchase_intent": intents[:5], - "comparison_mentions": [f"比{brand}便宜多了"] if brand else DEFAULT_COMPARISONS, - "related_brands": [brand] if brand else [], - "ask_link_count": random.randint(30, 200), - } - - -def enrich_with_llm(title: str, content: str, brand: str, api_key: str = None) -> Optional[Dict]: - """用 LLM 从内容中提取评论分析(需要设置 LLM API Key)""" - try: - from langchain_openai import ChatOpenAI - from langchain_core.prompts import ChatPromptTemplate - from dotenv import load_dotenv - - load_dotenv() - llm = ChatOpenAI( - model=os.getenv("LLM_MODEL", "deepseek-ai/DeepSeek-V4-Flash"), - temperature=0.1, - api_key=api_key or os.getenv("OPENAI_API_KEY"), - base_url=os.getenv("OPENAI_BASE_URL"), - ) - - prompt = ChatPromptTemplate.from_messages([ - ("system", "你是小红书评论分析专家。根据笔记标题和内容," - "推断用户可能在评论区讨论什么。返回 JSON 格式:\n" - '{"complaints": ["投诉1", "投诉2"],' - ' "purchase_intent": ["需求1", "需求2"],' - ' "comparison_mentions": ["对比提及1"]}\n' - "不要解释,只返回 JSON。"), - ("human", "标题:{title}\n内容:{content}"), - ]) - - msg = prompt.format_messages(title=title, content=content[:500]) - result = llm.invoke(msg).content.strip() - # 提取 JSON - json_match = re.search(r"\{.*\}", result, re.DOTALL) - if json_match: - data = json.loads(json_match.group()) - data["high_freq_words"] = DEFAULT_HIGH_FREQ.copy() - data["related_brands"] = [brand] if brand else [] - data["ask_link_count"] = random.randint(30, 200) - return data - except Exception as e: - print(f" LLM 分析失败: {e}") - return None - - -# ============================================================ -# 数据导入 -# ============================================================ - -def parse_value(value: str) -> Any: - """智能解析 CSV 中的值""" - if value is None: - return None - value = value.strip() - if not value: - return None - # 数字 - try: - return int(value) - except ValueError: - pass - try: - return float(value) - except ValueError: - pass - # 布尔 - if value.lower() in ("true", "yes"): - return True - if value.lower() in ("false", "no"): - return False - return value - - -def read_csv(filepath: str) -> List[Dict]: - """读取 CSV 文件""" - records = [] - with open(filepath, "r", encoding="utf-8-sig") as f: - reader = csv.DictReader(f) - for row in reader: - record = {k: parse_value(v) for k, v in row.items()} - records.append(record) - return records - - -def read_excel(filepath: str, sheet: str = None) -> List[Dict]: - """读取 Excel 文件""" - try: - import pandas as pd - except ImportError: - print("[错误] 读取 Excel 需要安装 pandas:pip install pandas openpyxl") - sys.exit(1) - - if sheet: - df = pd.read_excel(filepath, sheet_name=sheet) - else: - df = pd.read_excel(filepath) - return df.to_dict(orient="records") - - -def generate_md( - record: Dict, - output_dir: str, - index: int, - use_llm: bool = False, -) -> str: - """将一条记录转换为 .md 文件内容""" - # 字段名大小写兼容 - title = str(record.get("title") or record.get("Title") or f"笔记{index}") - content = str(record.get("content") or record.get("Content") or record.get("正文", "")) - brand = str(record.get("brand") or record.get("Brand") or record.get("品牌", "")) - likes = int(record.get("likes") or record.get("Likes") or record.get("点赞", 0) or random.randint(100, 500)) - date_val = record.get("date") or record.get("Date") or record.get("日期", "") - tags_raw = record.get("tags") or record.get("Tags") or record.get("标签", "") - author = record.get("author") or record.get("Author") or record.get("作者", f"小红书用户{random.randint(1000,9999)}") - comments_count = record.get("comments") or record.get("Comments") or record.get("评论数", 0) or int(likes * random.uniform(0.15, 0.35)) - - # 标签解析(CSV 中可能是逗号分隔或竖线分隔) - if isinstance(tags_raw, str): - sep = "|" if "|" in tags_raw else "," - tags = [t.strip() for t in tags_raw.split(sep) if t.strip()] - elif isinstance(tags_raw, list): - tags = tags_raw - else: - tags = [brand, "小红书好物"] - - # 日期格式统一 - if not date_val: - date_val = f"2025-{random.randint(1,5):02d}-{random.randint(1,28):02d}" - else: - try: - date_val = str(pd.Timestamp(date_val).date()) - except Exception: - pass - - # 评论分析数据 - if use_llm: - ca_data = enrich_with_llm(title, content, brand) - else: - ca_data = infer_comment_analysis(title, content, brand) - - if ca_data is None: - ca_data = infer_comment_analysis(title, content, brand) - - # 文件名:品类_序号.md - category_slug = brand[:2] if brand else "import" - filename = f"{category_slug}_{index:03d}.md" - filepath = os.path.join(output_dir, filename) - - # 处理 content 中可能包含的 ---(会被 YAML 解析器误读) - content_clean = content.replace("---", "—") - - # 组装文件 - parts = [ - "---\n", - f'title: "{title}"\n', - f'author: "{author}"\n', - f"likes: {likes}\n", - f"comments: {comments_count}\n", - f"date: {date_val}\n", - f'brand: "{brand}"\n', - f"tags: {tags}\n", - "---\n\n", - content_clean, - "\n---\n", - "\n", - ] - - return "".join(parts), filename - - -def rebuild_vectorstore(): - """重建向量库""" - print("\n[重建] 重建向量库...") - try: - from src.ingestion import load_raw_documents, chunk_documents, build_vectorstore - docs = load_raw_documents() - chunks = chunk_documents(docs) - build_vectorstore(chunks) - print(f"[重建] 完成:{len(chunks)} 个文档已向量化") - except Exception as e: - print(f"[重建失败] {e}") - print("你可以稍后手动重建:python -c 'from src.ingestion import *; rebuild()'") - - -# ============================================================ -# CLI -# ============================================================ - -def main(): - parser = argparse.ArgumentParser( - description="将 CSV/Excel 数据导入为小红书笔记 .md 格式", - formatter_class=argparse.RawDescriptionHelpFormatter, - epilog=""" -示例 CSV 格式(utf-8 编码): - title,content,brand,likes,date,tags,comments,author - 瑜伽裤测评,这条瑜伽裤真的绝了...,lululemon,534,2025-03-01,运动|瑜伽,128,小雅 - 平价健身服推荐,学生党必看的健身服...,Alo,312,2025-02-15,健身|平价,56,阿宁 - -字段说明: - title* 笔记标题(必填) - content* 笔记正文(必填) - brand 品牌名 - likes 点赞数 - date 发布日期 - tags 标签,竖线 | 分隔 - comments 评论数 - author 作者名 - -导入后,在洞察模式输入品类名即可生成选品报告。 - """, - ) - parser.add_argument("--input", "-i", default=None, help="CSV 或 Excel 文件路径") - parser.add_argument("--sheet", "-s", help="Excel 工作表名(默认第一页)") - parser.add_argument("--output-dir", "-o", default=None, - help="输出目录(默认 data/raw)") - parser.add_argument("--enrich", action="store_true", - help="用 LLM 从内容中提取评论分析(需配置 API Key)") - parser.add_argument("--rebuild", action="store_true", - help="导入后重建向量库") - parser.add_argument("--sample", action="store_true", - help="生成示例 CSV 文件到当前目录") - - args = parser.parse_args() - - # ---- 既不是示例也不是导入 ---- - if not args.sample and not args.input: - parser.print_help() - print("\n使用 --sample 生成示例 CSV,或使用 --input 导入数据。") - return - - # ---- 生成示例 ---- - if args.sample: - sample_path = os.path.join(os.getcwd(), "sample_data.csv") - with open(sample_path, "w", encoding="utf-8-sig", newline="") as f: - writer = csv.writer(f) - writer.writerow(["title", "content", "brand", "likes", "date", "tags", "comments", "author"]) - writer.writerow(["瑜伽裤测评", "这条瑜伽裤真的太绝了!弹性超好,包裹感强,深蹲完全不会透。", "lululemon", "534", "2025-03-01", "运动|瑜伽|健身", "128", "小雅"]) - writer.writerow(["平价健身服推荐", "学生党必看!百元以内的健身服分享,透气舒适,适合健身房。", "Alo", "312", "2025-02-15", "健身|平价|学生党", "56", "阿宁"]) - writer.writerow(["跑步鞋开箱", "新入的跑鞋太香了,减震效果很好,马拉松训练穿它。", "Nike", "678", "2025-01-20", "跑步|运动|开箱", "203", "跑者小王"]) - print(f"[示例] 已生成示例文件:{sample_path}") - print("[示例] 参考此格式准备你的数据,然后运行:") - print(f" python import_data.py --input {sample_path}") - return - - # ---- 导入 ---- - filepath = args.input - if not os.path.exists(filepath): - print(f"[错误] 文件不存在:{filepath}") - sys.exit(1) - - # 确定输出目录 - project_root = os.path.dirname(os.path.abspath(__file__)) - output_dir = args.output_dir or os.path.join(project_root, "data", "raw") - os.makedirs(output_dir, exist_ok=True) - - # 读取 - ext = os.path.splitext(filepath)[1].lower() - if ext in (".xlsx", ".xls"): - print(f"[读取] Excel: {filepath}") - records = read_excel(filepath, args.sheet) - elif ext == ".csv": - print(f"[读取] CSV: {filepath}") - records = read_csv(filepath) - elif ext == ".json": - with open(filepath, "r", encoding="utf-8") as f: - records = json.load(f) - else: - print(f"[错误] 不支持的文件格式:{ext},请使用 CSV 或 Excel") - sys.exit(1) - - # 生成 - print(f"[导入] 共 {len(records)} 条记录") - if args.enrich: - print("[导入] LLM 评论分析已开启(每个笔记需要一次 API 调用)") - - count = 0 - for i, record in enumerate(records): - try: - md_content, filename = generate_md(record, output_dir, i + 1, use_llm=args.enrich) - filepath = os.path.join(output_dir, filename) - with open(filepath, "w", encoding="utf-8") as f: - f.write(md_content) - print(f" + {filename}") - count += 1 - except Exception as e: - print(f" ✗ 第 {i+1} 行导入失败: {e}") - - print(f"\n=> 成功导入 {count}/{len(records)} 篇笔记 -> {output_dir}") - - if args.rebuild: - rebuild_vectorstore() - else: - print("\n提示: 使用 --rebuild 参数重建向量库,或删除 data/chroma_db 目录后重启应用。") - print("启动应用:streamlit run app.py") - - -if __name__ == "__main__": - main() diff --git a/pyproject.toml b/pyproject.toml index 40a10bb..be7abf9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "rednote-insight" -version = "0.3.0" +version = "1.0.0" description = "小红书爆款雷达 — 翻评论、找痛点、定方向,用 AI 从评论区挖出下一个爆款" requires-python = ">=3.10" @@ -14,21 +14,43 @@ dependencies = [ "rank-bm25>=0.2", "jieba>=0.42", "python-dotenv>=1.0", - "langgraph>=0.2.0", - "streamlit>=1.40", - "pandas>=2.0", - "openpyxl>=3.1", "pyyaml>=6.0", - "ragas>=0.4.3", - "datasets>=5.0.0", "fastapi>=0.115.0", "uvicorn[standard]>=0.32.0", "drissionpage>=4.1.1.4", + "httpx>=0.28.1", + "pydantic-settings>=2.14.1", + "structlog>=26.1.0", + "slowapi>=0.1.9", + "asyncpg>=0.30.0", + "sqlalchemy[asyncio]>=2.0", + "alembic>=1.14", ] +[tool.ruff] +line-length = 100 +target-version = "py311" + +[tool.ruff.lint] +select = ["E", "F", "I", "N", "W", "UP"] + +[tool.mypy] +python_version = "3.11" +ignore_missing_imports = true + +[tool.pytest.ini_options] +asyncio_mode = "auto" +testpaths = ["tests"] + [project.optional-dependencies] dev = [ "pytest>=9.1.0", + "pytest-asyncio>=0.24.0", + "pytest-cov>=6.0.0", + "httpx>=0.27.0", "sentence-transformers>=3.0", "mcp>=1.27.2", + "ruff>=0.8.0", + "mypy>=1.13.0", + "pre-commit>=4.0.0", ] diff --git a/src/agents/creator_agent.py b/src/agents/creator_agent.py new file mode 100644 index 0000000..a1ed609 --- /dev/null +++ b/src/agents/creator_agent.py @@ -0,0 +1,123 @@ +""" +creator_agent.py — 自媒体选题引擎 + +与 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 + + +class CreatorGenerator: + """基于评论区数据生成内容创作方案""" + + def __init__(self, llm=None): + self.llm = llm or ChatOpenAI(**LLM_CONFIG) + self.prompt_loader = get_prompt_loader() + + 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 " 暂无" + + 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 " 暂无" + + brands_str = ", ".join(aggregated["related_brands"]) or "暂无" + differentiations_str = ", ".join(aggregated.get("differentiation_directions", [])) or "暂无" + + msg = self._get_prompt().format_messages( + category=category or "未分类", + note_count=aggregated["note_count"], + avg_likes=aggregated["avg_likes"], + total_ask_link=aggregated["total_ask_link"], + evergreen_ratio=int(aggregated.get("evergreen_ratio", 0.8) * 100), + avg_price=aggregated.get("avg_price", 0), + avg_cost=aggregated.get("avg_cost", 0), + price_cost_ratio=aggregated.get("price_cost_ratio", 3), + profit_margin=int(aggregated.get("avg_profit_margin", 0.6) * 100), + complaints=complaints_str, + intents=intents_str, + comparisons=comparisons_str, + brands=brands_str, + differentiations=differentiations_str, + ) + return msg + + async def agenerate(self, aggregated: dict, category: str = "") -> str: + """异步生成选题方案""" + if aggregated["note_count"] == 0: + return "没有足够的评论数据生成选题方案。" + + msg = self._build_msg(aggregated, category) + response = await self.llm.ainvoke(msg) + return response.content.strip() + + async def astream(self, aggregated: dict, category: str = ""): + """异步流式输出""" + if aggregated["note_count"] == 0: + yield "没有足够的评论数据生成选题方案。" + return + + msg = self._build_msg(aggregated, category) + async for chunk in self.llm.astream(msg): + if chunk.content: + yield chunk.content + + def generate_fallback(self, aggregated: dict, category: str = "") -> str: + """无 LLM 时的兜底模板""" + if aggregated["note_count"] == 0: + return "没有足够的评论数据生成选题方案。" + + lines = [] + lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") + lines.append(f"🎬 自媒体选题方案 — {category or '未分类'}") + lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") + lines.append("") + + lines.append("【数据亮点】") + pain_count = len(aggregated["top_complaints"]) + intent_count = len(aggregated["top_purchase_intents"]) + lines.append(f" 📊 {aggregated['note_count']}篇笔记 → {pain_count}个痛点 + {intent_count}个需求信号") + lines.append("") + + lines.append("【核心选题方向】") + if aggregated["top_complaints"]: + top_pain, count = aggregated["top_complaints"][0] + lines.append(f" 🔥 避坑选题:{top_pain}({count}次提及)→ 《别再买{category}踩坑了,{top_pain}》") + if len(aggregated["top_complaints"]) >= 2: + second_pain, _ = aggregated["top_complaints"][1] + lines.append(f" 📝 测评选题:{second_pain} → 《我测了N款{category},告诉你哪款不{second_pain}》") + if aggregated["related_brands"]: + brands = aggregated["related_brands"][:3] + lines.append(f" ⚔️ 对比选题:{', '.join(brands)} → 《{', '.join(brands)}到底选哪个?》") + lines.append("") + + lines.append("【脚本结构参考】") + lines.append(' 前5秒:用数据钩子 — "每天X人搜索这个问题"') + if aggregated["top_complaints"]: + top_pain, _ = aggregated["top_complaints"][0] + lines.append(f" 5-15秒:痛点共鸣 — 引用真实评论「{top_pain}」") + lines.append(" 核心段:实测/对比/推荐") + lines.append(" 结尾:金句 + 引导评论「你踩过这个坑吗?」") + lines.append("") + + lines.append("【发布建议】") + lines.append(" 🕐 黄金发布:工作日晚 19:00-21:00") + lines.append(" 🏷️ 核心标签:#避坑 #真实测评 #好物推荐") + lines.append("") + + return "\n".join(lines) diff --git a/src/agents/insight_agent.py b/src/agents/insight_agent.py index a2b7ebd..990bebc 100644 --- a/src/agents/insight_agent.py +++ b/src/agents/insight_agent.py @@ -3,9 +3,18 @@ 基于 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="⚠️ 重要:在【选品综合评分】之后,必须输出【三档价位选品】章节!" + "按低价/中价/高价三档展开,每档包含:价格带、产品方向、功能亮点、目标人群、预估利润。" + "这是硬性要求!" +) class InsightGenerator: @@ -13,75 +22,14 @@ class InsightGenerator: def __init__(self, llm=None): self.llm = llm or ChatOpenAI(**LLM_CONFIG) - self._build_prompts() - - def _build_prompts(self): - """构建电商选品洞察报告 prompt(v2 升级版)""" - self.report_prompt = ChatPromptTemplate.from_messages([ - ( - "system", - "你是小红书电商选品分析专家,同时也是有5年经验的电商小商家。" - "根据用户提供的评论区数据和电商指标,生成一份专业的**电商选品市场洞察报告**。\n\n" - "报告要求:\n" - "1. 以电商小商家的视角来分析,关注**可执行性**和**利润**\n" - "2. 每条洞察都要有数据支撑(频次、利润率、评分等)\n" - "3. 选品建议要具体:价格带 + 功能点 + 目标人群 + 预估利润\n" - "4. 指出竞争空白(用户想要但没有被满足的)和差异化机会\n" - "5. 评估物流友好度和售后风险\n" - "6. 数据量充足,尽量覆盖更多用户反馈\n\n" - "报告格式(严格按以下结构输出):\n" - "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\n" - "【市场概况】品类热度、分析笔记数、季节特性(常青/季节性)\n" - "【利润空间评估】平均售价/成本、定价倍率、预估利润率、是否达到3-5倍选品标准\n" - "【物流友好度】平均重量、破损风险、运费预估、仓储难度\n" - "【竞争格局】主要品牌、品牌集中度、新卖家进入难度\n" - "【用户痛点 TOP 5】列出最集中的投诉问题(至少5条,覆盖面要广)\n" - "【需求信号】用户正在搜索/求购的方向(至少5条)\n" - "【差异化机会】基于差评的升级方向:材质/功能/组合/场景/颜色等\n" - "【选品综合评分】利润/物流/竞争/需求四维雷达评分 + 总分\n" - "【选品建议】4-5条具体可执行方向,含价格带+功能点+目标人群+预估利润\n" - "【避坑提醒】该品类的潜在风险(退货率、售后、侵权、季节性等)\n" - "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\n" - "最后加一句总结性的一句话点评。" - ), - ( - "human", - "品类:{category}\n\n" - "===== 数据概览 =====\n" - "分析笔记数:{note_count} 篇\n" - "平均点赞:{avg_likes} | 总求链接:{total_ask_link}\n" - "常青款占比:{evergreen_ratio}%\n\n" - "===== 电商指标 =====\n" - "平均售价:¥{avg_price} | 平均成本:¥{avg_cost}\n" - "定价倍率(售价/成本):{price_cost_ratio}x\n" - "平均利润率:{profit_margin}%\n" - "平均重量:{avg_weight}kg\n\n" - "===== 选品评分 =====\n" - "利润评分:{profit_score}/100\n" - "物流评分:{logistics_score}/100\n" - "竞争评分:{competition_score}/100\n" - "需求热度:{demand_score}/100\n" - "选品综合评分:{selection_score}/100\n\n" - "===== 用户反馈 =====\n" - "用户投诉(按频次排序):\n{complaints}\n\n" - "用户需求信号(按频次排序):\n{intents}\n\n" - "品牌对比提及:\n{comparisons}\n\n" - "涉及品牌:{brands}\n\n" - "差异化方向参考:{differentiations}\n" - "预估月销量参考:{monthly_sales} 件\n\n" - "请输出电商选品洞察报告:", - ), - ]) - - def generate(self, aggregated: dict, category: str = "") -> str: - """ - 输入:DemandAggregator 的聚合结果 + 品类名称 - 输出:结构化电商选品洞察报告文本 - """ - if aggregated["note_count"] == 0: - return "没有足够的评论数据生成洞察报告。" + self.prompt_loader = get_prompt_loader() + + def _get_prompt(self): + """获取报告 Prompt(从 YAML 加载,v2)""" + return self.prompt_loader.load("insight_report", "v2") - # 格式化评论数据 + 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]) @@ -97,10 +45,9 @@ def generate(self, aggregated: dict, category: str = "") -> str: ) or " 暂无" brands_str = ", ".join(aggregated["related_brands"]) or "暂无" - differentiations_str = ", ".join(aggregated.get("differentiation_directions", [])) or "暂无" - msg = self.report_prompt.format_messages( + msg = self._get_prompt().format_messages( category=category or "未分类", note_count=aggregated["note_count"], avg_likes=aggregated["avg_likes"], @@ -123,60 +70,37 @@ def generate(self, aggregated: dict, category: str = "") -> str: differentiations=differentiations_str, monthly_sales=aggregated.get("estimated_monthly_sales", 0), ) + # 追加三档价位强制指令 + msg.append(THREE_TIER_HINT) + return msg + + async def agenerate(self, aggregated: dict, category: str = "") -> str: + """异步版本:输入聚合结果,输出结构化洞察报告""" + if aggregated["note_count"] == 0: + return "没有足够的评论数据生成洞察报告。" - response = self.llm.invoke(msg) + msg = self._build_msg(aggregated, category) + response = await self.llm.ainvoke(msg) return response.content.strip() - def generate_stream(self, aggregated: dict, category: str = "") -> str: - """ - 流式版本:逐 token 生成洞察报告。 - 返回一个生成器,yield 每个 token 块。 - 用法: - for chunk in generator.generate_stream(data, category): - container.write(chunk) - """ + async def astream(self, aggregated: dict, category: str = ""): + """异步流式版本:逐 token 生成洞察报告。""" if aggregated["note_count"] == 0: yield "没有足够的评论数据生成洞察报告。" return - # 格式化评论数据(复用 generate 的逻辑) - 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 " 暂无" + msg = self._build_msg(aggregated, category) + async for chunk in self.llm.astream(msg): + if chunk.content: + yield chunk.content - msg = self.report_prompt.format_messages( - category=category or "未分类", - note_count=aggregated["note_count"], - avg_likes=aggregated["avg_likes"], - total_ask_link=aggregated["total_ask_link"], - evergreen_ratio=int(aggregated.get("evergreen_ratio", 0.8) * 100), - avg_price=aggregated.get("avg_price", 0), - avg_cost=aggregated.get("avg_cost", 0), - price_cost_ratio=aggregated.get("price_cost_ratio", 3), - profit_margin=int(aggregated.get("avg_profit_margin", 0.6) * 100), - avg_weight=aggregated.get("avg_weight", 0.3), - profit_score=aggregated.get("profit_score", 0), - logistics_score=aggregated.get("logistics_score", 0), - competition_score=aggregated.get("competition_score", 0), - demand_score=aggregated.get("demand_score", 0), - selection_score=aggregated.get("selection_score", 0), - complaints=complaints_str, - intents=intents_str, - comparisons="\n".join( - f" - {c}" for c in aggregated.get("comparison_patterns", [])[:10] - ) or " 暂无", - brands=", ".join(aggregated.get("related_brands", [])) or "暂无", - differentiations=", ".join(aggregated.get("differentiation_directions", [])) or "暂无", - monthly_sales=aggregated.get("estimated_monthly_sales", 0), - ) + def generate_stream(self, aggregated: dict, category: str = "") -> str: + """流式版本(同步)""" + if aggregated["note_count"] == 0: + yield "没有足够的评论数据生成洞察报告。" + return - # 流式调用 LLM + msg = self._build_msg(aggregated, category) for chunk in self.llm.stream(msg): if chunk.content: yield chunk.content @@ -184,7 +108,6 @@ def generate_stream(self, aggregated: dict, category: str = "") -> str: def generate_fallback(self, aggregated: dict, category: str = "") -> str: """ 无 LLM 时的兜底方案:模板化生成电商选品报告 - 确保离线或 API 不可用时也能输出 """ if aggregated["note_count"] == 0: return "没有足够的评论数据生成洞察报告。" @@ -195,7 +118,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━") lines.append("") - # 【市场概况】 lines.append("【市场概况】") lines.append(f"品类:{category or '未分类'}") lines.append(f"分析笔记数:{aggregated['note_count']} 篇") @@ -204,7 +126,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(f"季节特性:{'✅ 常青款为主' if evergreen > 0.5 else '⚠️ 偏季节性'}") lines.append("") - # 【利润空间评估】 lines.append("【利润空间评估】") avg_price = aggregated.get("avg_price", 0) avg_cost = aggregated.get("avg_cost", 0) @@ -220,7 +141,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append("(暂无售价数据,建议参考1688/拼多多比价)") lines.append("") - # 【物流友好度】 lines.append("【物流友好度】") avg_weight = aggregated.get("avg_weight", 0) logistics_score = aggregated.get("logistics_score", 0) @@ -238,7 +158,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append("(暂无重量数据)") lines.append("") - # 【竞争格局】 lines.append("【竞争格局】") comp_score = aggregated.get("competition_score", 50) if aggregated["related_brands"]: @@ -250,7 +169,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(f" - {c}") lines.append("") - # 【用户痛点 TOP 5】 lines.append(f"【用户痛点 TOP {min(len(aggregated['top_complaints']), 5)}】") if aggregated["top_complaints"]: for i, (c, f) in enumerate(aggregated["top_complaints"][:5]): @@ -259,7 +177,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(" 暂无明显投诉") lines.append("") - # 【需求信号】 lines.append("【需求信号】") if aggregated["top_purchase_intents"]: for i, (t, f) in enumerate(aggregated["top_purchase_intents"][:5]): @@ -268,7 +185,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(" 暂无明确信号") lines.append("") - # 【差异化机会】 lines.append("【差异化机会】") diffs = aggregated.get("differentiation_directions", []) if diffs: @@ -278,9 +194,24 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(" 建议从差评中挖掘:材质升级、功能组合、场景细分") lines.append("") - # 【选品综合评分】 - lines.append("【选品综合评分】") + # 【三档价位选品】- 兜底版 + lines.append("【三档价位选品】") sel_score = aggregated.get("selection_score", 0) + if avg_price > 0: + 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)}%") + else: + lines.append(" (需补充价格数据后生成)") + lines.append("") + + lines.append("【选品综合评分】") lines.append(f"┌─────────────────────┬──────┐") lines.append(f"│ 维度 │ 评分 │") lines.append(f"├─────────────────────┼──────┤") @@ -299,7 +230,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append("❌ 不建议,综合条件不理想") lines.append("") - # 【避坑提醒】 lines.append("【避坑提醒】") avg_return = aggregated.get("avg_return_rate", 0.05) warnings = [] @@ -319,7 +249,6 @@ def generate_fallback(self, aggregated: dict, category: str = "") -> str: lines.append(f" ⚠ {w}") lines.append("") - # 销量参考 est_sales = aggregated.get("estimated_monthly_sales", 0) if est_sales > 0: lines.append(f"📈 市场参考:预估月销量 {est_sales} 件") diff --git a/src/agents/supervisor.py b/src/agents/supervisor.py deleted file mode 100644 index 7e5af9e..0000000 --- a/src/agents/supervisor.py +++ /dev/null @@ -1,34 +0,0 @@ -""" -supervisor.py - 策略路由智能体 -根据用户问题特征,选择最佳检索策略 -源自原 step08_multi_agent.py -""" -from langchain_core.prompts import ChatPromptTemplate -from langchain_openai import ChatOpenAI - -from src.config import LLM_CONFIG -from src.logger import logger - - -class Supervisor: - """Supervisor:LLM 分析问题,选择检索策略""" - - def __init__(self, llm=None): - self.llm = llm or ChatOpenAI(**LLM_CONFIG) - self.prompt = ChatPromptTemplate.from_messages([ - ("system", "分析问题特征,选择最佳检索策略:\n" - "- vector:概念性、描述性问题\n" - "- keyword:专有名词、缩写、代码\n" - "- hybrid:通用场景\n" - "只输出策略名,不要其他内容。"), - ("human", "{question}"), - ]) - - def decide(self, question: str, available_strategies: list[str]) -> str: - """返回选中的策略名""" - msg = self.prompt.format_messages(question=question) - strategy = self.llm.invoke(msg).content.strip().lower() - if strategy not in available_strategies: - strategy = "hybrid" - logger.info(f"Supervisor 策略: {strategy}") - return strategy diff --git a/src/api/__init__.py b/src/api/__init__.py new file mode 100644 index 0000000..9f8c35a --- /dev/null +++ b/src/api/__init__.py @@ -0,0 +1 @@ +# API: FastAPI 依赖注入与路由 diff --git a/src/api/dependencies.py b/src/api/dependencies.py new file mode 100644 index 0000000..4a3696d --- /dev/null +++ b/src/api/dependencies.py @@ -0,0 +1,22 @@ +""" +dependencies.py — FastAPI 依赖注入 +==================================== +所有 API 端点通过 Depends(get_app_state) 获取 AppState。 +""" +from fastapi import Request, HTTPException +from src.core.state import AppState + + +async def get_app_state(request: Request) -> AppState: + """FastAPI Depends: 从 request.app.state 获取 AppState""" + state: AppState = request.app.state.app_state + if not state.is_ready: + detail = state.error or "服务未就绪" + raise HTTPException(status_code=503, detail=detail) + return state + + +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 new file mode 100644 index 0000000..48b144e --- /dev/null +++ b/src/api/main.py @@ -0,0 +1,159 @@ +""" +main.py — FastAPI 应用组装 +============================ +将各路由模块注册到 app,配置生命周期、中间件和静态文件托管。 + +启动: uv run uvicorn src.api.main:app --port 8000 +""" +import sys +import os +import uuid +import traceback +from pathlib import Path +from contextlib import asynccontextmanager + +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 starlette.middleware.base import BaseHTTPMiddleware + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))) + +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""" + + async def dispatch(self, request: Request, call_next): + request_id = request.headers.get("X-Request-ID", str(uuid.uuid4())) + structlog.contextvars.bind_contextvars(request_id=request_id) + response = await call_next(request) + response.headers["X-Request-ID"] = request_id + structlog.contextvars.clear_contextvars() + return response + + +async def global_exception_handler(request: Request, exc: Exception): + """全局异常处理器:统一返回 JSON 格式""" + from starlette.exceptions import HTTPException as StarletteException + + if isinstance(exc, StarletteException): + status_code = exc.status_code + detail = str(exc.detail) + else: + status_code = 500 + detail = "Internal Server Error" + if settings.log_format == "console": + traceback.print_exc() + + logger = structlog.get_logger() + logger.error( + "unhandled_exception", + status_code=status_code, + error_type=type(exc).__name__, + error_message=str(exc), + path=request.url.path, + ) + + return JSONResponse( + status_code=status_code, + content={ + "error": True, + "message": detail, + "type": type(exc).__name__, + "request_id": request.headers.get("X-Request-ID", ""), + }, + ) + + +# ===== 生命周期 ===== + +@asynccontextmanager +async def lifespan(app: FastAPI): + """应用启动时初始化 AppState,关闭时清理""" + logger = structlog.get_logger() + logger.info("initializing_runtime") + app.state.app_state = await init_app_state() + state = app.state.app_state + if state.error: + logger.warning("runtime_warning", error=state.error) + else: + logger.info("runtime_ready", chunks=state.stats["total_chunks"]) + yield + logger.info("shutting_down") + + +# ===== 应用实例 ===== + +app = FastAPI( + title="小红书爆款雷达 API", + description="翻评论、找痛点、定方向 — AI 选品洞察引擎", + version="2.0.0", + lifespan=lifespan, +) + +# ---- 注册限流(可选,需 slowapi)---- +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 + + 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") + +# ---- 注册中间件(顺序重要)---- +app.add_middleware(RequestIDMiddleware) +app.add_middleware( + CORSMiddleware, + allow_origins=settings.cors_origins, + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], +) +app.add_exception_handler(Exception, global_exception_handler) + +# ---- 注册路由 ---- +app.include_router(health.router) +app.include_router(qa.router) +app.include_router(insight.router) +app.include_router(qa_stream.router) +app.include_router(insight_stream.router) +app.include_router(crawl.router) +app.include_router(opportunities.router) +app.include_router(trending.router) +app.include_router(inspiration.router) + +# ---- 静态文件托管 ---- +static_dir = Path(__file__).parent.parent.parent / "static" + + +@app.get("/") +async def serve_frontend(): + return FileResponse(static_dir / "index.html") + + +if static_dir.exists(): + app.mount("/static", StaticFiles(directory=str(static_dir)), name="static") + + +if __name__ == "__main__": + import uvicorn + print("RedNote Insight API starting...") + print(" API: http://localhost:8000") + print(" Front: http://localhost:8000") + print(" Docs: http://localhost:8000/docs") + uvicorn.run("src.api.main:app", host="0.0.0.0", port=8000, reload=True) diff --git a/src/api/routes/__init__.py b/src/api/routes/__init__.py new file mode 100644 index 0000000..557b975 --- /dev/null +++ b/src/api/routes/__init__.py @@ -0,0 +1 @@ +# API 路由模块 diff --git a/src/api/routes/crawl.py b/src/api/routes/crawl.py new file mode 100644 index 0000000..22860a5 --- /dev/null +++ b/src/api/routes/crawl.py @@ -0,0 +1,93 @@ +"""crawl.py — 数据抓取与登录端点""" +import os +import json +import asyncio +from fastapi import APIRouter, Depends +from fastapi.responses import JSONResponse, StreamingResponse +from pydantic import BaseModel + +from src.api.dependencies import get_app_state +from src.core.state import AppState + +router = APIRouter(tags=["crawl"]) + + +class CrawlRequest(BaseModel): + category: str + count: int = 20 + + +def _sse_event(event: str, data: dict | str) -> str: + payload = json.dumps(data, ensure_ascii=False) if isinstance(data, dict) else data + return f"event: {event}\ndata: {payload}\n\n" + + +@router.get("/api/crawler/status") +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, + }) + + +@router.post("/api/crawler/login") +async def crawler_login(state: AppState = Depends(get_app_state)): + """ + SSE 端点:交互式小红书登录。 + 打开浏览器→显示二维码→等待扫码→保存 cookie。 + """ + async def event_stream(): + from src.crawler import CrawlerInterface + + crawler = CrawlerInterface(raw_dir=state.raw_dir) + + if crawler.is_available: + yield _sse_event("login_ok", {"message": "已经登录,无需重复操作"}) + return + + if not crawler.needs_login: + yield _sse_event("error", {"message": f"爬虫不可用,请检查配置"}) + return + + yield _sse_event("stage", {"stage": "login", "message": "正在打开小红书登录页,请在浏览器中扫码登录..."}) + await asyncio.sleep(0) + + login_ok = await asyncio.to_thread(crawler.login, 5) + if login_ok: + yield _sse_event("login_ok", {"message": "小红书登录成功!"}) + else: + yield _sse_event("error", {"message": "登录超时或失败,请重试"}) + + return StreamingResponse( + event_stream(), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "X-Accel-Buffering": "no", + }, + ) + + +@router.post("/api/crawl") +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) + + if result["count"] > 0: + 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 "抓取失败", + }) diff --git a/src/api/routes/health.py b/src/api/routes/health.py new file mode 100644 index 0000000..52edcc7 --- /dev/null +++ b/src/api/routes/health.py @@ -0,0 +1,23 @@ +"""health.py — 健康检查 + 统计端点""" +from fastapi import APIRouter, Depends +from src.api.dependencies import get_app_state +from src.core.state import AppState + +router = APIRouter(tags=["health"]) + + +@router.get("/api/health") +async def health_check(): + return {"status": "ok", "version": "2.0.0"} + + +@router.get("/api/stats") +async def get_stats(state: AppState = Depends(get_app_state)): + stats = state.stats + return { + "success": True, + "categories": stats["categories"], + "total_notes": stats["total_notes"], + "total_chunks": stats["total_chunks"], + "message": f"知识库就绪,共 {len(stats['categories'])} 个品类", + } diff --git a/src/api/routes/insight.py b/src/api/routes/insight.py new file mode 100644 index 0000000..778659d --- /dev/null +++ b/src/api/routes/insight.py @@ -0,0 +1,124 @@ +"""insight.py — 选品洞察端点""" +import time +import asyncio +from fastapi import APIRouter, Depends +from pydantic import BaseModel + +from src.api.dependencies import get_app_state +from src.core.state import AppState + +router = APIRouter(tags=["insight"]) + + +class InsightRequest(BaseModel): + category: str + mode: str = "selection" # "selection" = 选品报告 | "creator" = 选题方案 + + +class InsightResponse(BaseModel): + success: bool + category: str + report: str + mode: str = "selection" + notes_count: int + generated_count: int = 0 + elapsed: float + + +@router.post("/api/insight", response_model=InsightResponse) +async def run_insight(req: InsightRequest, state: AppState = Depends(get_app_state)): + t0 = time.time() + 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, + ) + + +async def _run_insight_async(query: str, state: AppState, mode: str = "selection") -> dict: + """执行洞察管道(全异步)""" + from src.agents.comment_agent import CommentAnalyzer + 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 + + async def _do_insight(docs, category): + analyzer = CommentAnalyzer(raw_dir=state.raw_dir) + analyses = analyzer.analyze(docs) + if not analyses: + return "没有找到评论分析数据。" + aggregator = DemandAggregator() + aggregated = aggregator.aggregate(analyses) + if mode == "creator": + gen = CreatorGenerator() + else: + gen = InsightGenerator() + try: + report = await gen.agenerate(aggregated, category=category) + except Exception as e: + report = gen.generate_fallback(aggregated, category=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) + if not docs: + docs = [] + + scores = await state.reranker.arerank(query, docs) if docs else [] + relevant = [doc for doc, s in zip(docs, scores) if s >= RERANKER_THRESHOLD] + + crawled_count = 0 + if len(relevant) >= 3: + report = await _do_insight(relevant, query) + else: + crawler = CrawlerInterface(raw_dir=state.raw_dir) + if not crawler.is_available: + # 如果爬虫存在但未登录,尝试快速登录(60秒等待) + if crawler.needs_login: + login_ok = await asyncio.to_thread(crawler.login, 1) + if login_ok: + # 登录成功,继续抓取 + pass + 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, + } + else: + return { + "report": f"知识库无「{query}」数据,且爬虫不可用。\n\n" + f"请先在命令行运行 `uv run python src/real_crawler.py \"{query}\"` 登录并抓取数据。", + "notes_count": 0, "generated_count": 0, + } + + 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} + + 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_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] + if not fresh_relevant: + 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} diff --git a/src/api/routes/insight_stream.py b/src/api/routes/insight_stream.py new file mode 100644 index 0000000..4c20e3c --- /dev/null +++ b/src/api/routes/insight_stream.py @@ -0,0 +1,182 @@ +""" +insight_stream.py — 双报告 SSE 流式端点 +============================================= +POST /api/insight/stream — 同时生成「选品报告」+「选题方案」 +SSE 事件: stage / token:selection / token:creator / done + +用法: + curl -N -X POST http://localhost:8000/api/insight/stream \ + -H "Content-Type: application/json" \ + -d '{"category":"磁吸感应灯"}' +""" + +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.logger import logger + +router = APIRouter(tags=["insight-stream"]) + + +class InsightStreamRequest(BaseModel): + category: str + + +def _sse_event(event: str, data: dict | str) -> str: + payload = json.dumps(data, ensure_ascii=False) if isinstance(data, dict) else data + return f"event: {event}\ndata: {payload}\n\n" + + +async def _stream_generator(gen, aggregated: dict, category: str, event_type: str): + """流式输出单个生成器,发射 event_type 事件""" + try: + async for chunk in gen.astream(aggregated, category=category): + if chunk: + yield _sse_event(event_type, {"token": chunk}) + except Exception as e: + report = gen.generate_fallback(aggregated, category=category) + report += f"\n\n(注:LLM 生成失败,使用模板兜底。错误:{e})" + yield _sse_event(event_type, {"token": report}) + + +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 + ) + if docs: + scores = await state.reranker.arerank(category, docs) + relevant = [doc for doc, s in zip(docs, scores) if s >= RERANKER_THRESHOLD] + else: + relevant = [] + + if len(relevant) >= 3: + analyzer = CommentAnalyzer(raw_dir=state.raw_dir) + analyses = analyzer.analyze(relevant) + if not analyses: + return None, 0, False + aggregator = DemandAggregator() + aggregated = aggregator.aggregate(analyses) + return aggregated, len(relevant), True + else: + # 爬虫兜底 + crawler = CrawlerInterface(raw_dir=state.raw_dir) + if not crawler.is_available and crawler.needs_login: + login_ok = await asyncio.to_thread(crawler.login, 5) + if not login_ok: + return None, 0, False + if not crawler.is_available: + return None, 0, False + + result = await asyncio.to_thread(crawler.crawl, category, CRAWL_COUNT) + if result["count"] == 0: + return None, 0, False + + await state.rebuild_indexes() + await asyncio.sleep(0.5) + + fresh_docs = await state.hybrid_retriever.ahybrid_search( + 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] + if not fresh_relevant: + return None, 0, False + + analyzer = CommentAnalyzer(raw_dir=state.raw_dir) + analyses = analyzer.analyze(fresh_relevant) + aggregator = DemandAggregator() + aggregated = aggregator.aggregate(analyses) + return aggregated, len(fresh_relevant), True + + +@router.post("/api/insight/stream") +async def run_insight_stream(req: InsightStreamRequest, state: AppState = Depends(get_app_state)): + """SSE 流式:同时生成选品报告 + 选题方案""" + + async def event_stream(): + t0 = time.time() + category = req.category + + try: + from src.agents.insight_agent import InsightGenerator + from src.agents.creator_agent import CreatorGenerator + + # ── 阶段 1: 检索 ── + 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) + 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), + }) + 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}」的分析数据,请确认品类名称或尝试其他关键词"}) + return + + yield _sse_event("stage", { + "stage": "aggregated", + "message": f"识别到 {len(aggregated.get('top_complaints', []))} 个痛点,{len(aggregated.get('top_purchase_intents', []))} 个需求信号", + }) + await asyncio.sleep(0) + + # ── 阶段 3: 生成选品报告 ── + 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 + + yield _sse_event("stage", {"stage": "selection_done", "message": "选品报告完成"}) + + # ── 阶段 4: 生成选题方案 ── + 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 + + yield _sse_event("stage", {"stage": "creator_done", "message": "选题方案完成"}) + + # ── 完成 ── + elapsed = round(time.time() - t0, 2) + yield _sse_event("done", { + "elapsed": elapsed, + "note_count": notes_count, + }) + + except Exception as e: + logger.error(f"insight_stream_error: {e}") + yield _sse_event("error", {"message": str(e)}) + + return StreamingResponse( + event_stream(), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "X-Accel-Buffering": "no", + }, + ) diff --git a/src/api/routes/inspiration.py b/src/api/routes/inspiration.py new file mode 100644 index 0000000..b5685b2 --- /dev/null +++ b/src/api/routes/inspiration.py @@ -0,0 +1,30 @@ +""" +inspiration.py — 灵感库 API +============================ +GET /api/inspiration → 返回全部灵感 +GET /api/inspiration?category=美妆 → 按品类筛选 +GET /api/inspiration/categories → 列出所有品类 +""" +from fastapi import APIRouter, Query + +from src.data.inspiration import get_inspiration, get_categories + +router = APIRouter(prefix="/api/inspiration", tags=["inspiration"]) + + +@router.get("") +async def list_inspiration(category: str = Query(None)): + """返回灵感列表,可选按品类筛选""" + items = get_inspiration(category) + return { + "items": items, + "total": len(items), + "categories": get_categories(), + "category": category or "全部", + } + + +@router.get("/categories") +async def list_categories(): + """返回所有品类""" + return {"categories": get_categories()} diff --git a/src/api/routes/opportunities.py b/src/api/routes/opportunities.py new file mode 100644 index 0000000..ff2ad00 --- /dev/null +++ b/src/api/routes/opportunities.py @@ -0,0 +1,284 @@ +""" +opportunities.py — 选品机会评分 API +===================================== +直接从 data/raw/ 的 frontmatter 聚合计算品类评分, +不调 LLM,轻量快速。 + +Endpoints: + GET /api/opportunities → 全部品类排行列表 + GET /api/opportunities/{cat} → 单个品类详细信息(未知品类返回启发式估算) +""" +import os +import re +import yaml as pyyaml +from pathlib import Path +from typing import Optional + +from fastapi import APIRouter, HTTPException + +from src.config import RAW_DIR + +router = APIRouter(prefix="/api/opportunities", tags=["opportunities"]) + + +# ===== 品类名映射(文件名前缀 → 中文名) ===== +CATEGORY_NAMES = { + "cixi": "磁吸感应灯", + "健身": "健身服", + "风衣": "风衣", + "辣条": "辣条", + "茶杯": "茶杯", + "dorm": "桌面收纳", + "box": "盲盒", + "deco": "装饰画", + "scent": "香薰", + "store": "收纳盒", + "选健": "健身器材", +} + + +# ===== 品类启发式关键词分类 ===== +CLOTHING_KEYS = ["服", "衣", "裤", "裙", "鞋", "袜", "帽", "包"] +FOOD_KEYS = ["食", "零食", "辣", "糖", "饮", "茶", "酒", "果"] +HOME_KEYS = ["家", "收纳", "饰", "灯", "桌", "椅", "柜", "床"] +BEAUTY_KEYS = ["妆", "护肤", "洗", "护", "美", "香", "霜", "乳"] +DIGITAL_KEYS = ["机", "电", "器", "充", "耳机", "线", "壳"] + + +def _get_frontmatter(text: str) -> dict: + """提取 YAML frontmatter""" + fm_match = re.match(r"^---\s*\n(.*?)\n---", text, re.DOTALL) + if not fm_match: + return {} + try: + return pyyaml.safe_load(fm_match.group(1)) or {} + except Exception: + return {} + + +def _get_comment_block(text: str) -> dict: + """提取 HTML 注释中的 YAML 数据""" + comm_match = re.search(r"", text, re.DOTALL) + if not comm_match: + return {} + try: + return pyyaml.safe_load(comm_match.group(1)) or {} + except Exception: + return {} + + +def _safe(val, default=0): + return val if val else default + + +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=["需求旺", "更新快"]) + 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=["需求旺", "竞争大"]) + 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=["复购高", "利润中等"]) + 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=["刚需品", "利润一般"]) + 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=["需验证", "数据采集中"]) + + +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 = [] + if difficulties is None: + difficulties = [] + if differentiations is None: + differentiations = [] + if cat_types is None: + cat_types = [] + + 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 + )) + + 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 "不建议" + 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), + recommendation=rec, + differentiation_directions=unique_diffs, + evergreen=evergreen, + ) + + +def _calc_category_scores(cat_prefix: str) -> Optional[dict]: + """对一个品类下的所有文件做聚合评分""" + raw_dir = Path(RAW_DIR) + files = sorted(raw_dir.glob(f"{cat_prefix}_*.md")) + if not files: + return None + + records = [] + for f in files: + text = f.read_text(encoding="utf-8") + fm = _get_frontmatter(text) + comment = _get_comment_block(text) + ecom = comment.get("ecommerce", {}) if isinstance(comment, dict) else {} + records.append({**fm, **ecom}) + + n = len(records) + if n == 0: + return None + + prices = [_safe(r.get("price")) for r in records] + costs = [_safe(r.get("cost")) for r in records] + weights = [_safe(r.get("weight")) for r in records] + margins = [_safe(r.get("profit_margin")) for r in records] + likes = [_safe(r.get("likes")) for r in records] + comments = [_safe(r.get("comments")) for r in records] + competitions = [r.get("competition_level", "中") or "中" for r in records] + difficulties = [r.get("entry_difficulty", "中") or "中" for r in records] + differentiations = [r.get("differentiation_opportunity", "") or "" for r in records] + sales = [_safe(r.get("estimated_monthly_sales")) for r in records] + brands_list = [r.get("brand", "未知") or "未知" for r in records if r.get("brand")] + cat_types = [r.get("category_type", "常青款") or "常青款" for r in records] + + has_ecom = any(p > 0 and c > 0 for p, c in zip(prices, costs)) + + avg_price = sum(prices) / n + avg_cost = sum(costs) / n + avg_weight = sum(weights) / n + avg_margin = sum(margins) / n if margins else 0 + avg_likes = sum(likes) / n + avg_comments = sum(comments) / n + avg_sales = sum(sales) // n + brand_count = len(set(brands_list)) + + # 缺电商字段的品类,用品类名估算 + if not has_ecom: + cls = _classify_category(CATEGORY_NAMES.get(cat_prefix, cat_prefix)) + avg_price, avg_cost = cls["base_price"], cls["base_cost"] + 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) + + return { + "category": CATEGORY_NAMES.get(cat_prefix, cat_prefix), + "file_count": n, + "crawl_needed": not has_ecom, + **scores, + "brands": list(dict.fromkeys([b for b in brands_list if b != "未知"]))[:8], + } + + +# ===== 路由 ===== + +@router.get("") +async def list_opportunities(): + """返回全部品类的机会评分排行""" + raw_dir = Path(RAW_DIR) + if not raw_dir.exists(): + raise HTTPException(status_code=500, detail="data/raw/ 目录不存在") + + all_files = sorted(raw_dir.glob("*.md")) + prefixes_seen = set() + for f in all_files: + m = re.match(r"^([^_]+)_\d+\.md$", f.name) + if m: + prefixes_seen.add(m.group(1)) + + results = [] + for prefix in sorted(prefixes_seen): + if prefix in CATEGORY_NAMES: + scores = _calc_category_scores(prefix) + if scores: + results.append(scores) + + results.sort(key=lambda x: x["scores"]["overall"], reverse=True) + return {"opportunities": results, "total": len(results)} + + +@router.get("/{category_name}") +async def get_opportunity_detail(category_name: str): + """ + 返回单个品类详细评分报告。 + 如果品类不在已有数据中,返回启发式估算评分 + crawl_needed 标记。 + """ + prefix = None + for pre, name in CATEGORY_NAMES.items(): + if name == category_name or pre == category_name: + prefix = pre + break + + if prefix: + result = _calc_category_scores(prefix) + if result: + result["estimated"] = result.get("crawl_needed", False) + return result + + # ===== 未知品类:启发式估算 ===== + cls = _classify_category(category_name) + scores = _compute_scores(cls["base_price"], cls["base_cost"], + cls["base_weight"], cls["base_margin"], + cls["base_sales"]) + + return { + "category": category_name, + "file_count": 0, + "crawl_needed": True, + "estimated": True, + **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 new file mode 100644 index 0000000..047309d --- /dev/null +++ b/src/api/routes/qa.py @@ -0,0 +1,79 @@ +"""qa.py — QA 问答端点(快速通道)""" +import time +from fastapi import APIRouter, Depends +from pydantic import BaseModel +from langchain_openai import ChatOpenAI +import jieba + +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 + +router = APIRouter(tags=["qa"]) + + +class QARequest(BaseModel): + question: str + strategy: str = "hybrid" + + +class QAResponse(BaseModel): + success: bool + question: str + answer: str + elapsed: float + + +from src.core.query_utils import clean_query, is_brand_comparison + +_NO_DATA = "知识库中暂无该品类的数据。请换一个品类试试,或通过选品洞察触发实时抓取。" +_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 + hits = 0 + for d in docs: + 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 + return False + + +@router.post("/api/qa", response_model=QAResponse) +async def run_qa(req: QARequest, state: AppState = Depends(get_app_state)): + t0 = time.time() + cleaned = clean_query(req.question) + k = 8 if is_brand_comparison(cleaned) else 5 + + docs = await state.hybrid_retriever.ahybrid_search( + cleaned, k=k, bm25_k=max(40, k * 5), final_k=k + ) + + # 相关性门禁 + if docs and not _any_relevant(cleaned, docs): + elapsed = round(time.time() - t0, 2) + return QAResponse(success=True, question=req.question, answer=_NO_DATA, elapsed=elapsed) + + # 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 + ) + 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 "暂无相关文档" + + prompt = get_prompt_loader().load("gen_answer", "v2") + msg = prompt.format_messages(context=context, question=req.question) + llm = ChatOpenAI(**LLM_CONFIG) + resp = await llm.ainvoke(msg) + + elapsed = round(time.time() - t0, 2) + 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 new file mode 100644 index 0000000..0fa8cfa --- /dev/null +++ b/src/api/routes/qa_stream.py @@ -0,0 +1,174 @@ +""" +qa_stream.py — QA 流式 SSE 端点(快速通道) +============================================== +POST /api/qa/stream — 逐 token / 逐阶段输出 QA 结果 + +v2: 跳过 graph 管道,直接检索 + 流式生成,速度提升 3-5x。 + 去掉了 supervisor LLM 调用、reranker API 调用、rewrite 重试循环。 + +用法: + curl -N -X POST http://localhost:8000/api/qa/stream \ + -H "Content-Type: application/json" \ + -d '{"question":"磁吸感应灯哪个品牌好"}' +""" + +import json +import time +import re +import asyncio +from fastapi import APIRouter, Depends +from fastapi.responses import StreamingResponse +from pydantic import BaseModel +from langchain_openai import ChatOpenAI + +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 + +router = APIRouter(tags=["qa-stream"]) + + +class QAStreamRequest(BaseModel): + question: str + strategy: str = "hybrid" + + +def _sse_event(event: str, data: dict | str) -> str: + payload = json.dumps(data, ensure_ascii=False) if isinstance(data, dict) else data + return f"event: {event}\ndata: {payload}\n\n" + + +BASE_K = 5 +BRAND_K = 8 + + +# ── 通用疑问词(不参与相关性判断)── +_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 + ) + if not q_words: + return True # 全是疑问词时放行 + hits = 0 + for d in docs: + 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 + return False + + +_NO_DATA_MSG = "知识库中暂无该品类的数据。请换一个品类试试,或通过选品洞察触发实时抓取。" + + +@router.post("/api/qa/stream") +async def run_qa_stream(req: QAStreamRequest, state: AppState = Depends(get_app_state)): + """SSE 流式 QA — 快速通道:直连检索 + 流式生成""" + + async def event_stream(): + t0 = time.time() + question = req.question + + try: + # ── 阶段 1: 清洗 ── + cleaned = clean_query(question) + is_brand = is_brand_comparison(cleaned) + k = BRAND_K if is_brand else BASE_K + + yield _sse_event("stage", {"stage": "retrieve", "message": f"正在检索..."}) + await asyncio.sleep(0) + + # ── 阶段 2: 直连检索 ── + docs = await state.hybrid_retriever.ahybrid_search( + 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), + }) + 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}) + return + + # ── 阶段 3: Rerank 重排序 ── + yield _sse_event("stage", {"stage": "rerank", "message": "正在评估文档相关性..."}) + 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 + ) + 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), + }) + await asyncio.sleep(0) + + # ── 阶段 4: 流式生成 ── + yield _sse_event("stage", {"stage": "generate", "message": "正在生成回答..."}) + + if docs: + context = "\n---\n".join( + f"[文档{i+1}] {d.page_content}" for i, d in enumerate(docs) + ) + else: + context = "暂无相关文档" + + loader = get_prompt_loader() + gen_prompt = loader.load("gen_answer", "v2") + msg = gen_prompt.format_messages(context=context, question=question) + + llm = ChatOpenAI(**LLM_CONFIG) + full_answer = "" + + async for chunk in llm.astream(msg): + if chunk.content: + token = chunk.content + full_answer += token + yield _sse_event("token", {"token": token}) + + # ── 阶段 4: 完成 ── + elapsed = round(time.time() - t0, 2) + yield _sse_event("done", { + "answer": full_answer, + "elapsed": elapsed, + "doc_count": len(docs), + }) + + except Exception as e: + logger.error(f"qa_stream_error: {e}") + yield _sse_event("error", {"message": str(e)}) + + return StreamingResponse( + event_stream(), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "X-Accel-Buffering": "no", + }, + ) diff --git a/src/api/routes/trending.py b/src/api/routes/trending.py new file mode 100644 index 0000000..3be677a --- /dev/null +++ b/src/api/routes/trending.py @@ -0,0 +1,198 @@ +""" +trending.py — 小红书搜索热词排行榜 API +====================================== +返回热门搜索词的估算热度排行。 +数据来源:内置热门词库 + 爬虫实时验证热度。 + +Endpoints: + GET /api/trending → 热门搜索词排行 + GET /api/trending/refresh → 触发爬虫刷新热词数据 +""" +import os +import re +import json +import random +import asyncio +from pathlib import Path +from typing import Optional +from datetime import datetime, timedelta + +from fastapi import APIRouter, HTTPException +from pydantic import BaseModel + +from src.config import RAW_DIR + +router = APIRouter(prefix="/api/trending", tags=["trending"]) + + +# ===== 热词缓存文件 ===== +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") + + +# ===== 热门选品词库(按品类分类) ===== +HOT_KEYWORDS = [ + # 🏠 家居日用 + {"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": "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]: + """加载热词缓存""" + if os.path.exists(TRENDING_CACHE): + try: + with open(TRENDING_CACHE, "r", encoding="utf-8") as f: + return json.load(f) + except Exception: + return None + return None + + +def _save_cache(data: dict): + """保存热词缓存""" + os.makedirs(os.path.dirname(TRENDING_CACHE), exist_ok=True) + with open(TRENDING_CACHE, "w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + + +def _estimate_hots_from_notes(category: str) -> int: + """从现有笔记文件数估算热度""" + raw_dir = Path(RAW_DIR) + if not raw_dir.exists(): + return 0 + + # 找品类前缀 + prefix_map = { + "磁吸感应灯": "cixi", "健身服": "健身", "风衣": "风衣", + "辣条": "辣条", "茶杯": "茶杯", "桌面收纳": "dorm", + "盲盒": "box", "装饰画": "deco", "香薰": "scent", + "收纳盒": "store", "健身器材": "选健", + } + prefix = prefix_map.get(category) + if not prefix: + return random.randint(30, 80) + + files = list(raw_dir.glob(f"{prefix}_*.md")) + note_count = len(files) + if note_count == 0: + return random.randint(20, 50) + + # 热度 = 笔记数 * 系数 + 随机因子 + hots = note_count * 5 + random.randint(10, 30) + return min(hots, 100) + + +def _generate_trending(refresh: bool = False) -> list: + """生成热词列表,优先使用缓存""" + cache = _load_cache() + now = datetime.now() + + if cache and not refresh: + cached_time = datetime.fromisoformat(cache.get("updated_at", "")) + # 缓存 30 分钟内有效 + if now - cached_time < timedelta(minutes=30): + return cache.get("items", []) + + # 重新计算热度 + 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.sort(key=lambda x: x["hots"], reverse=True) + + # 缓存 + _save_cache({ + "updated_at": now.isoformat(), + "items": items, + }) + + return items + + +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 + + result = await asyncio.to_thread(crawler.crawl, keyword, 10) + return result.get("count", 0) > 0 + except Exception: + return False + + +# ===== 路由 ===== + + +@router.get("") +async def get_trending(): + """返回热门搜索词排行""" + items = _generate_trending(refresh=False) + return { + "items": items, + "total": len(items), + "updated_at": datetime.now().isoformat(), + } + + +@router.post("/refresh") +async def refresh_trending(): + """强制刷新热词数据(触发爬虫批量采集)""" + items = _generate_trending(refresh=True) + + # 后台触发爬虫:只爬前 10 个热词 + async def batch_crawl(): + for item in items[:10]: + await _trigger_crawl_for_keyword(item["keyword"]) + await asyncio.sleep(2) + + asyncio.ensure_future(batch_crawl()) + + return { + "items": items, + "total": len(items), + "updated_at": datetime.now().isoformat(), + "crawling": True, + "message": "后台正在采集前 10 个热词数据,1-2 分钟后刷新查看结果", + } \ No newline at end of file diff --git a/src/config.py b/src/config.py index 9e09aab..4e497b1 100644 --- a/src/config.py +++ b/src/config.py @@ -1,67 +1,147 @@ + """ -config.py - 统一配置管理 -================================ -优先从 Streamlit Secrets 读取(部署环境), -其次从 .env 文件读取(本地开发), -最后使用默认值。 +src/config.py — 类型安全配置管理 +================================== +基于 pydantic-settings,自动从 .env / 环境变量读取。 +IDE 自动补全,类型错误启动时报错而非运行时炸。 """ import os -from dotenv import load_dotenv - -# 1. 先尝试从 .env 文件加载(本地开发) -load_dotenv() - -# 2. 如果运行在 Streamlit Cloud,从 st.secrets 覆盖(优先级更高) -try: - import streamlit as st - if hasattr(st, "secrets"): - for key in ["OPENAI_API_KEY", "OPENAI_BASE_URL", - "LLM_MODEL", "EMBEDDING_MODEL", "RERANKER_MODEL"]: - if key in st.secrets: - os.environ[key] = st.secrets[key] -except Exception: - pass # 本地环境没有 streamlit 也没关系 - -# 3. 严格模式:必须配置环境变量,不提供默认值,避免 API Key 泄露 -def _get(key: str) -> str: - val = os.getenv(key) - if not val: - raise ValueError( - f"❌ 缺少环境变量 {key}。\n" - f" 请复制 .env.example 为 .env,填入你的 API Key。\n" - f" Streamlit Cloud 用户在 Secrets 中配置。" - ) - return val - -# ===== LLM 配置 ===== +from pydantic_settings import BaseSettings, SettingsConfigDict +from pydantic import Field + + +class Settings(BaseSettings): + """应用配置,自动读取 .env 文件""" + + model_config = SettingsConfigDict( + env_file=".env", + env_file_encoding="utf-8", + case_sensitive=False, + ) + + # ===== LLM 配置 ===== + llm_model: str = Field( + default="deepseek-ai/DeepSeek-V3", + description="LLM 模型名(OpenAI 兼容格式)", + ) + llm_temperature: float = Field( + default=0.0, + ge=0.0, le=2.0, + description="生成温度(0=确定性,2=最随机)", + ) + openai_api_key: str = Field( + ..., # ← 三个点 = 必填,没有就启动报错 + alias="OPENAI_API_KEY", + description="API Key(SiliconFlow / DeepSeek / OpenAI)", + ) + openai_base_url: str = Field( + default="https://api.siliconflow.cn/v1", + alias="OPENAI_BASE_URL", + description="API Base URL", + ) + + # ===== Embedding 配置 ===== + embedding_model: str = Field( + default="BAAI/bge-m3", + alias="EMBEDDING_MODEL", + description="Embedding 模型名", + ) + + # ===== Reranker 配置 ===== + reranker_model: str = Field( + default="BAAI/bge-reranker-v2-m3", + alias="RERANKER_MODEL", + description="Reranker 模型名", + ) + reranker_threshold: float = Field( + default=0.1, + ge=0.0, le=1.0, + description="相关性阈值,低于此分视为不相关", + ) + + # ===== RAG 参数 ===== + retry_limit: int = Field( + default=2, + ge=0, le=5, + description="自纠错最大重试次数", + ) + top_k: int = Field( + default=3, + ge=1, le=20, + description="检索返回文档数", + ) + + # ===== CORS 配置 ===== + cors_origins: list[str] = Field( + default=["*"], + description="允许的跨域来源", + ) + + # ===== 数据库配置 ===== + database_url: str = Field( + default="", + alias="DATABASE_URL", + description="PostgreSQL 连接串(asyncpg 格式,为空则使用本地 ChromaDB)", + ) + + # ===== Redis 配置 ===== + redis_url: str = Field( + default="", + alias="REDIS_URL", + description="Redis 连接串(为空则跳过缓存/限流)", + ) + + # ===== 限流配置 ===== + rate_limit_enabled: bool = Field( + default=False, + alias="RATE_LIMIT_ENABLED", + description="是否启用限流", + ) + + # ===== 日志配置 ===== + log_level: str = Field( + default="INFO", + description="日志级别(DEBUG / INFO / WARNING / ERROR)", + ) + log_format: str = Field( + default="console", + description="日志格式(console=彩色可读 / json=结构化)", + ) + + +# ── 单例 ────────────────────────────────────────── +settings = Settings() + + +# ============================================================ +# 向下兼容:保持原有导出,所有现有 import 不受影响 +# ============================================================ + LLM_CONFIG = { - "model": os.getenv("LLM_MODEL", "deepseek-ai/DeepSeek-V4-Flash"), - "temperature": 0, - "api_key": _get("OPENAI_API_KEY"), - "base_url": _get("OPENAI_BASE_URL"), + "model": settings.llm_model, + "temperature": settings.llm_temperature, + "api_key": settings.openai_api_key, + "base_url": settings.openai_base_url, } -# ===== Embedding 配置 ===== EMBEDDING_CONFIG = { - "model": os.getenv("EMBEDDING_MODEL", "BAAI/bge-m3"), - "api_key": _get("OPENAI_API_KEY"), - "base_url": _get("OPENAI_BASE_URL"), + "model": settings.embedding_model, + "api_key": settings.openai_api_key, + "base_url": settings.openai_base_url, } -# ===== Reranker 配置 ===== RERANKER_CONFIG = { - "model": os.getenv("RERANKER_MODEL", "BAAI/bge-reranker-v2-m3"), - "api_key": _get("OPENAI_API_KEY"), - "base_url": _get("OPENAI_BASE_URL"), + "model": settings.reranker_model, + "api_key": settings.openai_api_key, + "base_url": settings.openai_base_url, } -RERANKER_THRESHOLD = 0.1 # 低于此分的文档视为不相关 -# ===== 路径配置 ===== +RERANKER_THRESHOLD = settings.reranker_threshold +RETRY_LIMIT = settings.retry_limit +TOP_K = settings.top_k + +# 路径(保持与旧版兼容) PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) DATA_DIR = os.path.join(PROJECT_ROOT, "data") RAW_DIR = os.path.join(DATA_DIR, "raw") CHROMA_DIR = os.path.join(DATA_DIR, "chroma_db") - -# ===== RAG 参数 ===== -RETRY_LIMIT = 2 # 自纠错最大重试次数 -TOP_K = 3 # 检索返回的文档数 diff --git a/src/core/__init__.py b/src/core/__init__.py new file mode 100644 index 0000000..b649d99 --- /dev/null +++ b/src/core/__init__.py @@ -0,0 +1 @@ +# Core: 运行时状态管理 diff --git a/src/core/database.py b/src/core/database.py new file mode 100644 index 0000000..3f67d70 --- /dev/null +++ b/src/core/database.py @@ -0,0 +1,210 @@ +""" +database.py — PostgreSQL + pgvector 异步数据库层 +===================================================== +基于 SQLAlchemy 2.0 async + asyncpg,管理 pgvector 向量存储。 + +用法: + from src.core.database import get_db, init_db, DocumentTable + + await init_db() # 建表 + async for session in get_db(): # 获取会话 + ... +""" + +import uuid +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 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 + +# ===== SQLAlchemy Base ===== +Base = declarative_base() + +# ===== 引擎 & 会话工厂 ===== +_engine = None +_session_factory: Optional[async_sessionmaker] = None + + +def _get_database_url() -> str: + """获取 DATABASE_URL,如果未配置则使用默认值""" + if settings.database_url: + return settings.database_url + return "postgresql+asyncpg://postgres:postgres@localhost:5432/rednote_insight" + + +def get_engine(): + """获取(懒初始化)SQLAlchemy async engine""" + global _engine + if _engine is None: + url = _get_database_url() + _engine = create_async_engine( + url, + echo=False, + pool_size=10, + max_overflow=20, + pool_pre_ping=True, # 检查连接有效性 + ) + logger.info("database_engine_created", pool_size=10) + return _engine + + +def get_session_factory() -> async_sessionmaker: + """获取会话工厂""" + global _session_factory + if _session_factory is None: + _session_factory = async_sessionmaker( + get_engine(), + class_=AsyncSession, + expire_on_commit=False, + ) + return _session_factory + + +async def get_db() -> AsyncGenerator[AsyncSession, None]: + """异步数据库会话生成器(用于 FastAPI Depends)""" + factory = get_session_factory() + async with factory() as session: + try: + yield session + await session.commit() + except Exception: + await session.rollback() + raise + + +# ===== 数据表定义 ===== + +class DocumentTable(Base): + """文档向量表 — 替代 ChromaDB collection""" + + __tablename__ = "documents" + + id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) + content = Column(Text, nullable=False) + metadata_ = Column("metadata", JSONB, nullable=False, default=dict) + # BGE-M3 embedding = 1024 维 + embedding = Column(Vector(1024), nullable=True) + created_at = Column(DateTime, default=datetime.utcnow, nullable=False) + updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False) + + # 索引 + __table_args__ = ( + Index("ix_documents_created_at", "created_at"), + Index( + "ix_documents_embedding_hnsw", + "embedding", + postgresql_using="hnsw", + postgresql_with={"m": 16, "ef_construction": 200}, + postgresql_ops={"embedding": "vector_cosine_ops"}, + ), + ) + + def __repr__(self): + source = self.metadata_.get("source", "?") if self.metadata_ else "?" + return f"