mirror of
https://github.com/msoedov/agentic_security.git
synced 2026-09-30 19:59:34 +02:00
fix(tests runtime):
This commit is contained in:
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":
|
||||||
|
|||||||
@@ -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 = [
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
Reference in new issue
Block a user