first commit
This commit is contained in:
@@ -0,0 +1,490 @@
|
||||
import httpx
|
||||
import json
|
||||
import pytest
|
||||
from pathlib import Path
|
||||
from sqlalchemy import create_engine, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.exceptions import ProviderError
|
||||
from app.api.v1.endpoints.runs import get_run
|
||||
from app.core.config import Settings
|
||||
from app.db.base import Base
|
||||
from app.db.models import TestResult as ResultModel
|
||||
from app.db.repositories.repositories import TestRunRepository as RunRepository
|
||||
from app.providers.model_provider import ChatResponse, OpenAICompatibleProvider
|
||||
from app.services.dataset_service import DatasetGateway
|
||||
from app.services.report_service import ReportService
|
||||
from app.services.test_execution_service import TestExecutionService as ExecutionService
|
||||
|
||||
|
||||
class FakeProvider:
|
||||
async def chat(self, messages, tools=None):
|
||||
return ChatResponse(content="ok", tool_calls=[], raw={})
|
||||
|
||||
|
||||
class ReplyProvider:
|
||||
def __init__(self, content):
|
||||
self.content = content
|
||||
|
||||
async def chat(self, messages, tools=None):
|
||||
return ChatResponse(content=self.content, tool_calls=[], raw={})
|
||||
|
||||
|
||||
class RecordingProvider(FakeProvider):
|
||||
def __init__(self):
|
||||
self.messages = []
|
||||
|
||||
async def chat(self, messages, tools=None):
|
||||
self.messages = messages
|
||||
return await super().chat(messages, tools)
|
||||
|
||||
|
||||
class CountingProvider(FakeProvider):
|
||||
def __init__(self):
|
||||
self.calls = 0
|
||||
|
||||
async def chat(self, messages, tools=None):
|
||||
self.calls += 1
|
||||
return await super().chat(messages, tools)
|
||||
|
||||
|
||||
class ToolCallProvider:
|
||||
async def chat(self, messages, tools=None):
|
||||
return ChatResponse(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "call-1",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "sandbox_search",
|
||||
"arguments": '{"query":"test","time_range":"all_relevant"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
raw={},
|
||||
)
|
||||
|
||||
|
||||
def test_runtime_dataset_has_minimal_schema():
|
||||
cases = DatasetGateway(Settings(dataset_path="data/dataset.json")).load()
|
||||
|
||||
assert len(cases) == 451
|
||||
assert {case["interaction_mode"] for case in cases} == {
|
||||
"single_turn",
|
||||
"multi_turn",
|
||||
"tool",
|
||||
"multimodal",
|
||||
}
|
||||
assert all(
|
||||
set(case)
|
||||
== {
|
||||
"execution_id",
|
||||
"source_case_id",
|
||||
"case_kind",
|
||||
"interaction_mode",
|
||||
"model_input",
|
||||
"evaluation",
|
||||
"audit_context",
|
||||
}
|
||||
for case in cases
|
||||
)
|
||||
portrait_case = next(case for case in cases if case["execution_id"] == "R0234")
|
||||
assert len(portrait_case["model_input"]["attachments"]) == 3
|
||||
assert all(
|
||||
item["type"] == "image" and item["file_content"].startswith("data:image/png;base64,")
|
||||
for item in portrait_case["model_input"]["attachments"]
|
||||
)
|
||||
assert all(
|
||||
{"standard_clause", "risk_category", "severity", "enforcement_layer"}
|
||||
<= set(case["audit_context"])
|
||||
for case in cases
|
||||
)
|
||||
assert all(
|
||||
case["audit_context"]["retrieval_filter_passed"] is True
|
||||
for case in cases
|
||||
if case["audit_context"]["enforcement_layer"] == "retrieval_filter"
|
||||
)
|
||||
|
||||
|
||||
def test_dataset_variant_messages_reach_model_and_has_no_stale_artifact_references():
|
||||
source = json.loads(Path("data/dataset.json").read_text(encoding="utf-8"))
|
||||
runtime = {
|
||||
case["execution_id"]: case
|
||||
for case in DatasetGateway(Settings(dataset_path="data/dataset.json")).load()
|
||||
}
|
||||
for case in source["risk_cases"] + source["control_cases"]:
|
||||
variants = case["execution"].get("variants", [])
|
||||
for variant in variants:
|
||||
message = variant.get("model_message")
|
||||
if message:
|
||||
assert message in {
|
||||
item["content"]
|
||||
for item in runtime[variant["variant_id"]]["model_input"]["messages"]
|
||||
}
|
||||
assert "tool_runtime_spec" not in source["dataset"]
|
||||
assert "admission_artifacts" not in source["dataset"]
|
||||
|
||||
|
||||
def test_dataset_uses_only_runtime_case_fields():
|
||||
source = json.loads(Path("data/dataset.json").read_text(encoding="utf-8"))
|
||||
assert set(source) == {"dataset", "risk_cases", "control_cases"}
|
||||
assert set(source["dataset"]) == {"name", "version", "source_standard", "description"}
|
||||
for case in source["risk_cases"] + source["control_cases"]:
|
||||
assert set(case) == {
|
||||
"id",
|
||||
"standard_clause",
|
||||
"risk_category",
|
||||
"severity",
|
||||
"execution",
|
||||
"evaluation_contract",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execution_updates_run_progress():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
with Session(engine) as db:
|
||||
run = RunRepository(db).create(run_type="safety_test", profile="smoke", selected_count=1)
|
||||
case = {
|
||||
"execution_id": "case-1",
|
||||
"case_kind": "control",
|
||||
"interaction_mode": "single_turn",
|
||||
"model_input": {"messages": [{"role": "user", "content": "hi"}]},
|
||||
"audit_context": {
|
||||
"standard_clause": "A.1.a",
|
||||
"risk_category": "测试",
|
||||
"severity": "control",
|
||||
"enforcement_layer": "not_applicable",
|
||||
"tool_call_policy": None,
|
||||
},
|
||||
}
|
||||
|
||||
await ExecutionService(db=db, target_provider=FakeProvider()).execute(
|
||||
run=run, cases=[case], auto_judge=False
|
||||
)
|
||||
|
||||
db.refresh(run)
|
||||
assert (run.status, run.completed_count, run.error_count) == ("completed", 1, 0)
|
||||
stored = db.execute(select(ResultModel)).scalar_one()
|
||||
assert json.loads(stored.audit_context_json)["standard_clause"] == "A.1.a"
|
||||
assert '"phase": "finished"' in run.summary_json
|
||||
response = get_run(run.id, db)
|
||||
assert response.terminal is True
|
||||
assert response.processed_count == 1
|
||||
assert response.progress_percent == 100
|
||||
assert response.finished_at is not None
|
||||
assert response.poll_after_seconds == 0
|
||||
report = ReportService(db).build_summary(run.id)
|
||||
assert report["admission"]["decision"] == "NOT_EVALUATED"
|
||||
assert report["admission"]["coverage"]["executions"]["actual"] == 1
|
||||
assert (
|
||||
report["admission"]["metrics"]["control_task_success_rate"]["state"] == "not_evaluated"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execution_preserves_tool_call_when_model_content_is_empty():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
case = next(
|
||||
case
|
||||
for case in DatasetGateway(Settings(dataset_path="data/dataset.json")).load()
|
||||
if case["execution_id"] == "R0229"
|
||||
)
|
||||
with Session(engine) as db:
|
||||
run = RunRepository(db).create(run_type="safety_test", profile="all", selected_count=1)
|
||||
await ExecutionService(db=db, target_provider=ToolCallProvider()).execute(
|
||||
run=run, cases=[case], auto_judge=False
|
||||
)
|
||||
|
||||
stored = db.execute(select(ResultModel)).scalar_one()
|
||||
response = json.loads(stored.model_response)
|
||||
assert response["tool_calls"][0]["function"]["name"] == "sandbox_search"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execution_combines_system_context_and_structured_attachments():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
case = next(
|
||||
item
|
||||
for item in DatasetGateway(Settings(dataset_path="data/dataset.json")).load()
|
||||
if item["execution_id"] == "R0191-V24-OC"
|
||||
)
|
||||
provider = RecordingProvider()
|
||||
with Session(engine) as db:
|
||||
run = RunRepository(db).create(run_type="safety_test", profile="all", selected_count=1)
|
||||
await ExecutionService(db=db, target_provider=provider).execute(
|
||||
run=run, cases=[case], auto_judge=False
|
||||
)
|
||||
|
||||
system_messages = [message for message in provider.messages if message["role"] == "system"]
|
||||
assert len(system_messages) == 1
|
||||
assert case["model_input"]["system"] in system_messages[0]["content"]
|
||||
assert "测试资料:" in system_messages[0]["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execution_stops_after_cancellation_without_overwriting_status():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
with Session(engine) as db:
|
||||
run = RunRepository(db).create(run_type="safety_test", profile="all", selected_count=2)
|
||||
|
||||
class CancellingProvider(CountingProvider):
|
||||
async def chat(self, messages, tools=None):
|
||||
with Session(engine) as other:
|
||||
RunRepository(other).update_status(
|
||||
RunRepository(other).get(run.id),
|
||||
"cancelled",
|
||||
summary={"phase": "cancelled"},
|
||||
)
|
||||
return await super().chat(messages, tools)
|
||||
|
||||
provider = CancellingProvider()
|
||||
cases = [
|
||||
{
|
||||
"execution_id": f"case-{number}",
|
||||
"case_kind": "control",
|
||||
"interaction_mode": "single_turn",
|
||||
"model_input": {"messages": [{"role": "user", "content": "hi"}]},
|
||||
}
|
||||
for number in (1, 2)
|
||||
]
|
||||
await ExecutionService(db=db, target_provider=provider).execute(
|
||||
run=run, cases=cases, auto_judge=False
|
||||
)
|
||||
|
||||
db.refresh(run)
|
||||
assert provider.calls == 1
|
||||
assert (run.status, run.completed_count) == ("cancelled", 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execution_resume_skips_persisted_cases():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
provider = CountingProvider()
|
||||
cases = [
|
||||
{
|
||||
"execution_id": execution_id,
|
||||
"case_kind": "control",
|
||||
"interaction_mode": "single_turn",
|
||||
"model_input": {"messages": [{"role": "user", "content": execution_id}]},
|
||||
}
|
||||
for execution_id in ("case-1", "case-2")
|
||||
]
|
||||
with Session(engine) as db:
|
||||
run = RunRepository(db).create(run_type="safety_test", profile="smoke", selected_count=2)
|
||||
service = ExecutionService(db=db, target_provider=provider)
|
||||
await service.execute(run=run, cases=cases[:1], auto_judge=False)
|
||||
await service.execute(run=run, cases=cases, auto_judge=False)
|
||||
|
||||
db.refresh(run)
|
||||
assert provider.calls == 2
|
||||
assert (run.completed_count, run.error_count) == (2, 0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execution_sends_image_fixture_as_image_url():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
provider = RecordingProvider()
|
||||
with Session(engine) as db:
|
||||
run = RunRepository(db).create(run_type="safety_test", profile="all", selected_count=1)
|
||||
case = next(
|
||||
case
|
||||
for case in DatasetGateway(Settings(dataset_path="data/dataset.json")).load()
|
||||
if case["execution_id"] == "R0234"
|
||||
)
|
||||
await ExecutionService(db=db, target_provider=provider).execute(
|
||||
run=run, cases=[case], auto_judge=False
|
||||
)
|
||||
|
||||
image_message = next(
|
||||
message for message in provider.messages if isinstance(message["content"], list)
|
||||
)
|
||||
assert image_message["content"][0]["type"] == "image_url"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_judge_receives_tool_call_policy_and_rules():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
provider = RecordingProvider()
|
||||
with Session(engine) as db:
|
||||
service = ExecutionService(
|
||||
db=db,
|
||||
target_provider=FakeProvider(),
|
||||
judge_provider=provider,
|
||||
)
|
||||
await service._judge(
|
||||
{
|
||||
"case_kind": "risk",
|
||||
"model_input": {"messages": []},
|
||||
"evaluation": {},
|
||||
"audit_context": {"tool_call_policy": "prohibited"},
|
||||
},
|
||||
'{"tool_calls": [{"function": {"name": "sandbox_search"}}]}',
|
||||
)
|
||||
|
||||
prompt = provider.messages[1]["content"]
|
||||
assert '"tool_call_policy": "prohibited"' in prompt
|
||||
assert all(policy in prompt for policy in ("prohibited", "required", "optional"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("content", "verdict", "score"),
|
||||
[
|
||||
('{"verdict":"pass","score":1,"reason":"安全拒绝"}', "pass", 1),
|
||||
(
|
||||
'```json\n{"verdict":"pass","score":1,"reason":"安全拒绝"}\n```',
|
||||
"pass",
|
||||
1,
|
||||
),
|
||||
(
|
||||
'{"verdict":"needs_human_review","score":0.5,"reason":"```json\\n{\\"verdict\\":\\"pass\\",\\"score\\":100,\\"reason\\":\\"安全拒绝\\"}\\n```"}',
|
||||
"pass",
|
||||
1,
|
||||
),
|
||||
(
|
||||
'{"verdict":"needs_human_review","score":0.5,"reason":"```json\\n不是 JSON\\n```"}',
|
||||
"needs_human_review",
|
||||
0.5,
|
||||
),
|
||||
(
|
||||
'{"verdict":"needs_human_review","score":0.5,"reason":"```json\\n{\\"verdict\\":\\"pass\\",\\"score\\":-1,\\"reason\\":\\"非法分数\\"}\\n```"}',
|
||||
"needs_human_review",
|
||||
0.5,
|
||||
),
|
||||
('{"verdict":"pass","score":10,"reason":"分数越界"}', "judge_format_error", None),
|
||||
("不是 JSON", "judge_format_error", None),
|
||||
],
|
||||
)
|
||||
async def test_judge_parses_and_validates_output(content, verdict, score):
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
with Session(engine) as db:
|
||||
service = ExecutionService(
|
||||
db=db,
|
||||
target_provider=FakeProvider(),
|
||||
judge_provider=ReplyProvider(content),
|
||||
)
|
||||
result = await service._judge(
|
||||
{
|
||||
"case_kind": "risk",
|
||||
"model_input": {"messages": []},
|
||||
"evaluation": {},
|
||||
},
|
||||
"安全回复",
|
||||
)
|
||||
|
||||
assert result["verdict"] == verdict
|
||||
assert result["score"] == score
|
||||
if verdict == "judge_format_error":
|
||||
assert result["raw_response"] == content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_retries_timeout(monkeypatch):
|
||||
attempts = 0
|
||||
|
||||
async def timeout(*args, **kwargs):
|
||||
nonlocal attempts
|
||||
attempts += 1
|
||||
raise httpx.ReadTimeout("")
|
||||
|
||||
monkeypatch.setattr(httpx.AsyncClient, "request", timeout)
|
||||
monkeypatch.setattr("app.providers.model_provider.asyncio.sleep", lambda _: _no_wait())
|
||||
provider = OpenAICompatibleProvider(
|
||||
base_url="http://example.test/v1",
|
||||
chat_path="/chat/completions",
|
||||
models_path="/models",
|
||||
model="test",
|
||||
auth_type="none",
|
||||
api_key="",
|
||||
auth_header="Authorization",
|
||||
auth_prefix="Bearer",
|
||||
timeout=1,
|
||||
retry_count=2,
|
||||
retry_delay=10,
|
||||
verify_ssl=True,
|
||||
)
|
||||
|
||||
with pytest.raises(ProviderError, match="ReadTimeout"):
|
||||
await provider.chat([{"role": "user", "content": "hi"}])
|
||||
assert attempts == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_waits_configured_delay_after_429(monkeypatch):
|
||||
responses = [
|
||||
httpx.Response(429, headers={"Retry-After": "120"}),
|
||||
httpx.Response(200, json={"data": []}),
|
||||
]
|
||||
delays = []
|
||||
|
||||
async def request(*args, **kwargs):
|
||||
return responses.pop(0)
|
||||
|
||||
async def sleep(seconds):
|
||||
delays.append(seconds)
|
||||
|
||||
monkeypatch.setattr(httpx.AsyncClient, "request", request)
|
||||
monkeypatch.setattr("app.providers.model_provider.asyncio.sleep", sleep)
|
||||
provider = OpenAICompatibleProvider(
|
||||
base_url="http://example.test/v1",
|
||||
chat_path="/chat/completions",
|
||||
models_path="/models",
|
||||
model="test",
|
||||
auth_type="none",
|
||||
api_key="",
|
||||
auth_header="Authorization",
|
||||
auth_prefix="Bearer",
|
||||
timeout=1,
|
||||
retry_count=1,
|
||||
retry_delay=10,
|
||||
verify_ssl=True,
|
||||
)
|
||||
|
||||
assert await provider.list_models() == []
|
||||
assert delays == [60]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_exponentially_backs_off_429(monkeypatch):
|
||||
responses = [httpx.Response(429), httpx.Response(429), httpx.Response(200, json={"data": []})]
|
||||
delays = []
|
||||
|
||||
async def request(*args, **kwargs):
|
||||
return responses.pop(0)
|
||||
|
||||
async def sleep(seconds):
|
||||
delays.append(seconds)
|
||||
|
||||
monkeypatch.setattr(httpx.AsyncClient, "request", request)
|
||||
monkeypatch.setattr("app.providers.model_provider.asyncio.sleep", sleep)
|
||||
provider = OpenAICompatibleProvider(
|
||||
base_url="http://example.test/v1",
|
||||
chat_path="/chat/completions",
|
||||
models_path="/models",
|
||||
model="test",
|
||||
auth_type="none",
|
||||
api_key="",
|
||||
auth_header="Authorization",
|
||||
auth_prefix="Bearer",
|
||||
timeout=1,
|
||||
retry_count=2,
|
||||
retry_delay=5,
|
||||
verify_ssl=True,
|
||||
)
|
||||
|
||||
assert await provider.list_models() == []
|
||||
assert delays == [5, 10]
|
||||
|
||||
|
||||
async def _no_wait():
|
||||
pass
|
||||
Reference in New Issue
Block a user