From 74683ab0a022a24841a009ed262f83196499b332 Mon Sep 17 00:00:00 2001 From: guix4ever <111418839+guix4ever@users.noreply.github.com> Date: Fri, 24 Jul 2026 13:52:26 +0800 Subject: [PATCH] fix: inherit agent model in skill forge --- raven/config/raven.py | 4 +-- raven/context_engine/factory.py | 5 +++- raven/memory_engine/skill_forge/rewriter.py | 3 ++ tests/test_phase_a_default_engine.py | 32 +++++++++++++++++++-- tests/test_skill_forge_rewriter.py | 7 +++++ 5 files changed, 46 insertions(+), 5 deletions(-) diff --git a/raven/config/raven.py b/raven/config/raven.py index 877a168..e1714d9 100644 --- a/raven/config/raven.py +++ b/raven/config/raven.py @@ -964,8 +964,8 @@ class SkillForgeConfig(_Base): with RRF output size (local_pool_top_k + mass_pool_top_k dedupe).""" llm_gate_model: str | None = None - """Optional model override for gate calls. ``None`` → use the - provider's default chat model (typically the agent's main model).""" + """Optional model override for gate calls. ``None`` uses the active + agent model.""" llm_gate_temperature: float = 0.0 """Sampling temperature for gate calls. 0.0 for deterministic diff --git a/raven/context_engine/factory.py b/raven/context_engine/factory.py index 11f522d..bbe13a3 100644 --- a/raven/context_engine/factory.py +++ b/raven/context_engine/factory.py @@ -108,6 +108,7 @@ def build_context_engine( rewriter, gate = _build_rewriter_and_gate( provider=provider, + model=model, skill_forge_config=skill_forge_config, skill_forge_router_config=skill_forge_router_config, ) @@ -216,6 +217,7 @@ def _build_router( def _build_rewriter_and_gate( *, provider: LLMProvider, + model: str, skill_forge_config: "SkillForgeConfig | None", skill_forge_router_config: "SkillForgeRouterConfig", ) -> "tuple[QueryRewriter | None, LLMGateFilter | None]": @@ -236,6 +238,7 @@ def _build_rewriter_and_gate( if bool(getattr(skill_forge_config, "rewrite_enabled", False)): rewriter = QueryRewriter( provider, + model=model, max_tokens=int(getattr(skill_forge_config, "rewrite_max_tokens", 8192) or 8192), ) @@ -248,7 +251,7 @@ def _build_rewriter_and_gate( provider, max_select=int(getattr(skill_forge_config, "llm_gate_max_select", 2) or 2), legacy_top_k=int(skill_forge_router_config.top_k or 5), - model=getattr(skill_forge_config, "llm_gate_model", None) or None, + model=getattr(skill_forge_config, "llm_gate_model", None) or model, temperature=float(getattr(skill_forge_config, "llm_gate_temperature", 0.0)), max_tokens=int(getattr(skill_forge_config, "llm_gate_max_tokens", 8192) or 8192), ) diff --git a/raven/memory_engine/skill_forge/rewriter.py b/raven/memory_engine/skill_forge/rewriter.py index 7009e03..b897875 100644 --- a/raven/memory_engine/skill_forge/rewriter.py +++ b/raven/memory_engine/skill_forge/rewriter.py @@ -67,10 +67,12 @@ def __init__( self, provider: "LLMProvider", *, + model: str | None = None, max_tokens: int = 8192, temperature: float = 0.3, ) -> None: self._provider = provider + self._model = model self._max_tokens = max_tokens self._temperature = temperature @@ -85,6 +87,7 @@ async def analyze(self, query: str) -> RewriteResult: resp = await asyncio.wait_for( self._provider.chat_with_retry( messages=[{"role": "user", "content": prompt}], + model=self._model or None, max_tokens=self._max_tokens, temperature=self._temperature, ), diff --git a/tests/test_phase_a_default_engine.py b/tests/test_phase_a_default_engine.py index 505c4a3..8045932 100644 --- a/tests/test_phase_a_default_engine.py +++ b/tests/test_phase_a_default_engine.py @@ -25,6 +25,7 @@ ContextConfig, HubSourceConfig, MemoryConfig, + SkillForgeConfig, SkillForgeRouterConfig, ) from raven.context_engine import ContextAssembler @@ -81,6 +82,8 @@ def _build_engine( backend=None, hub_endpoint: str | None = None, memory_config: MemoryConfig | None = None, + model: str = "stub", + skill_forge_config: SkillForgeConfig | None = None, ) -> ContextAssembler: builder = ContextBuilder(workspace=tmp_path) engine = build_context_engine( @@ -88,11 +91,12 @@ def _build_engine( config=ContextConfig(), builder=builder, provider=_StubProvider(), - model="stub", + model=model, context_window_tokens=8192, get_tool_definitions=_stub_get_defs, backend=backend, memory_config=memory_config or MemoryConfig(), + skill_forge_config=skill_forge_config, skill_forge_router_config=SkillForgeRouterConfig( hub=HubSourceConfig(endpoint=hub_endpoint), ), @@ -102,10 +106,14 @@ def _build_engine( def _router_sources(engine: ContextAssembler): - skills = next(b for b in engine._builders if isinstance(b, SkillsSegmentBuilder)) + skills = _skills_builder(engine) return [type(s) for s in skills._router._sources], skills._router._sources +def _skills_builder(engine: ContextAssembler) -> SkillsSegmentBuilder: + return next(b for b in engine._builders if isinstance(b, SkillsSegmentBuilder)) + + def _memory_builder(engine: ContextAssembler) -> MemorySegmentBuilder: return next(b for b in engine._builders if isinstance(b, MemorySegmentBuilder)) @@ -124,6 +132,26 @@ def test_returns_assembler_without_backend(self, tmp_path: Path) -> None: assert isinstance(engine, ContextAssembler) assert _memory_builder(engine)._backend is None + def test_skill_forge_llms_inherit_agent_model(self, tmp_path: Path) -> None: + engine = _build_engine( + tmp_path, + model="minimax/minimax-m3", + skill_forge_config=SkillForgeConfig(), + ) + skills = _skills_builder(engine) + assert skills._rewriter._model == "minimax/minimax-m3" + assert skills._gate._model == "minimax/minimax-m3" + + def test_gate_model_override_takes_priority(self, tmp_path: Path) -> None: + engine = _build_engine( + tmp_path, + model="minimax/minimax-m3", + skill_forge_config=SkillForgeConfig(llm_gate_model="anthropic/claude-sonnet-4-5"), + ) + skills = _skills_builder(engine) + assert skills._rewriter._model == "minimax/minimax-m3" + assert skills._gate._model == "anthropic/claude-sonnet-4-5" + # --------------------------------------------------------------------------- # SkillForgeRouter assembly — which sources are present diff --git a/tests/test_skill_forge_rewriter.py b/tests/test_skill_forge_rewriter.py index ffa6b26..aa4a2be 100644 --- a/tests/test_skill_forge_rewriter.py +++ b/tests/test_skill_forge_rewriter.py @@ -64,6 +64,13 @@ async def test_analyze_returns_rewritten_query() -> None: assert result.rewritten_query == "generate pdf reports" +async def test_analyze_forwards_configured_model() -> None: + provider = _StubProvider(json.dumps({"need_retrieval": False})) + rewriter = QueryRewriter(provider, model="minimax/minimax-m3") + await rewriter.analyze("hello there") + assert provider.calls[0]["model"] == "minimax/minimax-m3" + + async def test_analyze_handles_code_fence_wrapping() -> None: provider = _StubProvider('```json\n{"need_retrieval": true, "rewritten_query": "trim"}\n```') result = await QueryRewriter(provider).analyze("verbose query")