81 lines
3.4 KiB
Python
81 lines
3.4 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.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))
|