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="发票?", history=['2026-07-10T15:42:13 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())