summaryrefslogtreecommitdiff
path: root/tests/test_prompts.py
diff options
context:
space:
mode:
authorOscar Wan <oscarwan@stanford.edu>2026-07-24 20:45:42 -0700
committerOscar Wan <oscarwan@stanford.edu>2026-07-24 20:45:42 -0700
commit15efc30e9e7179accd30375d3edb2e34a3b4dc5f (patch)
treeba1c49eb128e7906ba451141723d763e1dcda53a /tests/test_prompts.py
parent708f2af9c6985e9cb5cd53e434a7d3b8dfa2b4ac (diff)
updated generation process
Diffstat (limited to 'tests/test_prompts.py')
-rw-r--r--tests/test_prompts.py73
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)