Files
baozaotumao2025 14722be770 first commit
2026-07-18 21:00:26 +08:00

491 lines
17 KiB
Python

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