naming updates; fix static analysis script

This commit is contained in:
Adam Wilson
2025-07-05 13:01:28 -06:00
parent a9db321597
commit 640c261b26
12 changed files with 77 additions and 44 deletions
@@ -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,7 +1,7 @@
import abc
class AbstractLanguageModelResponseService(abc.ABC):
class AbstractTextGenerationCompletionService(abc.ABC):
@abc.abstractmethod
def invoke(self, user_prompt: str) -> str:
raise NotImplementedError
@@ -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,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,
@@ -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?