59 lines
2.8 KiB
Python
59 lines
2.8 KiB
Python
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())
|