Ver código fonte

feat(langgraph): config, models, RAG, Java client, RecommendAgent

- config.py: pydantic-settings 配置管理
- models/: 请求/响应 Pydantic 模型
- rag/: ChromaDB + Embedding + 混合检索 (BM25+向量)
- tools/: JavaClient HTTP 封装 + product_tools (3 个 @tool)
- graphs/recommend_graph.py: RecommendAgent (LLM + Tool Calling)
- api/recommend.py: POST /api/v1/recommend 端点
- main.py: 注册路由 + startup/shutdown 钩子
iwt 2 meses atrás
pai
commit
317d27eb38

+ 0 - 0
cfc-langgraph/app/agents/__init__.py


+ 40 - 0
cfc-langgraph/app/api/recommend.py

@@ -0,0 +1,40 @@
+from fastapi import APIRouter
+from app.models.recommend import RecommendRequest, RecommendResponse, RecommendItem
+from app.graphs.recommend_graph import RecommendAgent
+import logging
+
+logger = logging.getLogger(__name__)
+router = APIRouter(prefix="/api/v1", tags=["recommend"])
+
+_agent: RecommendAgent = None
+
+
+def get_agent() -> RecommendAgent:
+    global _agent
+    if _agent is None:
+        _agent = RecommendAgent()
+    return _agent
+
+
+@router.post("/recommend", response_model=RecommendResponse)
+async def recommend(req: RecommendRequest):
+    """营养推荐: Agent 搜索+LLM 解释"""
+    try:
+        agent = get_agent()
+        result = await agent.run(query=req.query, tags=req.tags, limit=req.limit)
+
+        items_data = result.get("items", [])
+        items = []
+        for item in items_data:
+            items.append(RecommendItem(
+                type=item.get("type", "product"),
+                id=item.get("id", 0),
+                name=item.get("name", ""),
+                description=item.get("description", ""),
+                reason=item.get("reason", ""),
+            ))
+
+        return RecommendResponse(items=items, source="agent")
+    except Exception as e:
+        logger.error("RecommendAgent 调用失败: %s", e, exc_info=True)
+        return RecommendResponse(items=[], source="error")

+ 48 - 0
cfc-langgraph/app/config.py

@@ -0,0 +1,48 @@
+from pydantic_settings import BaseSettings
+from typing import Optional
+
+
+class Settings(BaseSettings):
+    # LLM
+    llm_api_key: str
+    llm_base_url: str = "https://api.deepseek.com/v1"
+    llm_model: str = "deepseek-chat"
+
+    # Embedding
+    embedding_api_key: Optional[str] = None
+    embedding_base_url: Optional[str] = None
+    embedding_model: str = "text-embedding-v3"
+
+    # Java Backend
+    java_base_url: str = "http://localhost:9082"
+    java_context_url: Optional[str] = None
+
+    # LangSmith
+    langchain_tracing_v2: bool = False
+    langchain_api_key: Optional[str] = None
+    langchain_project: str = "cfc-langgraph"
+
+    # Service
+    service_host: str = "0.0.0.0"
+    service_port: int = 9000
+    log_level: str = "info"
+
+    # Chroma
+    chroma_db_path: str = "./data/chroma_db"
+
+    model_config = {"env_file": ".env", "env_file_encoding": "utf-8"}
+
+    @property
+    def effective_embedding_api_key(self) -> str:
+        return self.embedding_api_key or self.llm_api_key
+
+    @property
+    def effective_embedding_base_url(self) -> str:
+        return self.embedding_base_url or self.llm_base_url
+
+    @property
+    def effective_java_context_url(self) -> str:
+        return self.java_context_url or f"{self.java_base_url}/api/ai/context"
+
+
+settings = Settings()

+ 0 - 0
cfc-langgraph/app/graphs/__init__.py


+ 77 - 0
cfc-langgraph/app/graphs/recommend_graph.py

@@ -0,0 +1,77 @@
+from langchain_openai import ChatOpenAI
+from langchain_core.messages import SystemMessage, HumanMessage
+from app.tools.product_tools import search_product_by_keyword, search_activity_by_keyword, search_article_by_keyword
+from app.config import settings
+import json
+import logging
+
+logger = logging.getLogger(__name__)
+
+SYSTEM_PROMPT = """你是一个儿童成长营养推荐助手。根据用户的需求和营养标签, 推荐合适的商品、活动或文章。
+
+推荐原则:
+1. 首先尝试使用搜索工具查找匹配的内容
+2. 如果搜索结果为空, 基于你的知识给出建议
+3. 每项推荐必须附带推荐理由
+4. 以 JSON 格式输出推荐结果
+
+输出格式:
+{
+    "items": [
+        {
+            "source": "tool" 或 "knowledge",
+            "type": "product" / "activity" / "article",
+            "id": 数字,
+            "name": "名称",
+            "description": "描述",
+            "reason": "为什么推荐这个"
+        }
+    ]
+}
+"""
+
+
+class RecommendAgent:
+    def __init__(self):
+        self.llm = ChatOpenAI(
+            model=settings.llm_model,
+            api_key=settings.llm_api_key,
+            base_url=settings.llm_base_url,
+            temperature=0.3,
+        )
+        self.tools = [
+            search_product_by_keyword,
+            search_activity_by_keyword,
+            search_article_by_keyword,
+        ]
+        self.llm_with_tools = self.llm.bind_tools(self.tools)
+
+    async def run(self, query: str, tags: list[str], limit: int = 5) -> dict:
+        """执行推荐 Agent, 返回推荐结果"""
+        # 如果传入了 tags, 构造搜索关键词
+        search_query = query or " ".join(tags)
+
+        messages = [
+            SystemMessage(content=SYSTEM_PROMPT),
+            HumanMessage(content=f"用户需求: {search_query}\n最大返回数量: {limit}\n请搜索并推荐合适的内容。"),
+        ]
+
+        # LangChain Tool calling 自动完成: LLM 决定调哪个 Tool → 工具返回结果 → LLM 组织回答
+        response = await self.llm_with_tools.ainvoke(messages)
+
+        # 尝试解析 JSON 输出
+        content = response.content
+        try:
+            # 提取 JSON 块
+            if "```json" in content:
+                json_str = content.split("```json")[1].split("```")[0].strip()
+            elif "```" in content:
+                json_str = content.split("```")[1].split("```")[0].strip()
+            else:
+                json_str = content.strip()
+            result = json.loads(json_str)
+            return result
+        except (json.JSONDecodeError, IndexError):
+            # 非 JSON 输出, 包装为文本回答
+            logger.warning("Agent 输出非 JSON, raw: %s", content[:200])
+            return {"items": [], "text": content}

+ 8 - 3
cfc-langgraph/app/main.py

@@ -1,16 +1,21 @@
 from fastapi import FastAPI
 from fastapi import FastAPI
-from app.api import health
+from app.api import health, recommend
 
 
 app = FastAPI(title="cfc-langgraph", version="0.1.0")
 app = FastAPI(title="cfc-langgraph", version="0.1.0")
 
 
 app.include_router(health.router)
 app.include_router(health.router)
+app.include_router(recommend.router)
 
 
 
 
 @app.on_event("startup")
 @app.on_event("startup")
 async def startup():
 async def startup():
-    pass  # 后续 Phase 在此初始化 RAG / Agent
+    from app.rag.retriever import RagRetriever
+    retriever = RagRetriever()
+    await retriever.initialize()
 
 
 
 
 @app.on_event("shutdown")
 @app.on_event("shutdown")
 async def shutdown():
 async def shutdown():
-    pass  # 后续 Phase 在此清理资源
+    from app.tools.java_client import JavaClient
+    client = JavaClient()
+    await client.close()

+ 3 - 0
cfc-langgraph/app/models/__init__.py

@@ -0,0 +1,3 @@
+from .common import UserContext, SourceInfo
+from .chat import ChatRequest, ChatResponse
+from .recommend import RecommendRequest, RecommendResponse, RecommendItem

+ 18 - 0
cfc-langgraph/app/models/chat.py

@@ -0,0 +1,18 @@
+from pydantic import BaseModel
+from typing import Optional
+from .common import UserContext, SourceInfo
+
+
+class ChatRequest(BaseModel):
+    query: str
+    user_id: int
+    conversation_id: str = ""
+    context: Optional[UserContext] = None
+
+
+class ChatResponse(BaseModel):
+    answer: str
+    conversation_id: str
+    sources: list[SourceInfo] = []
+    tasks: list[dict] = []
+    trace_id: str = ""

+ 18 - 0
cfc-langgraph/app/models/common.py

@@ -0,0 +1,18 @@
+from pydantic import BaseModel
+from typing import Optional
+
+
+class UserContext(BaseModel):
+    """Java 侧传来的业务上下文"""
+    child_id: Optional[int] = None
+    report_id: Optional[int] = None
+    survey_id: Optional[int] = None
+    family_id: Optional[int] = None
+    mascot_code: Optional[str] = None
+
+
+class SourceInfo(BaseModel):
+    """回答引用来源"""
+    type: str  # knowledge / tool
+    title: str
+    score: Optional[float] = None

+ 29 - 0
cfc-langgraph/app/models/recommend.py

@@ -0,0 +1,29 @@
+from pydantic import BaseModel
+from typing import Optional
+
+
+class RecommendItem(BaseModel):
+    type: str  # product / activity / article
+    id: int
+    name: str
+    description: str = ""
+    cover_image: str = ""
+    price: Optional[float] = None
+    url: str = ""
+    score: float = 0.0
+    reason: str = ""
+
+
+class RecommendRequest(BaseModel):
+    user_id: int
+    query: str = ""
+    tags: list[str] = []
+    types: Optional[list[str]] = None
+    limit: int = 5
+    context: Optional[dict] = None
+
+
+class RecommendResponse(BaseModel):
+    items: list[RecommendItem]
+    source: str = ""  # sql / vector / hybrid
+    trace_id: str = ""

+ 2 - 0
cfc-langgraph/app/rag/__init__.py

@@ -0,0 +1,2 @@
+from .embeddings import get_embeddings
+from .retriever import RagRetriever

+ 15 - 0
cfc-langgraph/app/rag/embeddings.py

@@ -0,0 +1,15 @@
+from langchain_openai import OpenAIEmbeddings
+from app.config import settings
+
+_embeddings = None
+
+
+def get_embeddings():
+    global _embeddings
+    if _embeddings is None:
+        _embeddings = OpenAIEmbeddings(
+            model=settings.embedding_model,
+            api_key=settings.effective_embedding_api_key,
+            base_url=settings.effective_embedding_base_url,
+        )
+    return _embeddings

+ 95 - 0
cfc-langgraph/app/rag/retriever.py

@@ -0,0 +1,95 @@
+from langchain_chroma import Chroma
+from langchain.retrievers import EnsembleRetriever
+from langchain_community.retrievers import BM25Retriever
+from langchain.retrievers.document_compressors import LLMChainExtractor
+from langchain.retrievers import ContextualCompressionRetriever
+from .embeddings import get_embeddings
+from app.config import settings
+from app.tools.java_client import JavaClient
+from typing import Optional
+import logging
+
+logger = logging.getLogger(__name__)
+
+
+class RagRetriever:
+    """混合检索器: ChromaDB 向量 + BM25 关键词 + 可选 LLM 压缩"""
+
+    def __init__(self, collection_name: str = "cfc_knowledge"):
+        embeddings = get_embeddings()
+        self.vectorstore = Chroma(
+            collection_name=collection_name,
+            embedding_function=embeddings,
+            persist_directory=settings.chroma_db_path,
+        )
+        self.java_client = JavaClient()
+        self._bm25_retriever = None
+        self._bm25_texts = []
+
+    async def initialize(self):
+        """从 Java 侧拉取知识库数据, 构建 BM25 索引"""
+        try:
+            articles = await self.java_client.get_published_articles()
+            self._bm25_texts = [
+                f"{a['title']} {a['summary']} {a.get('tags', '')}"
+                for a in articles
+            ]
+            if self._bm25_texts:
+                self._bm25_retriever = BM25Retriever.from_texts(
+                    self._bm25_texts, metadatas=articles
+                )
+                logger.info("BM25 索引就绪: %d 条", len(self._bm25_texts))
+        except Exception as e:
+            logger.warning("BM25 初始化失败(不影响向量检索): %s", e)
+
+    async def retrieve(
+        self,
+        query: str,
+        filters: Optional[dict] = None,
+        k: int = 5,
+        use_compression: bool = False,
+    ) -> list[dict]:
+        """混合检索, 返回 [{content, metadata, score}]"""
+        retrievers = []
+
+        # 向量检索
+        vector_retriever = self.vectorstore.as_retriever(
+            search_kwargs={"k": k, "filter": filters}
+        )
+        retrievers.append(vector_retriever)
+
+        # BM25 检索
+        if self._bm25_retriever:
+            retrievers.append(self._bm25_retriever)
+
+        if len(retrievers) == 1:
+            docs = await retrievers[0].ainvoke(query)
+        else:
+            ensemble = EnsembleRetriever(
+                retrievers=retrievers, weights=[0.6, 0.4]
+            )
+            docs = await ensemble.ainvoke(query)
+
+        # 可选: LLM 压缩去噪
+        if use_compression and docs:
+            from langchain_openai import ChatOpenAI
+            llm = ChatOpenAI(
+                model=settings.llm_model,
+                api_key=settings.llm_api_key,
+                base_url=settings.llm_base_url,
+            )
+            compressor = LLMChainExtractor.from_llm(llm)
+            compression_retriever = ContextualCompressionRetriever(
+                base_compressor=compressor,
+                base_retriever=self.vectorstore.as_retriever(),
+            )
+            docs = await compression_retriever.ainvoke(query)
+
+        results = []
+        for doc in docs:
+            results.append({
+                "content": doc.page_content,
+                "metadata": doc.metadata,
+                "score": getattr(doc, "metadata", {}).get("score", 0),
+            })
+        return results[:k]

+ 1 - 0
cfc-langgraph/app/tools/__init__.py

@@ -0,0 +1 @@
+from .java_client import JavaClient

+ 81 - 0
cfc-langgraph/app/tools/java_client.py

@@ -0,0 +1,81 @@
+import httpx
+from typing import Optional
+from app.config import settings
+import logging
+
+logger = logging.getLogger(__name__)
+
+
+class JavaClient:
+    """Java 后端 HTTP 客户端 (所有 Python→Java 通信的单一入口)"""
+
+    def __init__(self):
+        self.base_url = settings.java_base_url
+        self._client: Optional[httpx.AsyncClient] = None
+
+    async def _get_client(self) -> httpx.AsyncClient:
+        if self._client is None:
+            self._client = httpx.AsyncClient(
+                base_url=self.base_url,
+                timeout=httpx.Timeout(10.0, connect=3.0),
+            )
+        return self._client
+
+    async def close(self):
+        if self._client:
+            await self._client.aclose()
+            self._client = None
+
+    async def get_published_articles(self) -> list[dict]:
+        """获取已发布的文章列表 (用于构建知识库)"""
+        client = await self._get_client()
+        resp = await client.post("/api/article/list", json={"status": "published", "limit": 1000})
+        data = resp.json()
+        if data.get("code") == 200:
+            return data.get("data", [])
+        return []
+
+    async def search_products(self, keyword: str, limit: int = 5) -> list[dict]:
+        """按关键词搜索上架商品"""
+        client = await self._get_client()
+        resp = await client.post("/api/product/search", json={
+            "keyword": keyword, "status": "上架", "limit": limit
+        })
+        data = resp.json()
+        if data.get("code") == 200:
+            return data.get("data", [])
+        return []
+
+    async def search_activities(self, keyword: str, limit: int = 5) -> list[dict]:
+        """按关键词搜索进行中的活动"""
+        client = await self._get_client()
+        resp = await client.post("/api/activity/search", json={
+            "keyword": keyword, "status": "published", "limit": limit
+        })
+        data = resp.json()
+        if data.get("code") == 200:
+            return data.get("data", [])
+        return []
+
+    async def search_articles(self, keyword: str, limit: int = 5) -> list[dict]:
+        """按关键词搜索已发布文章"""
+        client = await self._get_client()
+        resp = await client.post("/api/article/search", json={
+            "keyword": keyword, "status": "published", "limit": limit
+        })
+        data = resp.json()
+        if data.get("code") == 200:
+            return data.get("data", [])
+        return []
+
+    async def get_user_context(self, user_id: int, params: Optional[dict] = None) -> dict:
+        """获取用户上下文 (对应 Java AiContextService)"""
+        client = await self._get_client()
+        resp = await client.post(settings.effective_java_context_url, json={
+            "user_id": str(user_id),
+            "params": params or {},
+        })
+        data = resp.json()
+        if data.get("code") == 200:
+            return data.get("data", {})
+        return {}

+ 58 - 0
cfc-langgraph/app/tools/product_tools.py

@@ -0,0 +1,58 @@
+from langchain_core.tools import tool
+from app.tools.java_client import JavaClient
+import logging
+
+logger = logging.getLogger(__name__)
+_java = JavaClient()
+
+
+@tool
+async def search_product_by_keyword(keyword: str, limit: int = 5) -> str:
+    """按关键词搜索上架商品, 返回 JSON 商品列表"""
+    try:
+        products = await _java.search_products(keyword, limit)
+        if not products:
+            return "[]"
+        return str([{
+            "id": p["id"],
+            "name": p["name"],
+            "price": p.get("price"),
+            "description": p.get("intro") or p.get("description", ""),
+        } for p in products])
+    except Exception as e:
+        logger.warning("搜索商品失败: %s", e)
+        return "[]"
+
+
+@tool
+async def search_activity_by_keyword(keyword: str, limit: int = 5) -> str:
+    """按关键词搜索进行中的活动, 返回 JSON 活动列表"""
+    try:
+        activities = await _java.search_activities(keyword, limit)
+        if not activities:
+            return "[]"
+        return str([{
+            "id": a["id"],
+            "name": a["title"],
+            "description": a.get("description", ""),
+        } for a in activities])
+    except Exception as e:
+        logger.warning("搜索活动失败: %s", e)
+        return "[]"
+
+
+@tool
+async def search_article_by_keyword(keyword: str, limit: int = 5) -> str:
+    """按关键词搜索已发布文章, 返回 JSON 文章列表"""
+    try:
+        articles = await _java.search_articles(keyword, limit)
+        if not articles:
+            return "[]"
+        return str([{
+            "id": a["id"],
+            "name": a["title"],
+            "summary": a.get("summary", ""),
+        } for a in articles])
+    except Exception as e:
+        logger.warning("搜索文章失败: %s", e)
+        return "[]"