import os
from dotenv import load_dotenv
from langchain_openai import ChatOpenAI
from langchain_openai.embeddings import OpenAIEmbeddings

# 加载环境变量
load_dotenv()
os.environ["OPENAI_API_KEY"]=os.getenv("OPENAI_API_KEY")

# 初始化模型和嵌入
# 使用硅基流动的千问模型
llm = ChatOpenAI(
    model="Qwen/QwQ-32B",
    openai_api_base="https://api.siliconflow.cn/v1",
    openai_api_key=os.getenv("SILICONFLOW_API_KEY")
)

# 使用硅基流动的BAAI/bge-large-zh-v1.5 embedding模型
embeddings = OpenAIEmbeddings(
    model="BAAI/bge-large-zh-v1.5",
    openai_api_base="https://api.siliconflow.cn/v1",
    openai_api_key=os.getenv("SILICONFLOW_API_KEY"),  # 需要在.env文件中添加硅基流动的API密钥
    chunk_size=32  # 设置批处理大小为32,符合硅基流动的限制
)

from langchain_community.document_loaders import WebBaseLoader
from langchain_text_splitters import RecursiveCharacterTextSplitter

# 知识来源URL列表
urls=[
    "https://arxiv.org/abs/2509.16990",
    "https://arxiv.org/abs/2509.16990",
]

# 加载网页内容
docs = WebBaseLoader(web_paths=urls).load()

# 文档分块(调整为适应BAAI/bge-large-zh-v1.5的512 token限制)
# 一般来说,1个token约等于0.75个英文单词,中文字符约1-2个token
# 为了安全起见,设置chunk_size为400字符,overlap为50
splitter = RecursiveCharacterTextSplitter(
    chunk_size=400,  # 减小chunk size以适应512 token限制
    chunk_overlap=50,
    length_function=len,
    separators=["\n\n", "\n", " ", ""]
)
final_docs = splitter.split_documents(docs)

from langchain_community.vectorstores import Chroma

# 创建Chroma向量存储
vector_store = Chroma.from_documents(
    documents=final_docs,
    embedding=embeddings,
    collection_name="rag-chrome"
)

# 转换为检索器接口
retriever = vector_store.as_retriever()

from langchain.tools.retriever import create_retriever_tool

# 将检索器包装成工具
retriever_tool = create_retriever_tool(
    retriever=retriever, 
    name="retriever_blog_post", 
    description="搜索并返回关于Lilian Weng博客中LLM代理、提示工程等相关信息"
)

# 注册工具
tools = [retriever_tool]

# 将工具附加到LLM
llm_with_tool = llm.bind_tools(tools)

from typing import TypedDict, Sequence, Annotated
from langchain_core.messages import BaseMessage
from langgraph.graph import add_messages

class AgentState(TypedDict):
    messages: Annotated[Sequence[BaseMessage], add_messages]

def llm_decision_maker(state:AgentState):
    print("----调用LLM决策器----")
    
    # 获取最新用户消息
    last_message = state["messages"][-1]
    question = last_message.content
    response = llm_with_tool.invoke(question)
  
    # 返回新消息添加到状态
    return {"messages": [response]}


from langchain_core.prompts import PromptTemplate
from typing import Literal

def grade_documents(state:AgentState)->Literal["Output Generator", "Query Rewriter"]:
    print("----调用评估器检查相关性----")
    
    prompt=PromptTemplate(
        template="""你是一个评估器,判断文档是否与用户问题相关。
                    文档:{context}
                    用户问题:{question}
                    如果文档讨论或包含与用户问题相关的信息,标记为相关。
                    请只回答'yes'或'no',不要添加其他内容。""",
                    input_variables=["context", "question"]
    )
    
    chain = prompt | llm

    message = state["messages"]
    last_message = message[-1]
    docs = last_message.content
    question = message[0].content

    response = chain.invoke({"context": docs, "question": question})
    # 从响应中提取文本内容
    score = response.content.strip().lower()

    if "yes" in score:
        print("----决策:文档相关----")
        return "generator"
    else:
        print("----决策:文档不相关----")
        return "rewriter"


from langchain import hub

def generate(state:AgentState):
    print("----RAG输出生成----")
    
    message=state["messages"]
    question=message[0].content
    last_message = message[-1]
    docs = last_message.content

    prompt=hub.pull("rlm/rag-prompt")
    rag_chain=prompt | llm

    response=rag_chain.invoke({"context": docs, "question": question})
    
    print(f"这是我的响应:{response}")
    return {"messages": [response]}


from langchain_core.messages import HumanMessage

def tools_condition(state:AgentState):
    """检查是否需要调用工具"""
    messages = state["messages"]
    last_message = messages[-1]
    
    # 检查消息是否包含工具调用
    if hasattr(last_message, 'tool_calls') and last_message.tool_calls:
        return "tools"
    else:
        return END

def rewrite(state:AgentState):
    print("----转换查询----")
    message=state["messages"]
    
    question=message[0].content
    
    input= [HumanMessage(content=f"""分析输入并推理其语义意图。
                    原始问题:{question}
                    请提出一个改进的问题:""")
    ]

    response=llm.invoke(input)
    return {"messages": [response]}

    
from langgraph.graph import StateGraph, END, START
from langgraph.prebuilt import ToolNode

# 创建工具节点
retriever_node = ToolNode(tools)

# 初始化工作流
workflow = StateGraph(AgentState)

# 添加节点和边
workflow.add_node("LLM Decision Maker", llm_decision_maker)
workflow.add_node("Vector Retriever", retriever_node)
workflow.add_node("Output Generator", generate)
workflow.add_node("Query Rewriter", rewrite)
workflow.add_edge(START, "LLM Decision Maker")
workflow.add_conditional_edges("LLM Decision Maker", 
                               tools_condition,
                               {"tools": "Vector Retriever", 
                                END:END} )
workflow.add_conditional_edges("Vector Retriever", 
                               grade_documents,
                               {"generator": "Output Generator", 
                                "rewriter":"Query Rewriter"} )
workflow.add_edge("Output Generator", END)
workflow.add_edge("Query Rewriter", "LLM Decision Maker")

# 编译工作流
app = workflow.compile()

response = app.invoke({"messages":["What are agents and prompt engineering use ?"]})
print(response)

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐