dependency injection container

This commit is contained in:
Adam Wilson
2025-06-16 21:46:56 -06:00
parent b20631b9e8
commit 34ab1858c5
15 changed files with 205 additions and 54 deletions
@@ -0,0 +1,18 @@
import sys
from dependency_injector.wiring import Provide, inject
from src.text_generation.dependency_injection_container import DependencyInjectionContainer
from src.text_generation.entrypoints.server import RestApiServer
@inject
def main(
server: RestApiServer = Provide[DependencyInjectionContainer.rest_api_server]
) -> None:
server.listen()
if __name__ == '__main__':
container = DependencyInjectionContainer()
container.init_resources()
container.wire(modules=[__name__])
main()
@@ -6,13 +6,32 @@ from src.text_generation.services.language_models.retrieval_augmented_generation
from src.text_generation.services.logging.file_logging_service import FileLoggingService
class HttpApiController:
def __init__(self):
self.logger = FileLoggingService(filename='text_generation.controller.log').logger
def __init__(
self,
logging_service: FileLoggingService,
text_generation_response_service: TextGenerationResponseService,
rag_response_service: RetrievalAugmentedGenerationResponseService
):
self.logger = logging_service.logger
# TODO: temp debug
self.original_info = self.logger.info
self.logger.info = self.debug_info
self.text_generation_response_service = text_generation_response_service
self.rag_response_service = rag_response_service
self.routes = {}
# Register routes
self.register_routes()
self.text_generation_svc = TextGenerationResponseService()
self.rag_svc = RetrievalAugmentedGenerationResponseService()
def debug_info(self, msg, *args, **kwargs):
try:
return self.original_info(msg, *args, **kwargs)
except TypeError as e:
print(f"Logging error with message: {repr(msg)}")
print(f"Args: {args}")
print(f"Kwargs: {kwargs}")
raise e
def register_routes(self):
"""Register all API routes"""
@@ -58,7 +77,7 @@ class HttpApiController:
start_response('400 Bad Request', response_headers)
return [response_body]
response_text = self.text_generation_svc.invoke(user_prompt=prompt)
response_text = self.text_generation_response_service.invoke(user_prompt=prompt)
response_body = self.format_response(response_text)
http_status_code = 200 # make enum
@@ -84,7 +103,7 @@ class HttpApiController:
start_response('400 Bad Request', response_headers)
return [response_body]
response_text = self.rag_svc.invoke(user_prompt=prompt)
response_text = self.rag_response_service.invoke(user_prompt=prompt)
response_body = self.format_response(response_text)
http_status_code = 200 # make enum
+13 -12
View File
@@ -1,23 +1,24 @@
from wsgiref.simple_server import make_server
from src.text_generation.entrypoints.http_api_controller import HttpApiController
from src.text_generation.services.logging.file_logging_service import FileLoggingService
from wsgiref.simple_server import make_server
class RestApiServer:
def __init__(self):
logging_service = FileLoggingService(filename='text_generation.server.log')
def __init__(
self,
listening_port: int,
logging_service: FileLoggingService,
api_controller: HttpApiController
):
self.listening_port = listening_port
self.logger = logging_service.logger
self.api_controller = api_controller
def listen(self):
try:
port = 9999
controller = HttpApiController()
with make_server('', port, controller) as wsgi_srv:
print(f'listening on port {port}...')
with make_server('', self.listening_port, self.api_controller) as wsgi_srv:
print(f'listening on port {self.listening_port}...')
wsgi_srv.serve_forever()
except Exception as e:
self.logger.debug(e)
if __name__ == '__main__':
srv = RestApiServer()
srv.listen()
self.logger.debug(e)