Compare commits
5 Commits
c66d571e43
...
master
| Author | SHA1 | Date | |
|---|---|---|---|
| 54eb3680f1 | |||
| b04242bbb2 | |||
| c45ce61e9c | |||
| e0b7ed4364 | |||
| c136426589 |
5
.gitignore
vendored
5
.gitignore
vendored
@@ -65,4 +65,7 @@ media/
|
||||
.DS_Store
|
||||
|
||||
logs/
|
||||
packages/
|
||||
packages/
|
||||
*.db
|
||||
*.db-shm
|
||||
*.db-wal
|
||||
28
agent.py
28
agent.py
@@ -5,16 +5,24 @@ import sqlite3
|
||||
from langchain_tavily import TavilySearch
|
||||
from langgraph.checkpoint.sqlite import SqliteSaver
|
||||
|
||||
from services.db_service import init_db
|
||||
|
||||
from dotenv import load_dotenv
|
||||
import os
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# 连接sqlite
|
||||
connection = sqlite3.connect("chief.db", check_same_thread=False)
|
||||
connection = sqlite3.connect(os.getenv("SQLITE_DB_PATH"), check_same_thread=False)
|
||||
# 初始化checkpointer
|
||||
checkpointer = SqliteSaver(connection)
|
||||
# 自动建表
|
||||
checkpointer.setup()
|
||||
init_db()
|
||||
|
||||
# web搜索工具,使用tavily作为web搜索工具
|
||||
web_search = TavilySearch(
|
||||
tavily_api_key="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",
|
||||
tavily_api_key=os.getenv("TAVILY_API_KEY"),
|
||||
max_results=5,
|
||||
topic="general"
|
||||
)
|
||||
@@ -30,10 +38,10 @@ system_prompt = """
|
||||
"""
|
||||
|
||||
model = init_chat_model(
|
||||
model="qwen-turbo",
|
||||
model=os.getenv("MODEL_NAME"),
|
||||
model_provider="openai",
|
||||
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
api_key="sk-52bcd98e9c1d45908437c4e8706eefff",
|
||||
base_url=os.getenv("MODEL_BASE_URL"),
|
||||
api_key=os.getenv("MODEL_API_KEY"),
|
||||
temperature=1.5,
|
||||
)
|
||||
|
||||
@@ -43,3 +51,13 @@ agent = create_agent(
|
||||
checkpointer=checkpointer, # 记忆
|
||||
system_prompt=system_prompt # 系统提示词
|
||||
)
|
||||
|
||||
title_prompt = """
|
||||
请根据以下对话内容,生成一句不超过 20 字的标题,用于记录本次对话主题。
|
||||
只返回标题,不要解释。
|
||||
"""
|
||||
|
||||
title_agent = create_agent(
|
||||
model=model,
|
||||
system_prompt=title_prompt
|
||||
)
|
||||
|
||||
BIN
chief.db-shm
BIN
chief.db-shm
Binary file not shown.
BIN
chief.db-wal
BIN
chief.db-wal
Binary file not shown.
9
model.py
9
model.py
@@ -1,4 +1,11 @@
|
||||
from pydantic import BaseModel
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
message: str
|
||||
username: str
|
||||
message: str
|
||||
thread_id: str
|
||||
|
||||
class SessionMessageResponse(BaseModel):
|
||||
username: str
|
||||
title: str
|
||||
thread_id: str
|
||||
@@ -5,5 +5,6 @@ langchain-openai~=1.2.2
|
||||
langchain-tavily~=0.2.18
|
||||
langgraph~=1.2.2
|
||||
langgraph-checkpoint-sqlite~=3.1.0
|
||||
pydantic~=2.12.5
|
||||
langchain-core~=1.4.0
|
||||
pydantic~=2.13.4
|
||||
langchain-core~=1.4.0
|
||||
python-dotenv~=1.2.2
|
||||
21
router.py
21
router.py
@@ -1,10 +1,10 @@
|
||||
import uuid
|
||||
from typing import List
|
||||
|
||||
from fastapi import APIRouter
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from model import ChatRequest
|
||||
from service import get_messages, clear_messages, search_recipes
|
||||
from model import ChatRequest, SessionMessageResponse
|
||||
from services.agent_service import get_messages, clear_messages, search_recipes, get_user_session_messages
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -13,16 +13,21 @@ router = APIRouter()
|
||||
async def chat_endpoint(request: ChatRequest):
|
||||
"""流式对话"""
|
||||
return StreamingResponse(
|
||||
search_recipes(request.message, str(uuid.uuid4())),
|
||||
search_recipes(request),
|
||||
media_type="text/event-stream"
|
||||
)
|
||||
|
||||
|
||||
@router.get("/chat/messages")
|
||||
async def get_chat_messages(thread_id: str):
|
||||
@router.get("/chat/messages/session", response_model=List[SessionMessageResponse])
|
||||
def get_session_messages(username: str):
|
||||
"""获取历史消息"""
|
||||
messages = get_messages(thread_id)
|
||||
return {"messages": messages}
|
||||
return get_user_session_messages(username)
|
||||
|
||||
|
||||
@router.get("/chat/messages")
|
||||
def get_chat_messages(thread_id: str):
|
||||
"""获取历史消息"""
|
||||
return get_messages(thread_id)
|
||||
|
||||
|
||||
@router.delete("/chat/messages")
|
||||
|
||||
@@ -1,23 +1,38 @@
|
||||
from langchain_core.messages import HumanMessage, AIMessageChunk, AIMessage
|
||||
|
||||
from agent import agent, checkpointer
|
||||
from agent import agent, checkpointer, title_agent
|
||||
from model import ChatRequest
|
||||
from services.db_service import exists_session, save_session, query_session_messages
|
||||
from utils import dicts_to_messages
|
||||
|
||||
|
||||
async def search_recipes(prompt: str, thread_id: str):
|
||||
def generate_title(messages: list[dict[str, str]]) -> str:
|
||||
response = title_agent.invoke({
|
||||
"messages": dicts_to_messages(messages)
|
||||
})
|
||||
return response["messages"][-1].content.strip()
|
||||
|
||||
|
||||
async def search_recipes(request: ChatRequest):
|
||||
"""调用agent搜索食谱"""
|
||||
print(f"[用户]: {prompt}, thread_id: {thread_id}")
|
||||
print(f"[用户] {request.username}: {request.message}, thread_id: {request.thread_id}")
|
||||
try:
|
||||
message = HumanMessage(content=prompt)
|
||||
message = HumanMessage(content=request.message)
|
||||
|
||||
# 流式调用Agent
|
||||
for chunk, metadata in agent.stream(
|
||||
{"messages": [message]},
|
||||
{"configurable": {"thread_id": thread_id}},
|
||||
{"configurable": {"thread_id": request.thread_id}},
|
||||
stream_mode="messages"
|
||||
):
|
||||
if isinstance(chunk, AIMessageChunk) and chunk.content:
|
||||
yield chunk.content
|
||||
|
||||
# 总结对话标题并保存
|
||||
if not exists_session(request.thread_id, request.username):
|
||||
title = generate_title(get_messages(request.thread_id))
|
||||
save_session(request.thread_id, request.username, title)
|
||||
|
||||
except Exception as e:
|
||||
print(f"\n[错误]: {str(e)}")
|
||||
yield "信息检索失败,试试看手动输入食物列表?"
|
||||
@@ -60,6 +75,13 @@ def get_messages(thread_id: str) -> list[dict[str, str]]:
|
||||
if isinstance(msg, HumanMessage):
|
||||
result.append({"role": "user", "content": msg.content})
|
||||
elif isinstance(msg, AIMessage):
|
||||
result.append({"role": "assistant", "content": msg.content})
|
||||
result.append({"role": "ai", "content": msg.content})
|
||||
|
||||
return result
|
||||
|
||||
|
||||
# 查询历史会话消息
|
||||
def get_user_session_messages(username: str):
|
||||
"""查询历史会话消息"""
|
||||
print(f"查询历史会话消息,username: {username}")
|
||||
return query_session_messages(username)
|
||||
85
services/db_service.py
Normal file
85
services/db_service.py
Normal file
@@ -0,0 +1,85 @@
|
||||
import sqlite3
|
||||
import os
|
||||
from typing import List
|
||||
|
||||
from model import SessionMessageResponse
|
||||
|
||||
|
||||
def init_db():
|
||||
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
|
||||
|
||||
# 业务表:会话
|
||||
conn.execute("""
|
||||
CREATE TABLE IF NOT EXISTS session (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
thread_id TEXT NOT NULL UNIQUE,
|
||||
username TEXT NOT NULL,
|
||||
title TEXT NOT NULL,
|
||||
created_time DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
""")
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
|
||||
def query_session_messages(username: str) -> List[SessionMessageResponse]:
|
||||
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
|
||||
conn.row_factory = sqlite3.Row
|
||||
cursor = conn.cursor()
|
||||
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT username, title, thread_id
|
||||
FROM session
|
||||
WHERE username = ?
|
||||
ORDER BY created_time DESC
|
||||
""",
|
||||
(username,)
|
||||
)
|
||||
|
||||
rows = cursor.fetchall()
|
||||
conn.close()
|
||||
|
||||
return [
|
||||
SessionMessageResponse(
|
||||
username=row["username"],
|
||||
title=row["title"],
|
||||
thread_id=row["thread_id"],
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
|
||||
|
||||
def exists_session(thread_id: str, username: str) -> bool:
|
||||
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
|
||||
cursor = conn.cursor()
|
||||
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT 1
|
||||
FROM session
|
||||
WHERE thread_id = ? AND username = ?
|
||||
LIMIT 1
|
||||
""",
|
||||
(thread_id, username)
|
||||
)
|
||||
|
||||
result = cursor.fetchone()
|
||||
conn.close()
|
||||
|
||||
return result is not None
|
||||
|
||||
|
||||
def save_session(thread_id: str, username: str, title: str):
|
||||
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT OR REPLACE INTO session
|
||||
(thread_id, username, title)
|
||||
VALUES (?, ?, ?)
|
||||
""",
|
||||
(thread_id, username, title)
|
||||
)
|
||||
conn.commit()
|
||||
conn.close()
|
||||
5
test.py
Normal file
5
test.py
Normal file
@@ -0,0 +1,5 @@
|
||||
from agent import model
|
||||
|
||||
if __name__ == '__main__':
|
||||
response = model.invoke("西红柿、鸡蛋")
|
||||
print(response)
|
||||
11
utils.py
Normal file
11
utils.py
Normal file
@@ -0,0 +1,11 @@
|
||||
from langchain_core.messages import HumanMessage, AIMessage
|
||||
|
||||
|
||||
def dicts_to_messages(messages: list[dict]):
|
||||
result = []
|
||||
for m in messages:
|
||||
if m["role"] == "user":
|
||||
result.append(HumanMessage(content=m["content"]))
|
||||
elif m["role"] == "ai":
|
||||
result.append(AIMessage(content=m["content"]))
|
||||
return result
|
||||
Reference in New Issue
Block a user