fix(tests runtime):

This commit is contained in:
Alexander Myasoedov committed 2025-12-09 20:00:04 +02:00
1 parent b9dc5de708
commit d56b406e1a
5 files changed
+59 -17

No files matched your search

@@ -206,7 +206,11 @@ class QLearningPromptSelector(PromptSelectionInterface):
class Module: class Module:
def __init__( def __init__(
self, prompt_groups: list[str], tools_inbox: asyncio.Queue, opts: dict = {} self,
prompt_groups: list[str],
tools_inbox: asyncio.Queue,
opts: dict = {},
rl_model: PromptSelectionInterface | None = None,
): ):
self.tools_inbox = tools_inbox self.tools_inbox = tools_inbox
self.opts = opts self.opts = opts
@@ -214,7 +218,7 @@ class Module:
self.max_prompts = self.opts.get("max_prompts", 10) # Default max M prompts self.max_prompts = self.opts.get("max_prompts", 10) # Default max M prompts
self.run_id = U.uuid4().hex self.run_id = U.uuid4().hex
self.batch_size = self.opts.get("batch_size", 500) self.batch_size = self.opts.get("batch_size", 500)
self.rl_model = CloudRLPromptSelector( self.rl_model = rl_model or CloudRLPromptSelector(
prompt_groups, "https://mcp.metaheuristic.co", run_id=self.run_id prompt_groups, "https://mcp.metaheuristic.co", run_id=self.run_id
) )
@@ -33,11 +33,17 @@ def mock_requests() -> Mock:
@pytest.fixture @pytest.fixture
def mock_rl_selector() -> Mock: def mock_rl_selector(dataset_prompts) -> Mock:
return CloudRLPromptSelector( class StubSelector:
dataset_prompts, def __init__(self, prompts: list[str]):
api_url="https://mcp.metaheuristic.co", self.prompts = prompts
) self.idx = 0
def select_next_prompts(self, current_prompt: str, passed_guard: bool) -> list[str]:
self.idx = (self.idx + 1) % len(self.prompts)
return [self.prompts[self.idx]]
return StubSelector(dataset_prompts)
@pytest.fixture @pytest.fixture
@@ -91,7 +97,10 @@ class TestCloudRLPromptSelector:
next_prompt = selector.select_next_prompt("What is AI?", passed_guard=True) next_prompt = selector.select_next_prompt("What is AI?", passed_guard=True)
assert next_prompt in dataset_prompts assert next_prompt in dataset_prompts
def test_select_next_prompt_success_service(self, dataset_prompts): def test_select_next_prompt_success_service(self, dataset_prompts, mock_requests):
mock_requests.return_value.status_code = 200
mock_requests.return_value.json.return_value = {"next_prompts": ["What is AI?"]}
selector = CloudRLPromptSelector( selector = CloudRLPromptSelector(
dataset_prompts, dataset_prompts,
api_url="https://mcp.metaheuristic.co", api_url="https://mcp.metaheuristic.co",
@@ -99,7 +108,7 @@ class TestCloudRLPromptSelector:
next_prompt = selector.select_next_prompt( next_prompt = selector.select_next_prompt(
"How does RL work?", passed_guard=True "How does RL work?", passed_guard=True
) )
assert next_prompt assert next_prompt == "What is AI?"
# Tests for QLearningPromptSelector # Tests for QLearningPromptSelector
@@ -188,7 +197,7 @@ class TestModule:
async def test_apply_basic_flow( async def test_apply_basic_flow(
self, dataset_prompts, tools_inbox, mock_rl_selector self, dataset_prompts, tools_inbox, mock_rl_selector
): ):
module = Module(dataset_prompts, tools_inbox) module = Module(dataset_prompts, tools_inbox, rl_model=mock_rl_selector)
count = 0 count = 0
async for prompt in module.apply(): async for prompt in module.apply():
@@ -198,7 +207,9 @@ class TestModule:
break break
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_apply_rl_with_tools_inbox(self, dataset_prompts, tools_inbox): async def test_apply_rl_with_tools_inbox(
self, dataset_prompts, tools_inbox, mock_rl_selector
):
# Add a test message to the tools inbox # Add a test message to the tools inbox
test_message = { test_message = {
"message": "Test message", "message": "Test message",
@@ -207,7 +218,7 @@ class TestModule:
} }
await tools_inbox.put(test_message) await tools_inbox.put(test_message)
module = Module(dataset_prompts, tools_inbox) module = Module(dataset_prompts, tools_inbox, rl_model=mock_rl_selector)
async for output in module.apply(): async for output in module.apply():
if output == "Test message": if output == "Test message":
+8 -1
View File
@@ -76,14 +76,21 @@ async def test_perform_single_shot_scan_success(prepare_prompts_mock):
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("agentic_security.probe_data.msj_data.prepare_prompts")
@patch("agentic_security.probe_data.data.prepare_prompts") @patch("agentic_security.probe_data.data.prepare_prompts")
async def test_perform_many_shot_scan_probe_injection(prepare_prompts_mock): async def test_perform_many_shot_scan_probe_injection(
prepare_prompts_mock, msj_prepare_prompts_mock
):
# Mock main and probe prompt modules # Mock main and probe prompt modules
prepare_prompts_mock.side_effect = [ prepare_prompts_mock.side_effect = [
[MagicMock(dataset_name="main_module", prompts=["main_prompt1"], lazy=False)], [MagicMock(dataset_name="main_module", prompts=["main_prompt1"], lazy=False)],
[MagicMock(dataset_name="probe_module", prompts=["probe_prompt1"], lazy=False)], [MagicMock(dataset_name="probe_module", prompts=["probe_prompt1"], lazy=False)],
] ]
msj_prepare_prompts_mock.return_value = [
MagicMock(dataset_name="msj_probe_module", prompts=["msj_probe_prompt"], lazy=False)
]
# Mock request_factory # Mock request_factory
mock_response = AsyncMock() mock_response = AsyncMock()
mock_response.fn.side_effect = [ mock_response.fn.side_effect = [
+3 -1
View File
@@ -1,5 +1,6 @@
import base64 import base64
import io import io
import random
import httpx import httpx
import pytest import pytest
@@ -85,8 +86,9 @@ def test_data_config_endpoint():
def test_refusal_rate(): def test_refusal_rate():
"""Test that refusal rate is approximately 20%""" """Test that refusal rate is approximately 20%"""
random.seed(0)
refusal_count = 0 refusal_count = 0
total_trials = 1000 total_trials = 200
for _ in range(total_trials): for _ in range(total_trials):
response = client.post("/v1/self-probe", json={"prompt": "test"}) response = client.post("/v1/self-probe", json={"prompt": "test"})
+21 -3
View File
@@ -1,6 +1,7 @@
import importlib import importlib
import os import os
import signal import signal
import socket
import subprocess import subprocess
import tempfile import tempfile
import time import time
@@ -24,12 +25,29 @@ def test_server(request):
preexec_fn=lambda: signal.signal(signal.SIGINT, signal.SIG_IGN), preexec_fn=lambda: signal.signal(signal.SIGINT, signal.SIG_IGN),
) )
# Give the server time to start def wait_for_port(host: str, port: int, timeout: float = 5.0) -> bool:
time.sleep(2) start = time.time()
while time.time() - start < timeout:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.settimeout(0.2)
try:
sock.connect((host, port))
return True
except OSError:
time.sleep(0.1)
return False
if not wait_for_port("127.0.0.1", 9094):
server.kill()
pytest.skip("Test server failed to start within timeout")
def cleanup(): def cleanup():
server.terminate() server.terminate()
server.wait() try:
server.wait(timeout=3)
except subprocess.TimeoutExpired:
server.kill()
server.wait(timeout=2)
request.addfinalizer(cleanup) request.addfinalizer(cleanup)
return server return server