mirror of
https://github.com/lightbroker/llmsecops-research.git
synced 2026-08-05 19:08:42 +02:00
naming updates; fix static analysis script
This commit is contained in:
@@ -6,8 +6,8 @@ from src.text_generation.entrypoints.http_api_controller import HttpApiControlle
|
||||
from src.text_generation.entrypoints.server import RestApiServer
|
||||
from src.text_generation.services.logging.json_web_traffic_logging_service import JSONWebTrafficLoggingService
|
||||
from src.text_generation.services.nlp.semantic_similarity_service import SemanticSimilarityService
|
||||
from src.text_generation.services.nlp.text_generation_response_service import TextGenerationResponseService
|
||||
from src.text_generation.services.nlp.retrieval_augmented_generation_response_service import RetrievalAugmentedGenerationResponseService
|
||||
from src.text_generation.services.nlp.text_generation_completion_service import TextGenerationCompletionService
|
||||
from src.text_generation.services.nlp.retrieval_augmented_generation_completion_service import RetrievalAugmentedGenerationCompletionService
|
||||
from src.text_generation.services.guardrails.generated_text_guardrail_service import GeneratedTextGuardrailService
|
||||
from src.text_generation.services.guidelines.rag_guidelines_service import RetrievalAugmentedGenerationGuidelinesService
|
||||
from src.text_generation.services.utilities.response_processing_service import ResponseProcessingService
|
||||
@@ -40,7 +40,7 @@ class DependencyInjectionContainer(containers.DeclarativeContainer):
|
||||
)
|
||||
|
||||
rag_response_service = providers.Factory(
|
||||
RetrievalAugmentedGenerationResponseService,
|
||||
RetrievalAugmentedGenerationCompletionService,
|
||||
foundation_model=foundation_model,
|
||||
embedding_model=embedding_model,
|
||||
rag_guidelines_service=rag_guidelines_service,
|
||||
@@ -67,7 +67,7 @@ class DependencyInjectionContainer(containers.DeclarativeContainer):
|
||||
)
|
||||
|
||||
text_generation_response_service = providers.Factory(
|
||||
TextGenerationResponseService,
|
||||
TextGenerationCompletionService,
|
||||
foundation_model
|
||||
)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import json
|
||||
import traceback
|
||||
|
||||
from src.text_generation.services.logging.abstract_web_traffic_logging_service import AbstractWebTrafficLoggingService
|
||||
from src.text_generation.services.nlp.abstract_language_model_response_service import AbstractLanguageModelResponseService
|
||||
from src.text_generation.services.nlp.abstract_text_generation_completion_service import AbstractTextGenerationCompletionService
|
||||
from src.text_generation.services.guardrails.abstract_generated_text_guardrail_service import AbstractGeneratedTextGuardrailService
|
||||
|
||||
|
||||
@@ -11,8 +11,8 @@ class HttpApiController:
|
||||
def __init__(
|
||||
self,
|
||||
logging_service: AbstractWebTrafficLoggingService,
|
||||
text_generation_response_service: AbstractLanguageModelResponseService,
|
||||
rag_response_service: AbstractLanguageModelResponseService,
|
||||
text_generation_response_service: AbstractTextGenerationCompletionService,
|
||||
rag_response_service: AbstractTextGenerationCompletionService,
|
||||
generated_text_guardrail_service: AbstractGeneratedTextGuardrailService
|
||||
):
|
||||
self.logging_service = logging_service
|
||||
|
||||
@@ -7,5 +7,5 @@ class AbstractRetrievalAugmentedGenerationGuidelinesService(abc.ABC):
|
||||
raise NotImplementedError
|
||||
|
||||
@abc.abstractmethod
|
||||
def create_context(self, user_prompt: str) -> str:
|
||||
def create_guidelines_context(self, user_prompt: str) -> str:
|
||||
raise NotImplementedError
|
||||
@@ -58,7 +58,7 @@ class RetrievalAugmentedGenerationGuidelinesService(
|
||||
|
||||
# public methods
|
||||
|
||||
def create_context(self, user_prompt: str) -> str:
|
||||
def create_guidelines_context(self, user_prompt: str) -> str:
|
||||
return self._create_context(user_prompt)
|
||||
|
||||
def get_prompt_template(self):
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
import abc
|
||||
|
||||
|
||||
class AbstractLanguageModelResponseService(abc.ABC):
|
||||
class AbstractTextGenerationCompletionService(abc.ABC):
|
||||
@abc.abstractmethod
|
||||
def invoke(self, user_prompt: str) -> str:
|
||||
raise NotImplementedError
|
||||
+2
-2
@@ -1,7 +1,7 @@
|
||||
from src.text_generation.services.nlp.abstract_language_model_response_service import AbstractLanguageModelResponseService
|
||||
from src.text_generation.services.nlp.abstract_text_generation_completion_service import AbstractTextGenerationCompletionService
|
||||
|
||||
|
||||
class FakeLanguageModelResponseService(AbstractLanguageModelResponseService):
|
||||
class FakeTextGenerationCompletionService(AbstractTextGenerationCompletionService):
|
||||
|
||||
def invoke(self, user_prompt: str) -> str:
|
||||
|
||||
+3
-4
@@ -3,13 +3,12 @@ from langchain.prompts import PromptTemplate
|
||||
|
||||
from src.text_generation.ports.abstract_embedding_model import AbstractEmbeddingModel
|
||||
from src.text_generation.ports.abstract_foundation_model import AbstractFoundationModel
|
||||
from src.text_generation.services.nlp.abstract_language_model_response_service import AbstractLanguageModelResponseService
|
||||
from src.text_generation.services.nlp.abstract_text_generation_completion_service import AbstractTextGenerationCompletionService
|
||||
from src.text_generation.services.guidelines.abstract_rag_guidelines_service import AbstractRetrievalAugmentedGenerationGuidelinesService
|
||||
from src.text_generation.services.utilities.abstract_response_processing_service import AbstractResponseProcessingService
|
||||
|
||||
|
||||
class RetrievalAugmentedGenerationResponseService(AbstractLanguageModelResponseService):
|
||||
|
||||
class RetrievalAugmentedGenerationCompletionService(AbstractTextGenerationCompletionService):
|
||||
def __init__(
|
||||
self,
|
||||
foundation_model: AbstractFoundationModel,
|
||||
@@ -32,7 +31,7 @@ class RetrievalAugmentedGenerationResponseService(AbstractLanguageModelResponseS
|
||||
template=self.rag_guidelines_service.get_prompt_template(),
|
||||
input_variables=["context", "question"]
|
||||
)
|
||||
context = self.rag_guidelines_service.create_context(user_prompt)
|
||||
context = self.rag_guidelines_service.create_guidelines_context(user_prompt)
|
||||
chain = prompt | self.language_model_pipeline | StrOutputParser()
|
||||
raw_response = chain.invoke({
|
||||
"context": context,
|
||||
+6
-4
@@ -2,19 +2,21 @@ from langchain.prompts import PromptTemplate
|
||||
from langchain_core.output_parsers import StrOutputParser
|
||||
from langchain_core.runnables import RunnablePassthrough
|
||||
|
||||
from src.text_generation.services.nlp.abstract_language_model_response_service import AbstractLanguageModelResponseService
|
||||
from src.text_generation.common.constants import Constants
|
||||
from src.text_generation.services.nlp.abstract_text_generation_completion_service import AbstractTextGenerationCompletionService
|
||||
from src.text_generation.ports.abstract_foundation_model import AbstractFoundationModel
|
||||
|
||||
|
||||
class TextGenerationResponseService(AbstractLanguageModelResponseService):
|
||||
class TextGenerationCompletionService(AbstractTextGenerationCompletionService):
|
||||
|
||||
def __init__(self, foundation_model: AbstractFoundationModel):
|
||||
super().__init__()
|
||||
self.language_model_pipeline = foundation_model.create_pipeline()
|
||||
self.constants = Constants()
|
||||
|
||||
def _extract_assistant_response(self, text):
|
||||
if "<|assistant|>" in text:
|
||||
return text.split("<|assistant|>")[-1].strip()
|
||||
if self.constants.ASSISTANT_TOKEN in text:
|
||||
return text.split(self.constants.ASSISTANT_TOKEN)[-1].strip()
|
||||
return text
|
||||
|
||||
# TODO - get from config?
|
||||
Reference in New Issue
Block a user