88a1dab810
---ci--- phase: 7 milestone: v0.2 status: complete requirements: covered: [REQ-2-001, REQ-2-002, REQ-2-003, REQ-2-004, REQ-2-005, REQ-2-006, REQ-2-007, REQ-2-008, REQ-2-009, REQ-2-010, REQ-2-011, REQ-2-012] partial: [] ---/ci--- Milestone v0.2 (ai-tutor-architecture) merged to main. Escalation record (audit remediation, durable): P1 executor delegation failed twice (empty subagent results, zero files created); auto-resolved at full autonomy to inline execution with identical plan fidelity (commit 3271373, reflog-only after phase branch squash-delete).
157 lines
5.4 KiB
Python
157 lines
5.4 KiB
Python
"""Structured output defense tests — 4 layers (D-020), against mock providers."""
|
|
|
|
import pytest
|
|
from pydantic import BaseModel
|
|
|
|
from ai_service.agents.structured import (
|
|
StructuredOutputError,
|
|
extract_json_object,
|
|
parse_structured,
|
|
structured_completion,
|
|
)
|
|
from ai_service.llm.mock import MockProvider, ScriptedJSONProvider
|
|
from ai_service.llm.types import Message
|
|
|
|
|
|
class Score(BaseModel):
|
|
score: int
|
|
verdict: str
|
|
|
|
|
|
HINT = '{"score": <int 0-100>, "verdict": "<short verdict>"}'
|
|
|
|
|
|
def test_extract_json_plain():
|
|
assert extract_json_object('{"a": 1}') == '{"a": 1}'
|
|
|
|
|
|
def test_extract_json_fenced():
|
|
text = '```json\n{"a": 1}\n```'
|
|
assert extract_json_object(text) == '{"a": 1}'
|
|
|
|
|
|
def test_extract_json_with_prose_around():
|
|
text = 'Sure! Here is my answer: {"a": {"b": "x } y"}, "c": 2} hope that helps'
|
|
assert extract_json_object(text) == '{"a": {"b": "x } y"}, "c": 2}'
|
|
|
|
|
|
def test_extract_json_no_object_raises():
|
|
with pytest.raises(StructuredOutputError):
|
|
extract_json_object("no json here")
|
|
|
|
|
|
def test_extract_json_unbalanced_raises():
|
|
with pytest.raises(StructuredOutputError):
|
|
extract_json_object('{"a": 1')
|
|
|
|
|
|
def test_parse_structured_valid():
|
|
result = parse_structured('{"score": 88, "verdict": "solid"}', Score)
|
|
assert result.score == 88
|
|
|
|
|
|
def test_parse_structured_invalid_schema_raises():
|
|
with pytest.raises(StructuredOutputError):
|
|
parse_structured('{"wrong": "shape"}', Score)
|
|
|
|
|
|
async def test_structured_completion_happy_path():
|
|
provider = ScriptedJSONProvider({"score": 91, "verdict": "excellent work"})
|
|
messages = [Message(role="user", content="grade my artifact")]
|
|
result = await structured_completion(
|
|
provider, messages, model="m", schema=Score, schema_hint=HINT
|
|
)
|
|
assert result.score == 91
|
|
assert result.verdict == "excellent work"
|
|
|
|
|
|
async def test_structured_completion_retries_then_raises():
|
|
# Plain MockProvider returns non-schema JSON for json_object requests →
|
|
# both attempts fail validation → StructuredOutputError after ONE retry.
|
|
provider = MockProvider()
|
|
provider.received_calls = []
|
|
messages = [Message(role="user", content="grade me")]
|
|
with pytest.raises(StructuredOutputError):
|
|
await structured_completion(
|
|
provider, messages, model="m", schema=Score, schema_hint=HINT
|
|
)
|
|
|
|
|
|
async def test_structured_completion_retry_succeeds_after_invalid_first_response():
|
|
# Layer 4 recovery: first reply is wrong-schema fenced JSON, retry is valid.
|
|
class FlakyProvider(MockProvider):
|
|
def __init__(self) -> None:
|
|
super().__init__()
|
|
self.n = 0
|
|
self.retry_request: list[Message] = []
|
|
|
|
async def chat(self, messages, *, model, temperature=0.7, response_format=None):
|
|
self.n += 1
|
|
if self.n == 1:
|
|
return '```json\n{"summary": "wrong shape"}\n```'
|
|
self.retry_request = list(messages)
|
|
return '{"score": 75, "verdict": "recovered"}'
|
|
|
|
provider = FlakyProvider()
|
|
messages = [Message(role="user", content="grade me")]
|
|
result = await structured_completion(
|
|
provider, messages, model="m", schema=Score, schema_hint=HINT
|
|
)
|
|
assert result.score == 75
|
|
assert result.verdict == "recovered"
|
|
assert provider.n == 2
|
|
# The retry must feed the validation error back to the model.
|
|
retry_contents = " ".join(m.content for m in provider.retry_request)
|
|
assert "previous response was invalid" in retry_contents
|
|
assert HINT in retry_contents
|
|
|
|
|
|
async def test_structured_completion_is_bounded_to_one_retry():
|
|
# Permanently-invalid provider: exactly two provider calls, then raise.
|
|
class CountingProvider(MockProvider):
|
|
def __init__(self) -> None:
|
|
super().__init__()
|
|
self.calls = 0
|
|
|
|
async def chat(self, messages, *, model, temperature=0.7, response_format=None):
|
|
self.calls += 1
|
|
return await super().chat(
|
|
messages, model=model, response_format=response_format
|
|
)
|
|
|
|
provider = CountingProvider()
|
|
messages = [Message(role="user", content="grade me")]
|
|
with pytest.raises(StructuredOutputError, match="after retry"):
|
|
await structured_completion(
|
|
provider, messages, model="m", schema=Score, schema_hint=HINT
|
|
)
|
|
assert provider.calls == 2
|
|
|
|
|
|
async def test_structured_completion_sends_schema_instruction():
|
|
"""Layer 2: the schema hint must reach the provider in the request."""
|
|
provider = ScriptedJSONProvider({"score": 70, "verdict": "passing"})
|
|
captured: list[list] = []
|
|
|
|
original = provider.chat
|
|
|
|
async def recording_chat(messages, *, model, temperature=0.7, response_format=None):
|
|
captured.append(list(messages))
|
|
return await original(
|
|
messages, model=model, temperature=temperature, response_format=response_format
|
|
)
|
|
|
|
provider.chat = recording_chat
|
|
messages = [Message(role="user", content="grade")]
|
|
await structured_completion(provider, messages, model="m", schema=Score, schema_hint=HINT)
|
|
assert captured, "provider was never called"
|
|
last_user = next(m for m in reversed(captured[0]) if m.role == "user")
|
|
assert HINT in last_user.content
|
|
assert "ONLY" in last_user.content # JSON-only instruction present
|
|
|
|
# Determinism: same request yields same reply
|
|
again = await structured_completion(
|
|
provider, messages, model="m", schema=Score, schema_hint=HINT
|
|
)
|
|
assert again.score == 70
|