Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions raven/config/raven.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 4 additions & 1 deletion raven/context_engine/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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]":
Expand All @@ -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),
)

Expand All @@ -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),
)
Expand Down
3 changes: 3 additions & 0 deletions raven/memory_engine/skill_forge/rewriter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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,
),
Expand Down
32 changes: 30 additions & 2 deletions tests/test_phase_a_default_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
ContextConfig,
HubSourceConfig,
MemoryConfig,
SkillForgeConfig,
SkillForgeRouterConfig,
)
from raven.context_engine import ContextAssembler
Expand Down Expand Up @@ -81,18 +82,21 @@ 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(
workspace=tmp_path,
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),
),
Expand All @@ -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))

Expand All @@ -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
Expand Down
7 changes: 7 additions & 0 deletions tests/test_skill_forge_rewriter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down