sales-assistant-py-new/app/chat_qa_query/query_graph.py

59 lines
2.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import asyncio
from datetime import datetime, timedelta
from typing import List, Dict, Any
from langgraph.constants import START, END
from langgraph.graph import StateGraph
from app.chat_qa.nodes.finalize_node import finalize
from app.chat_qa.nodes.generate_qa_node import generate_qa
from app.chat_qa.nodes.quality_check_node import quality_check
from app.chat_qa.nodes.query_message_node import query_message
from app.chat_qa.nodes.route_by_quality_node import route_by_quality
from app.chat_qa_query.query_context import QueryQAContext
from app.chat_qa_query.query_nodes.node_answer_output import answer_output
from app.chat_qa_query.query_nodes.node_rewrite import rewrite_query
from app.chat_qa_query.query_nodes.node_search_embedding import search_embed
from app.chat_qa_query.query_qa_state import QueryQAState
from app.client.embedding_client_manager import embedding_client
from app.client.milvus_client_manager import milvus_client
from app.client.mysql_client_manager import db_assistant_mysql_client_manager
from app.repository.archive_messages_repository import ArchiveMessagesRepository
from app.repository.milvus.message_qa_repository import QARepository
from app.repository.milvus.summary_repository import SummaryRepository
graph_builder = StateGraph(state_schema=QueryQAState, context_schema=QueryQAContext)
graph_builder.add_node("rewrite_query", rewrite_query)
graph_builder.add_node("search_embed", search_embed)
graph_builder.add_node("answer_output", answer_output)
graph_builder.add_edge(START, "rewrite_query")
graph_builder.add_edge("rewrite_query", "search_embed")
graph_builder.add_edge("search_embed", "answer_output")
graph_builder.add_edge("answer_output", END)
qa_graph = graph_builder.compile()
if __name__ == '__main__':
async def test():
db_assistant_mysql_client_manager.init()
embedding_client.init()
milvus_client.init()
async with db_assistant_mysql_client_manager.session_factory() as db_session:
archive_messages_mysql_repository = ArchiveMessagesRepository(db_session)
qa_milvus_repository = QARepository()
state: QueryQAState = QueryQAState(original_query="老师你好我想咨询一下3d建模费用和就业情况", history=['2026-07-10T15:42:13 EXTERNAL 老师这个课程多少钱啊\n', '2026-07-10T15:40:34 EXTERNAL 老师 我想报班\n'])
context = QueryQAContext(meta_mysql_repository=archive_messages_mysql_repository,
qa_milvus_repository=qa_milvus_repository
)
async for chunk in qa_graph.astream(input=state, context=context, stream_mode="custom"):
for item in chunk:
print("=================="+item)
milvus_client.close()
await embedding_client.close()
asyncio.run(test())