30 lines
1.6 KiB
Python
30 lines
1.6 KiB
Python
from typing import Annotated
|
|
|
|
from fastapi import Depends
|
|
from langchain_huggingface import HuggingFaceEndpointEmbeddings
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.chat_qa_query.service.query_qa_service import QueryQAService
|
|
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
|
|
|
|
"""
|
|
普通对象 / 客户端 / 仓库 → 用 return不需要关闭、不需要释放、不需要上下文管理 → 直接 return 实例
|
|
数据库连接 / 会话 / 需要自动释放的资源 → 用 yield必须用完自动关闭 / 释放 / 回滚 → 必须用 yield
|
|
"""
|
|
|
|
async def get_assistant_session()->AsyncSession:
|
|
async with db_assistant_mysql_client_manager.session_factory() as db_assistant_session:
|
|
yield db_assistant_session
|
|
|
|
async def get_archive_messages_repository(session: Annotated[AsyncSession, Depends(get_assistant_session)])->ArchiveMessagesRepository:
|
|
return ArchiveMessagesRepository(session)
|
|
|
|
async def get_qa_milvus_repository()->QARepository:
|
|
return QARepository()
|
|
|
|
async def get_qa_query_service(archive_mysql_repository: Annotated[ArchiveMessagesRepository, Depends(get_archive_messages_repository)],
|
|
qa_milvus_repository: Annotated[QARepository, Depends(get_qa_milvus_repository)]) -> QueryQAService:
|
|
return QueryQAService(archive_mysql_repository=archive_mysql_repository,
|
|
qa_milvus_repository=qa_milvus_repository) |