more test cases and notes

This commit is contained in:
Adam Wilson
2025-07-06 17:47:49 -06:00
parent ffa2d73ae0
commit 7d49d160b9
7 changed files with 73 additions and 27 deletions
@@ -0,0 +1,4 @@
import abc
class AbstractGuardrailAnalyzedResponse(abc.ABC):
pass
@@ -1,11 +1,15 @@
class GuardrailAnalyzedResponse:
def __init__(
self,
score: float,
cosine_similarity_risk_threshold: float,
original: str,
final: str):
self.score = score
self.is_malicious = score >= cosine_similarity_risk_threshold
self.original = original
self.final = final
from src.text_generation.domain.abstract_guardrail_analyzed_response import AbstractGuardrailAnalyzedResponse
class GuardrailAnalyzedResponse(AbstractGuardrailAnalyzedResponse):
def __init__(
self,
score: float,
cosine_similarity_risk_threshold: float,
original: str):
is_completion_malicious = score >= cosine_similarity_risk_threshold
self.score = score
self.original = original
self.is_completion_malicious = is_completion_malicious
self.final = "I can't answer that." if is_completion_malicious else original
@@ -69,7 +69,7 @@ class HttpApiController:
return [response_body]
response_text = self.text_generation_response_service.invoke(user_prompt=prompt)
score = self.generated_text_guardrail_service.is_text_malicious(response_text)
score = self.generated_text_guardrail_service.process_generated_text(response_text)
response_body = self.format_response(response_text)
http_status_code = 200 # make enum
@@ -96,7 +96,7 @@ class HttpApiController:
return [response_body]
response_text = self.rag_response_service.invoke(user_prompt=prompt)
score = self.generated_text_guardrail_service.is_text_malicious(response_text)
score = self.generated_text_guardrail_service.process_generated_text(response_text)
response_body = self.format_response(response_text)
http_status_code = 200 # make enum
@@ -1,7 +1,8 @@
import abc
from src.text_generation.domain.abstract_guardrail_analyzed_response import AbstractGuardrailAnalyzedResponse
class AbstractGeneratedTextGuardrailService(abc.ABC):
@abc.abstractmethod
def is_text_malicious(self, model_generated_text: str) -> float:
def process_generated_text(self, model_generated_text: str) -> AbstractGuardrailAnalyzedResponse:
raise NotImplementedError
@@ -1,3 +1,4 @@
from src.text_generation.domain.abstract_guardrail_analyzed_response import AbstractGuardrailAnalyzedResponse
from src.text_generation.domain.guardrail_analyzed_response import GuardrailAnalyzedResponse
from src.text_generation.services.guardrails.abstract_generated_text_guardrail_service import AbstractGeneratedTextGuardrailService
from src.text_generation.services.nlp.abstract_semantic_similarity_service import AbstractSemanticSimilarityService
@@ -13,12 +14,10 @@ class GeneratedTextGuardrailService(AbstractGeneratedTextGuardrailService):
self.semantic_similarity_service.use_comparison_texts(comparison_texts)
self.cosine_similarity_risk_threshold: float = 0.5
def is_text_malicious(self, model_generated_text: str) -> GuardrailAnalyzedResponse:
def process_generated_text(self, model_generated_text: str) -> AbstractGuardrailAnalyzedResponse:
score: float = self.semantic_similarity_service.analyze(text=model_generated_text)
response = GuardrailAnalyzedResponse(
score=score,
cosine_similarity_risk_threshold=self.cosine_similarity_risk_threshold,
original=model_generated_text,
final="test")
return response
original=model_generated_text)
return response