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