Add preamble parameter to chat agent and citations_history to ChatSessionState
This commit is contained in:
parent
5767158737
commit
84c48e1541
5 changed files with 72 additions and 3 deletions
|
|
@ -92,6 +92,7 @@ def prepare_chat_context(
|
||||||
def create_chat_agent(
|
def create_chat_agent(
|
||||||
config: AppConfig,
|
config: AppConfig,
|
||||||
features: list[str] | None = None,
|
features: list[str] | None = None,
|
||||||
|
preamble: str | None = None,
|
||||||
) -> Agent[ChatDeps, str]:
|
) -> Agent[ChatDeps, str]:
|
||||||
"""Create the chat agent with composed toolsets.
|
"""Create the chat agent with composed toolsets.
|
||||||
|
|
||||||
|
|
@ -100,6 +101,9 @@ def create_chat_agent(
|
||||||
features: List of features to enable. Defaults to DEFAULT_FEATURES
|
features: List of features to enable. Defaults to DEFAULT_FEATURES
|
||||||
(search, documents, qa). Available features: "search",
|
(search, documents, qa). Available features: "search",
|
||||||
"documents", "qa", "analysis".
|
"documents", "qa", "analysis".
|
||||||
|
preamble: Optional custom identity/rules section for the system prompt.
|
||||||
|
When provided, replaces the default identity prompt. Tool guidance,
|
||||||
|
feature rules, and closing are still appended by the builder.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The configured chat agent.
|
The configured chat agent.
|
||||||
|
|
@ -133,7 +137,7 @@ def create_chat_agent(
|
||||||
model,
|
model,
|
||||||
deps_type=ChatDeps,
|
deps_type=ChatDeps,
|
||||||
output_type=str,
|
output_type=str,
|
||||||
instructions=build_chat_prompt(features),
|
instructions=build_chat_prompt(features, preamble=preamble),
|
||||||
toolsets=toolsets,
|
toolsets=toolsets,
|
||||||
retries=3,
|
retries=3,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -26,7 +26,10 @@ _PROMPT_QA_CLOSING = (
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_chat_prompt(features: list[str]) -> str:
|
def build_chat_prompt(
|
||||||
|
features: list[str],
|
||||||
|
preamble: str | None = None,
|
||||||
|
) -> str:
|
||||||
"""Build a chat system prompt from the given feature list.
|
"""Build a chat system prompt from the given feature list.
|
||||||
|
|
||||||
Each feature adds its relevant tool guidance to the prompt.
|
Each feature adds its relevant tool guidance to the prompt.
|
||||||
|
|
@ -34,11 +37,14 @@ def build_chat_prompt(features: list[str]) -> str:
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
features: List of feature names (e.g., ["search", "documents", "qa"]).
|
features: List of feature names (e.g., ["search", "documents", "qa"]).
|
||||||
|
preamble: Optional custom identity/rules section. When provided,
|
||||||
|
replaces the default identity prompt. Tool guidance, feature
|
||||||
|
rules, and closing are still appended.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The composed system prompt string.
|
The composed system prompt string.
|
||||||
"""
|
"""
|
||||||
parts = [_PROMPT_BASE]
|
parts = [preamble if preamble is not None else _PROMPT_BASE]
|
||||||
|
|
||||||
# Add feature-specific critical rules
|
# Add feature-specific critical rules
|
||||||
if "qa" in features:
|
if "qa" in features:
|
||||||
|
|
|
||||||
|
|
@ -16,6 +16,7 @@ class ChatSessionState(BaseModel):
|
||||||
|
|
||||||
initial_context: str | None = None
|
initial_context: str | None = None
|
||||||
citations: list[Citation] = []
|
citations: list[Citation] = []
|
||||||
|
citations_history: list[list[Citation]] = []
|
||||||
qa_history: list["QAHistoryEntry"] = []
|
qa_history: list["QAHistoryEntry"] = []
|
||||||
session_context: SessionContext | None = None
|
session_context: SessionContext | None = None
|
||||||
document_filter: list[str] = []
|
document_filter: list[str] = []
|
||||||
|
|
|
||||||
|
|
@ -155,3 +155,35 @@ def test_build_chat_prompt_includes_analysis():
|
||||||
assert "analyze" in prompt
|
assert "analyze" in prompt
|
||||||
assert "search" in prompt
|
assert "search" in prompt
|
||||||
assert "ask" in prompt
|
assert "ask" in prompt
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_chat_prompt_with_preamble():
|
||||||
|
"""Custom preamble replaces the default identity section."""
|
||||||
|
custom = "You are a custom assistant."
|
||||||
|
prompt = build_chat_prompt(DEFAULT_FEATURES, preamble=custom)
|
||||||
|
|
||||||
|
assert prompt.startswith(custom)
|
||||||
|
# Tool guidance should still be appended
|
||||||
|
assert "search" in prompt
|
||||||
|
assert "ask" in prompt
|
||||||
|
# Default identity should NOT be present
|
||||||
|
assert "haiku.rag" not in prompt
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_chat_prompt_without_preamble_uses_default():
|
||||||
|
"""Without preamble, the default identity section is used."""
|
||||||
|
prompt = build_chat_prompt(DEFAULT_FEATURES)
|
||||||
|
|
||||||
|
assert "haiku.rag" in prompt
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_chat_agent_with_preamble():
|
||||||
|
"""create_chat_agent passes preamble through to build_chat_prompt."""
|
||||||
|
custom = "You are a domain expert."
|
||||||
|
agent = create_chat_agent(Config, preamble=custom)
|
||||||
|
|
||||||
|
assert agent is not None
|
||||||
|
# _instructions is the internal list of instruction strings/callables
|
||||||
|
assert any(
|
||||||
|
custom in instr for instr in agent._instructions if isinstance(instr, str)
|
||||||
|
)
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
from haiku.rag.agents.chat.state import ChatSessionState
|
from haiku.rag.agents.chat.state import ChatSessionState
|
||||||
|
from haiku.rag.agents.research.models import Citation
|
||||||
from haiku.rag.tools.session import SessionContext, SessionState
|
from haiku.rag.tools.session import SessionContext, SessionState
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -113,3 +114,28 @@ def test_chat_session_state_model_dump_json_serializes_datetime():
|
||||||
# datetime should be serialized as ISO string, not datetime object
|
# datetime should be serialized as ISO string, not datetime object
|
||||||
assert isinstance(snapshot["session_context"]["last_updated"], str)
|
assert isinstance(snapshot["session_context"]["last_updated"], str)
|
||||||
assert snapshot["session_context"]["last_updated"] == "2025-01-27T12:00:00"
|
assert snapshot["session_context"]["last_updated"] == "2025-01-27T12:00:00"
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_session_state_citations_history_default():
|
||||||
|
"""citations_history defaults to empty list."""
|
||||||
|
state = ChatSessionState()
|
||||||
|
assert state.citations_history == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_session_state_citations_history_roundtrip():
|
||||||
|
"""citations_history serializes and deserializes correctly."""
|
||||||
|
citation = Citation(
|
||||||
|
index=1,
|
||||||
|
document_id="d1",
|
||||||
|
chunk_id="c1",
|
||||||
|
document_uri="test://doc",
|
||||||
|
document_title="Doc",
|
||||||
|
page_numbers=[],
|
||||||
|
headings=None,
|
||||||
|
content="content",
|
||||||
|
)
|
||||||
|
state = ChatSessionState(citations_history=[[citation]])
|
||||||
|
data = state.model_dump(mode="json")
|
||||||
|
restored = ChatSessionState.model_validate(data)
|
||||||
|
assert len(restored.citations_history) == 1
|
||||||
|
assert restored.citations_history[0][0].chunk_id == "c1"
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue