Propagate domain_preamble to skill instructions and main agent preamble

This commit is contained in:
Yiorgis Gozadinos 2026-04-01 13:07:58 +03:00
parent fa92151311
commit c4c90b3930
No known key found for this signature in database
6 changed files with 101 additions and 6 deletions

View file

@ -16,7 +16,7 @@ from starlette.routing import Route
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config import load_yaml_config from haiku.rag.config import load_yaml_config
from haiku.rag.config.models import AppConfig from haiku.rag.config.models import AppConfig
from haiku.rag.skills.rag import AGENT_PREAMBLE, create_skill from haiku.rag.skills.rag import create_skill, get_agent_preamble
from haiku.rag.utils import get_model from haiku.rag.utils import get_model
from haiku.skills import ( from haiku.skills import (
SkillDeps, SkillDeps,
@ -74,7 +74,9 @@ toolset = SkillToolset(skills=[skill])
agent = Agent( agent = Agent(
get_model(Config.qa.model, Config), get_model(Config.qa.model, Config),
instructions=build_system_prompt(toolset.skill_catalog, preamble=AGENT_PREAMBLE), instructions=build_system_prompt(
toolset.skill_catalog, preamble=get_agent_preamble(Config)
),
toolsets=[toolset], toolsets=[toolset],
deps_type=SkillDeps, deps_type=SkillDeps,
) )

View file

@ -30,7 +30,7 @@ from textual.worker import Worker
from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config import get_config from haiku.rag.config import get_config
from haiku.rag.skills.rag import AGENT_PREAMBLE, RAGState from haiku.rag.skills.rag import RAGState, get_agent_preamble
from haiku.skills.agent import ( from haiku.skills.agent import (
SkillToolset, SkillToolset,
run_agui_stream, run_agui_stream,
@ -157,7 +157,8 @@ class ChatApp(App):
self._agent = Agent( self._agent = Agent(
self._model, self._model,
instructions=build_system_prompt( instructions=build_system_prompt(
self._toolset.skill_catalog, preamble=AGENT_PREAMBLE self._toolset.skill_catalog,
preamble=get_agent_preamble(self.config),
), ),
toolsets=[self._toolset], toolsets=[self._toolset],
) )

View file

@ -24,6 +24,13 @@ CRITICAL RULES:
_RAG_TOOLS = ["search", "list_documents", "get_document", "ask", "research"] _RAG_TOOLS = ["search", "list_documents", "get_document", "ask", "research"]
def get_agent_preamble(config: AppConfig) -> str:
"""Build the main agent preamble, prepending domain_preamble if configured."""
if config.prompts.domain_preamble:
return f"{config.prompts.domain_preamble}\n\n{AGENT_PREAMBLE}"
return AGENT_PREAMBLE
class RAGState(BaseModel): class RAGState(BaseModel):
citations: list[Citation] = Field(default_factory=list) citations: list[Citation] = Field(default_factory=list)
qa_history: list[QAHistoryEntry] = Field(default_factory=list) qa_history: list[QAHistoryEntry] = Field(default_factory=list)
@ -88,11 +95,15 @@ def create_skill(
tools = create_skill_tools(db_path, config, RAGState, _RAG_TOOLS) tools = create_skill_tools(db_path, config, RAGState, _RAG_TOOLS)
extras = create_skill_extras(db_path, config) extras = create_skill_extras(db_path, config)
skill_instructions = instructions()
if config.prompts.domain_preamble and skill_instructions:
skill_instructions = f"{config.prompts.domain_preamble}\n\n{skill_instructions}"
return Skill( return Skill(
metadata=skill_metadata(), metadata=skill_metadata(),
source=SkillSource.ENTRYPOINT, source=SkillSource.ENTRYPOINT,
path=_skill_path, path=_skill_path,
instructions=instructions(), instructions=skill_instructions,
tools=list(tools.values()), tools=list(tools.values()),
extras=extras, extras=extras,
state_type=STATE_TYPE, state_type=STATE_TYPE,

View file

@ -69,11 +69,15 @@ def create_skill(
tools = create_skill_tools(db_path, config, RLMState, ["analyze"]) tools = create_skill_tools(db_path, config, RLMState, ["analyze"])
extras = create_skill_extras(db_path, config) extras = create_skill_extras(db_path, config)
skill_instructions = instructions()
if config.prompts.domain_preamble and skill_instructions:
skill_instructions = f"{config.prompts.domain_preamble}\n\n{skill_instructions}"
return Skill( return Skill(
metadata=skill_metadata(), metadata=skill_metadata(),
source=SkillSource.ENTRYPOINT, source=SkillSource.ENTRYPOINT,
path=_skill_path, path=_skill_path,
instructions=instructions(), instructions=skill_instructions,
tools=list(tools.values()), tools=list(tools.values()),
extras=extras, extras=extras,
state_type=STATE_TYPE, state_type=STATE_TYPE,

View file

@ -2,6 +2,7 @@ from unittest.mock import AsyncMock
from haiku.rag.agents.research.models import Citation, ResearchReport from haiku.rag.agents.research.models import Citation, ResearchReport
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig
from haiku.rag.skills.rag import ( from haiku.rag.skills.rag import (
STATE_NAMESPACE, STATE_NAMESPACE,
STATE_TYPE, STATE_TYPE,
@ -52,6 +53,55 @@ class TestRAGModuleAPI:
assert skill.instructions == instructions() assert skill.instructions == instructions()
class TestGetAgentPreamble:
def test_without_domain_preamble(self):
from haiku.rag.skills.rag import AGENT_PREAMBLE, get_agent_preamble
config = AppConfig()
assert get_agent_preamble(config) == AGENT_PREAMBLE
def test_with_domain_preamble(self):
from haiku.rag.config.models import PromptsConfig
from haiku.rag.skills.rag import AGENT_PREAMBLE, get_agent_preamble
config = AppConfig(
prompts=PromptsConfig(
domain_preamble="This knowledge base contains C-146 aircraft documents."
)
)
result = get_agent_preamble(config)
assert result.startswith(
"This knowledge base contains C-146 aircraft documents."
)
assert AGENT_PREAMBLE in result
class TestDomainPreambleInSkillInstructions:
def test_create_skill_without_domain_preamble(self, test_app_config, temp_db_path):
from haiku.rag.skills.rag import create_skill, instructions
skill = create_skill(config=test_app_config, db_path=temp_db_path)
assert skill.instructions == instructions()
def test_create_skill_with_domain_preamble(self, temp_db_path):
from haiku.rag.config.models import PromptsConfig
from haiku.rag.skills.rag import create_skill, instructions
config = AppConfig(
prompts=PromptsConfig(
domain_preamble="This knowledge base contains C-146 aircraft documents."
)
)
skill = create_skill(config=config, db_path=temp_db_path)
assert skill.instructions is not None
assert skill.instructions.startswith(
"This knowledge base contains C-146 aircraft documents."
)
base_instructions = instructions()
assert base_instructions is not None
assert base_instructions in skill.instructions
class TestRAGSkillCreation: class TestRAGSkillCreation:
def test_create_skill_returns_valid_skill(self, test_app_config, temp_db_path): def test_create_skill_returns_valid_skill(self, test_app_config, temp_db_path):
from haiku.rag.skills.rag import create_skill from haiku.rag.skills.rag import create_skill

View file

@ -2,6 +2,7 @@ from unittest.mock import AsyncMock
from haiku.rag.agents.rlm.models import RLMResult from haiku.rag.agents.rlm.models import RLMResult
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig
from haiku.rag.skills.rlm import ( from haiku.rag.skills.rlm import (
STATE_NAMESPACE, STATE_NAMESPACE,
STATE_TYPE, STATE_TYPE,
@ -91,6 +92,32 @@ class TestRLMSkillCreation:
assert skill.metadata.name == "rag-rlm" assert skill.metadata.name == "rag-rlm"
class TestDomainPreambleInRLMSkillInstructions:
def test_create_skill_without_domain_preamble(self, test_app_config, temp_db_path):
from haiku.rag.skills.rlm import create_skill, instructions
skill = create_skill(config=test_app_config, db_path=temp_db_path)
assert skill.instructions == instructions()
def test_create_skill_with_domain_preamble(self, temp_db_path):
from haiku.rag.config.models import PromptsConfig
from haiku.rag.skills.rlm import create_skill, instructions
config = AppConfig(
prompts=PromptsConfig(
domain_preamble="This knowledge base contains C-146 aircraft documents."
)
)
skill = create_skill(config=config, db_path=temp_db_path)
assert skill.instructions is not None
assert skill.instructions.startswith(
"This knowledge base contains C-146 aircraft documents."
)
base_instructions = instructions()
assert base_instructions is not None
assert base_instructions in skill.instructions
class TestAnalyzeTool: class TestAnalyzeTool:
async def test_analyze_returns_result(self, rag_db, monkeypatch): async def test_analyze_returns_result(self, rag_db, monkeypatch):
from haiku.rag.skills.rlm import create_skill from haiku.rag.skills.rlm import create_skill