mirror of
https://github.com/lightbroker/llmsecops-research.git
synced 2026-08-16 16:20:36 +02:00
dependency injection container
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user