to 22
This commit is contained in:
@@ -1,6 +1,3 @@
|
||||
"""
|
||||
DI контейнер на основе dishka
|
||||
"""
|
||||
from dishka import Container, Provider, Scope, provide
|
||||
from fastapi import Request
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -19,6 +16,7 @@ from src.domain.repositories.document_repository import IDocumentRepository
|
||||
from src.domain.repositories.conversation_repository import IConversationRepository
|
||||
from src.domain.repositories.message_repository import IMessageRepository
|
||||
from src.domain.repositories.collection_access_repository import ICollectionAccessRepository
|
||||
from src.domain.repositories.vector_repository import IVectorRepository
|
||||
from src.infrastructure.external.yandex_ocr import YandexOCRService
|
||||
from src.infrastructure.external.deepseek_client import DeepSeekClient
|
||||
from src.application.services.document_parser_service import DocumentParserService
|
||||
@@ -28,7 +26,14 @@ from src.application.use_cases.document_use_cases import DocumentUseCases
|
||||
from src.application.use_cases.conversation_use_cases import ConversationUseCases
|
||||
from src.application.use_cases.message_use_cases import MessageUseCases
|
||||
from src.domain.entities.user import User
|
||||
|
||||
from src.shared.config import settings
|
||||
from qdrant_client import QdrantClient
|
||||
from src.infrastructure.repositories.qdrant.vector_repository import QdrantVectorRepository
|
||||
from src.application.services.embedding_service import EmbeddingService
|
||||
from src.application.services.reranker_service import RerankerService
|
||||
from src.application.services.rag_service import RAGService
|
||||
from src.application.services.text_splitter import TextSplitter
|
||||
from src.application.use_cases.rag_use_cases import RAGUseCases
|
||||
|
||||
class DatabaseProvider(Provider):
|
||||
@provide(scope=Scope.REQUEST)
|
||||
@@ -81,6 +86,44 @@ class ServiceProvider(Provider):
|
||||
return DocumentParserService(ocr_service)
|
||||
|
||||
|
||||
class VectorServiceProvider(Provider):
|
||||
@provide(scope=Scope.APP)
|
||||
def get_qdrant_client(self) -> QdrantClient:
|
||||
return QdrantClient(host=settings.QDRANT_HOST, port=settings.QDRANT_PORT)
|
||||
|
||||
@provide(scope=Scope.APP)
|
||||
def get_vector_repository(self, client: QdrantClient) -> IVectorRepository:
|
||||
return QdrantVectorRepository(client=client, vector_size=768)
|
||||
|
||||
@provide(scope=Scope.APP)
|
||||
def get_embedding_service(self) -> EmbeddingService:
|
||||
return EmbeddingService()
|
||||
|
||||
@provide(scope=Scope.APP)
|
||||
def get_reranker_service(self, embedding_service: EmbeddingService) -> RerankerService:
|
||||
return RerankerService(fallback_encoder=embedding_service.model)
|
||||
|
||||
@provide(scope=Scope.APP)
|
||||
def get_text_splitter(self) -> TextSplitter:
|
||||
return TextSplitter()
|
||||
|
||||
@provide(scope=Scope.APP)
|
||||
def get_rag_service(
|
||||
self,
|
||||
vector_repo: IVectorRepository,
|
||||
embedding_service: EmbeddingService,
|
||||
reranker_service: RerankerService,
|
||||
deepseek_client: DeepSeekClient,
|
||||
text_splitter: TextSplitter
|
||||
) -> RAGService:
|
||||
return RAGService(
|
||||
vector_repository=vector_repo,
|
||||
embedding_service=embedding_service,
|
||||
reranker_service=reranker_service,
|
||||
deepseek_client=deepseek_client,
|
||||
splitter=text_splitter,
|
||||
)
|
||||
|
||||
class AuthProvider(Provider):
|
||||
@provide(scope=Scope.REQUEST)
|
||||
async def get_current_user(self, request: Request, user_repo: IUserRepository) -> User:
|
||||
@@ -131,6 +174,16 @@ class UseCaseProvider(Provider):
|
||||
) -> MessageUseCases:
|
||||
return MessageUseCases(message_repo, conversation_repo)
|
||||
|
||||
@provide(scope=Scope.REQUEST)
|
||||
def get_rag_use_cases(
|
||||
self,
|
||||
rag_service: RAGService,
|
||||
document_repo: IDocumentRepository,
|
||||
conversation_repo: IConversationRepository,
|
||||
message_repo: IMessageRepository
|
||||
) -> RAGUseCases:
|
||||
return RAGUseCases(rag_service, document_repo, conversation_repo, message_repo)
|
||||
|
||||
|
||||
def create_container() -> Container:
|
||||
container = Container()
|
||||
@@ -139,5 +192,6 @@ def create_container() -> Container:
|
||||
container.add_provider(ServiceProvider())
|
||||
container.add_provider(AuthProvider())
|
||||
container.add_provider(UseCaseProvider())
|
||||
container.add_provider(VectorServiceProvider())
|
||||
return container
|
||||
|
||||
|
||||
Reference in New Issue
Block a user