diff --git a/agentic_research/agents/recursive_prover.py b/agentic_research/agents/recursive_prover.py index e4e1e72..6807bef 100644 --- a/agentic_research/agents/recursive_prover.py +++ b/agentic_research/agents/recursive_prover.py @@ -538,11 +538,21 @@ def _prove_parent_with_children( f"Suggested fix: {node.failure_diagnosis.suggested_fix}\n" ) + use_thinking = self._prover_config.parent_extended_thinking + log.info( + "parent_assembly_extended_thinking", + thinking_budget=self._prover_config.thinking_budget, + parent_extended_thinking=use_thinking, + ) + thinking_kwargs: dict[str, object] = {} + if use_thinking: + thinking_kwargs["thinking_budget"] = self._prover_config.thinking_budget response = self._llm.complete( system=PARENT_PROOF_SYSTEM, messages=[{"role": "user", "content": user_content}], - use_extended_thinking=self._prover_config.use_extended_thinking, + use_extended_thinking=use_thinking, use_cache=True, + **thinking_kwargs, ) tokens.input_tokens += response.token_usage.input_tokens tokens.output_tokens += response.token_usage.output_tokens diff --git a/agentic_research/models/agents.py b/agentic_research/models/agents.py index d9512d8..ffa939b 100644 --- a/agentic_research/models/agents.py +++ b/agentic_research/models/agents.py @@ -66,6 +66,7 @@ class ProverConfig(BaseModel): max_tokens: int = Field(default=16384, ge=1) use_extended_thinking: bool = False thinking_budget: int = Field(default=40000, ge=1000, description="Max thinking tokens for extended thinking mode") + parent_extended_thinking: bool = Field(default=True, description="Use extended thinking for parent assembly step in recursive prover") lean_timeout_seconds: int = Field(default=60, ge=1) diff --git a/tests/test_recursive_prover_extended_thinking.py b/tests/test_recursive_prover_extended_thinking.py new file mode 100644 index 0000000..8390868 --- /dev/null +++ b/tests/test_recursive_prover_extended_thinking.py @@ -0,0 +1,159 @@ +"""Tests for extended thinking in RecursiveProver parent assembly.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +from agentic_research.agents.recursive_prover import RecursiveProver +from agentic_research.models.agents import ( + LLMResponse, + ProverConfig, + TokenUsage, +) +from agentic_research.models.proof import ( + LemmaTree, + ProofNode, +) +from agentic_research.models.tools import CompilationResult, CompilationStatus, ToolStatus + + +def _mock_llm_response(content: str = "```lean\nsorry\n```") -> LLMResponse: + return LLMResponse( + content=content, + model="claude-opus-4-6-20250616", + stop_reason="end_turn", + token_usage=TokenUsage(input_tokens=100, output_tokens=50), + ) + + +def _make_tree_with_children() -> LemmaTree: + root = ProofNode( + node_id="root", + statement_nl="root statement", + statement_lean="theorem root : True", + depth=0, + children=["child1"], + ) + child = ProofNode( + node_id="child1", + statement_nl="child statement", + statement_lean="theorem child1 : True", + depth=1, + parent_id="root", + ) + return LemmaTree( + root_id="root", + nodes={"root": root, "child1": child}, + topological_order=["child1", "root"], + ) + + +def _make_ok_compilation() -> CompilationResult: + return CompilationResult( + status=ToolStatus.SUCCESS, + compilation_status=CompilationStatus.OK, + errors=[], + warnings=[], + ) + + +class TestParentAssemblyExtendedThinking: + def test_default_uses_extended_thinking(self) -> None: + mock_llm = MagicMock() + mock_repl = MagicMock() + mock_repl.execute.return_value = _make_ok_compilation() + mock_llm.complete.return_value = _mock_llm_response( + "```lean\ntheorem root : True := by trivial\n```" + ) + + config = ProverConfig() + assert config.parent_extended_thinking is True + + prover = RecursiveProver( + llm_client=mock_llm, + lean_repl=mock_repl, + prover_config=config, + ) + + tree = _make_tree_with_children() + tokens = TokenUsage() + prover._prove_parent_with_children(tree, tree.nodes["root"], tokens) + + complete_call = mock_llm.complete.call_args + assert complete_call.kwargs["use_extended_thinking"] is True + assert complete_call.kwargs["thinking_budget"] == 40000 + + def test_custom_thinking_budget_passed(self) -> None: + mock_llm = MagicMock() + mock_repl = MagicMock() + mock_repl.execute.return_value = _make_ok_compilation() + mock_llm.complete.return_value = _mock_llm_response( + "```lean\ntheorem root : True := by trivial\n```" + ) + + config = ProverConfig(thinking_budget=80000) + prover = RecursiveProver( + llm_client=mock_llm, + lean_repl=mock_repl, + prover_config=config, + ) + + tree = _make_tree_with_children() + tokens = TokenUsage() + prover._prove_parent_with_children(tree, tree.nodes["root"], tokens) + + complete_call = mock_llm.complete.call_args + assert complete_call.kwargs["thinking_budget"] == 80000 + + def test_parent_extended_thinking_false_disables(self) -> None: + mock_llm = MagicMock() + mock_repl = MagicMock() + mock_repl.execute.return_value = _make_ok_compilation() + mock_llm.complete.return_value = _mock_llm_response( + "```lean\ntheorem root : True := by trivial\n```" + ) + + config = ProverConfig(parent_extended_thinking=False) + prover = RecursiveProver( + llm_client=mock_llm, + lean_repl=mock_repl, + prover_config=config, + ) + + tree = _make_tree_with_children() + tokens = TokenUsage() + prover._prove_parent_with_children(tree, tree.nodes["root"], tokens) + + complete_call = mock_llm.complete.call_args + assert complete_call.kwargs["use_extended_thinking"] is False + assert "thinking_budget" not in complete_call.kwargs + + def test_prover_config_default_values(self) -> None: + config = ProverConfig() + assert config.parent_extended_thinking is True + assert config.thinking_budget == 40000 + assert config.use_extended_thinking is False + + def test_thinking_budget_from_config_not_default(self) -> None: + config = ProverConfig(thinking_budget=25000) + assert config.thinking_budget == 25000 + + mock_llm = MagicMock() + mock_repl = MagicMock() + mock_repl.execute.return_value = _make_ok_compilation() + mock_llm.complete.return_value = _mock_llm_response( + "```lean\ntheorem root : True := by trivial\n```" + ) + + prover = RecursiveProver( + llm_client=mock_llm, + lean_repl=mock_repl, + prover_config=config, + ) + + tree = _make_tree_with_children() + tokens = TokenUsage() + prover._prove_parent_with_children(tree, tree.nodes["root"], tokens) + + complete_call = mock_llm.complete.call_args + assert complete_call.kwargs["thinking_budget"] == 25000