diff options
Diffstat (limited to 'tests/test_prompts.py')
| -rw-r--r-- | tests/test_prompts.py | 73 |
1 files changed, 73 insertions, 0 deletions
diff --git a/tests/test_prompts.py b/tests/test_prompts.py index a2bf892..4f6568b 100644 --- a/tests/test_prompts.py +++ b/tests/test_prompts.py @@ -5,6 +5,20 @@ import inspect from gap_pipeline import prompts from gap_pipeline.clients import OpenAIJsonClient +from gap_pipeline.kernel_models import ( + DiffusedProof, + MethodPlan, + ProofDAG, + ReplacementPlan, +) +from gap_pipeline.kernel_prompts import ( + dag_user, + diffusion_user, + judge_user, + method_user, + render_user, + replacement_user, +) EXPECTED = { @@ -38,3 +52,62 @@ def test_prompt_values_are_byte_locked() -> None: def test_o3_adapter_does_not_send_temperature() -> None: source = inspect.getsource(OpenAIJsonClient.generate_json) assert '"temperature"' not in source + + +def test_literal_five_stage_prompts_render(item) -> None: + dag = ProofDAG.model_validate( + { + "nodes": [{"node_id": "n1", "claim": "claim", "dependencies": []}], + "terminal_node_id": "n1", + } + ) + methods = MethodPlan.model_validate( + {"nodes": [{"node_id": "n1", "method_label": "method"}]} + ) + replacements = ReplacementPlan.model_validate( + { + "changes": [ + { + "slot_id": "s1", + "source_node_id": "n1", + "description": "constant", + "original_value": "1", + "replacement_value": "2", + "guard_condition": "positive", + "guard_justification": "2 is positive", + } + ], + "closure_statement": "No undeclared changes.", + } + ) + diffused = DiffusedProof.model_validate( + { + "nodes": [ + { + "node_id": "n1", + "dependencies": [], + "method_label": "method", + "instantiated_claim": "new claim", + "justification": "valid", + } + ], + "terminal_node_id": "n1", + "terminal_answer": "answer", + } + ) + variant = { + "question": "question", + "solution": "[n1] solution", + "node_order": ["n1"], + "terminal_answer": "answer", + } + + rendered = [ + dag_user(item), + method_user(dag), + replacement_user(item, dag, methods), + diffusion_user(item, dag, methods, replacements), + render_user(replacements, diffused), + judge_user(item, dag, methods, replacements, diffused, variant), + ] + assert all("{" in value and "}" in value for value in rendered) |
