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.context import QAContext 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.qa_state import QAState 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=QAState, context_schema=QAContext) graph_builder.add_node("generate_qa", generate_qa) graph_builder.add_node("quality_check", quality_check) graph_builder.add_node("finalize", finalize) graph_builder.add_edge(START, "generate_qa") graph_builder.add_edge("generate_qa", "quality_check") graph_builder.add_conditional_edges("quality_check", route_by_quality, { "pass": "finalize", "retry": "generate_qa", "fail": "finalize" }) graph_builder.add_edge("finalize", END) graph = graph_builder.compile() if __name__ == '__main__': async def test(day_str: str): print(f"开始日期:{day_str}") 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() # 根据日期查询向量库 repo = SummaryRepository() message_list: List[Dict[str, Any]] = repo.query_by_date(day_str) print(f"向量库数据:{len(message_list)}") message_content_list = [message["message_context"] for message in message_list] for message_content in message_content_list: state: QAState = QAState(msg_time=day_str, chat_history=message_content) context = QAContext(meta_mysql_repository=archive_messages_mysql_repository, qa_milvus_repository=qa_milvus_repository ) async for chunk in graph.astream(input=state, context=context,stream_mode="custom"): print(chunk) milvus_client.close() await embedding_client.close() start = datetime(2026, 5, 29) end = datetime(2026, 6, 16) result = [] temp = start while temp <= end: # 按 年-月-日 无补零格式输出 date_format = f"{temp.year}-{str(temp.month).zfill(2)}-{str(temp.day).zfill(2)}" result.append(date_format) temp += timedelta(days=1) for date in result: asyncio.run(test(date))