mirror of
https://github.com/lightbroker/llmsecops-research.git
synced 2026-08-12 22:30:24 +02:00
support batch tests
This commit is contained in:
@@ -1,15 +1,13 @@
|
||||
"""
|
||||
Usage:
|
||||
$ pytest -k test_04_malicious_prompts_rag_and_cot -s --disable-warnings
|
||||
"""
|
||||
|
||||
import os
|
||||
import pytest
|
||||
from test_utils import run_prompt_analysis_test
|
||||
|
||||
from src.text_generation.common.model_id import ModelId
|
||||
from src.text_generation.services.nlp.abstract_text_generation_completion_service import AbstractTextGenerationCompletionService
|
||||
from tests.integration.test_utils import run_prompt_analysis_test, get_prompt_batch
|
||||
|
||||
|
||||
@pytest.mark.mitigation
|
||||
def test_04_malicious_prompts_rag_and_cot(
|
||||
def test_04_malicious_prompts_rag_and_cot_microsoft_phi_3_mini4k_instruct(
|
||||
text_generation_completion_service: AbstractTextGenerationCompletionService,
|
||||
semantic_similarity_service,
|
||||
malicious_prompts,
|
||||
@@ -17,14 +15,19 @@ def test_04_malicious_prompts_rag_and_cot(
|
||||
):
|
||||
def configure_service(service):
|
||||
return (service
|
||||
.with_chain_of_thought_guidelines()
|
||||
.with_rag_context_guidelines())
|
||||
.with_chain_of_thought_guidelines()
|
||||
.with_rag_context_guidelines())
|
||||
|
||||
prompt_batch, start, end = get_prompt_batch(malicious_prompts)
|
||||
|
||||
results = run_prompt_analysis_test(
|
||||
test_id=4,
|
||||
model_id=ModelId.MICROSOFT_PHI_3_MINI4K_INSTRUCT,
|
||||
text_generation_completion_service=text_generation_completion_service,
|
||||
semantic_similarity_service=semantic_similarity_service,
|
||||
prompts=malicious_prompts,
|
||||
prompts=prompt_batch,
|
||||
start=start,
|
||||
end=end,
|
||||
comparison_texts=prompt_injection_all_example_responses,
|
||||
service_configurator=configure_service
|
||||
)
|
||||
@@ -1,5 +1,7 @@
|
||||
import inspect
|
||||
import os
|
||||
from typing import List, Callable
|
||||
from src.text_generation.common.model_id import ModelId
|
||||
from src.text_generation.domain.semantic_similarity_result import SemanticSimilarityResult
|
||||
from src.text_generation.domain.text_generation_completion_result import TextGenerationCompletionResult
|
||||
from src.text_generation.services.logging.test_run_logging_service import TestRunLoggingService
|
||||
@@ -8,11 +10,40 @@ from src.text_generation.services.nlp.abstract_text_generation_completion_servic
|
||||
from src.text_generation.services.nlp.text_generation_completion_service import TextGenerationCompletionService
|
||||
|
||||
|
||||
|
||||
def get_prompt_batch(prompts: List[str], batch_size=10, env_var='PROMPT_BATCH'):
|
||||
|
||||
batch_size = int(os.getenv('BATCH_SIZE', '2'))
|
||||
batch_num = int(os.getenv('PROMPT_BATCH', '1'))
|
||||
|
||||
if 'BATCH_OFFSET' in os.environ:
|
||||
# Option 1: Fixed offset per workflow
|
||||
offset = int(os.getenv('BATCH_OFFSET', '0'))
|
||||
else:
|
||||
# Option 2: Configurable range
|
||||
prompt_range = int(os.getenv('PROMPT_RANGE', '1'))
|
||||
offset = (prompt_range - 1) * 20
|
||||
|
||||
# Calculate start and end indices
|
||||
start_idx = offset + (batch_num - 1) * batch_size
|
||||
end_idx = min(start_idx + batch_size, len(prompts))
|
||||
|
||||
# Get the subset of prompts for this batch
|
||||
prompt_subset: List[str] = prompts[start_idx:end_idx]
|
||||
|
||||
print(f"Running batch {batch_num} (range offset {offset}): prompts {start_idx+1}-{end_idx} ({len(prompt_subset)} prompts)")
|
||||
|
||||
return prompt_subset, (start_idx+1), end_idx
|
||||
|
||||
|
||||
def run_prompt_analysis_test(
|
||||
test_id: int,
|
||||
model_id: ModelId,
|
||||
text_generation_completion_service: AbstractTextGenerationCompletionService,
|
||||
semantic_similarity_service: AbstractSemanticSimilarityService,
|
||||
prompts: List,
|
||||
start: int,
|
||||
end: int,
|
||||
comparison_texts: List,
|
||||
service_configurator: Callable,
|
||||
max_prompts: int = 100
|
||||
@@ -42,7 +73,7 @@ def run_prompt_analysis_test(
|
||||
configured_service: TextGenerationCompletionService = service_configurator(text_generation_completion_service)
|
||||
|
||||
print(f'sending prompt {i} to LLM')
|
||||
completion_result: TextGenerationCompletionResult = configured_service.invoke(user_prompt=prompt)
|
||||
completion_result: TextGenerationCompletionResult = configured_service.invoke(user_prompt=prompt, model_id=model_id)
|
||||
|
||||
print(f'analyzing result')
|
||||
completion_result.finalize_completion_text()
|
||||
@@ -51,13 +82,18 @@ def run_prompt_analysis_test(
|
||||
print(f'{i}/{len(prompts)} Max Score: {result.max}')
|
||||
print(f'{i}/{len(prompts)} Avg Score: {result.mean}')
|
||||
|
||||
TestRunLoggingService(test_id=test_id).log_results(
|
||||
TestRunLoggingService(
|
||||
test_id=test_id,
|
||||
model_id=model_id,
|
||||
start=start,
|
||||
end=end
|
||||
).log_results(
|
||||
id=inspect.currentframe().f_back.f_code.co_name,
|
||||
text_generation_completion_result=completion_result,
|
||||
final_completion_text_score=result.max,
|
||||
is_rag_few_shot_enabled=text_generation_completion_service.is_rag_context_enabled(),
|
||||
is_cot_enabled=text_generation_completion_service.is_chain_of_thought_enabled(),
|
||||
is_reflexion_enabled=text_generation_completion_service.is_reflexion_enabled(),
|
||||
is_reflexion_enabled=False,
|
||||
original_llm_config=completion_result.original_result.llm_config
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user