matrix strategy for tests

This commit is contained in:
Adam Wilson
2025-08-18 16:12:31 -06:00
parent b1a2672d14
commit 010933aa59
6 changed files with 40 additions and 22 deletions
@@ -1,5 +1,7 @@
import pytest
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
from tests.integration.test_utils import get_prompt_batch, run_prompt_analysis_test
@pytest.mark.mitigation
@@ -18,7 +20,7 @@ def test_03_malicious_prompts_rag_microsoft_phi_3_mini4k_instruct(
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=get_prompt_batch(malicious_prompts),
comparison_texts=prompt_injection_all_example_responses,
service_configurator=configure_service
)
@@ -1,6 +1,6 @@
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
from tests.integration.test_utils import get_prompt_batch, run_prompt_analysis_test
import pytest
@@ -23,7 +23,7 @@ def test_04_malicious_prompts_rag_and_cot_apple_openelm_3b_instruct(
model_id=ModelId.APPLE_OPENELM_3B_INSTRUCT,
text_generation_completion_service=text_generation_completion_service,
semantic_similarity_service=semantic_similarity_service,
prompts=malicious_prompts[:1],
prompts=get_prompt_batch(malicious_prompts),
comparison_texts=prompt_injection_all_example_responses,
service_configurator=configure_service
)
@@ -1,6 +1,6 @@
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
from tests.integration.test_utils import run_prompt_analysis_test, get_prompt_batch
import pytest
@@ -23,7 +23,7 @@ def test_04_malicious_prompts_rag_and_cot_meta_llama_3_2_3b_instruct(
model_id=ModelId.META_LLAMA_3_2_3B_INSTRUCT,
text_generation_completion_service=text_generation_completion_service,
semantic_similarity_service=semantic_similarity_service,
prompts=malicious_prompts[:1],
prompts=get_prompt_batch(malicious_prompts),
comparison_texts=prompt_injection_all_example_responses,
service_configurator=configure_service
)
@@ -1,9 +1,9 @@
import os
import pytest
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
import pytest
from tests.integration.test_utils import run_prompt_analysis_test, get_prompt_batch
@pytest.mark.mitigation
@@ -17,13 +17,13 @@ def test_04_malicious_prompts_rag_and_cot_microsoft_phi_3_mini4k_instruct(
return (service
.with_chain_of_thought_guidelines()
.with_rag_context_guidelines())
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[:1],
prompts=get_prompt_batch(malicious_prompts),
comparison_texts=prompt_injection_all_example_responses,
service_configurator=configure_service
)
+15
View File
@@ -1,4 +1,5 @@
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
@@ -9,6 +10,20 @@ 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, batch_size=10, env_var='PROMPT_BATCH'):
"""
Returns a batch of prompts based on the PROMPT_BATCH environment variable.
Prints batch info for debugging.
"""
batch_num = int(os.getenv(env_var, '1'))
start_idx = (batch_num - 1) * batch_size
end_idx = min(start_idx + batch_size, len(prompts))
prompt_subset = prompts[start_idx:end_idx]
print(f"Running batch {batch_num}: prompts {start_idx+1}-{end_idx} ({len(prompt_subset)} prompts)")
return prompt_subset
def run_prompt_analysis_test(
test_id: int,
model_id: ModelId,