Unify Citation class and context formatting between chat and research agents

This commit is contained in:
Yiorgis Gozadinos 2026-01-22 14:38:24 +02:00
parent 74eac57ffa
commit 9ed8773672
No known key found for this signature in database
15 changed files with 253 additions and 89 deletions

View file

@ -1,7 +1,6 @@
from haiku.rag.agents.chat import ( from haiku.rag.agents.chat import (
ChatDeps, ChatDeps,
ChatSessionState, ChatSessionState,
CitationInfo,
QAResponse, QAResponse,
SearchAgent, SearchAgent,
SearchDeps, SearchDeps,
@ -9,6 +8,7 @@ from haiku.rag.agents.chat import (
) )
from haiku.rag.agents.qa import QuestionAnswerAgent, get_qa_agent from haiku.rag.agents.qa import QuestionAnswerAgent, get_qa_agent
from haiku.rag.agents.research import ( from haiku.rag.agents.research import (
Citation,
EvaluationResult, EvaluationResult,
ResearchContext, ResearchContext,
ResearchDependencies, ResearchDependencies,
@ -19,7 +19,6 @@ from haiku.rag.agents.research.graph import (
build_conversational_graph, build_conversational_graph,
build_research_graph, build_research_graph,
) )
from haiku.rag.agents.research.models import Citation
from haiku.rag.agents.research.state import ResearchDeps, ResearchState from haiku.rag.agents.research.state import ResearchDeps, ResearchState
__all__ = [ __all__ = [
@ -42,7 +41,6 @@ __all__ = [
"SearchAgent", "SearchAgent",
"ChatDeps", "ChatDeps",
"ChatSessionState", "ChatSessionState",
"CitationInfo",
"QAResponse", "QAResponse",
"SearchDeps", "SearchDeps",
] ]

View file

@ -5,7 +5,6 @@ from haiku.rag.agents.chat.state import (
AGUI_STATE_KEY, AGUI_STATE_KEY,
ChatDeps, ChatDeps,
ChatSessionState, ChatSessionState,
CitationInfo,
QAResponse, QAResponse,
SearchDeps, SearchDeps,
SessionContext, SessionContext,
@ -18,7 +17,6 @@ __all__ = [
"SearchAgent", "SearchAgent",
"ChatDeps", "ChatDeps",
"ChatSessionState", "ChatSessionState",
"CitationInfo",
"QAResponse", "QAResponse",
"SearchDeps", "SearchDeps",
"SessionContext", "SessionContext",

View file

@ -11,12 +11,12 @@ from haiku.rag.agents.chat.state import (
MAX_QA_HISTORY, MAX_QA_HISTORY,
ChatDeps, ChatDeps,
ChatSessionState, ChatSessionState,
CitationInfo,
QAResponse, QAResponse,
build_document_filter, build_document_filter,
) )
from haiku.rag.agents.research.dependencies import ResearchContext from haiku.rag.agents.research.dependencies import ResearchContext
from haiku.rag.agents.research.graph import build_conversational_graph from haiku.rag.agents.research.graph import build_conversational_graph
from haiku.rag.agents.research.models import Citation
from haiku.rag.agents.research.state import ResearchDeps, ResearchState from haiku.rag.agents.research.state import ResearchDeps, ResearchState
from haiku.rag.config.models import AppConfig from haiku.rag.config.models import AppConfig
from haiku.rag.utils import get_model from haiku.rag.utils import get_model
@ -97,7 +97,7 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
# Build citation infos for frontend display # Build citation infos for frontend display
citation_infos = [ citation_infos = [
CitationInfo( Citation(
index=i + 1, index=i + 1,
document_id=r.document_id or "", document_id=r.document_id or "",
chunk_id=r.chunk_id or "", chunk_id=r.chunk_id or "",
@ -217,7 +217,7 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
# Build citation infos for frontend and history # Build citation infos for frontend and history
citation_infos = [ citation_infos = [
CitationInfo( Citation(
index=i + 1, index=i + 1,
document_id=c.document_id, document_id=c.document_id,
chunk_id=c.chunk_id, chunk_id=c.chunk_id,

View file

@ -4,6 +4,7 @@ from typing import Any
from pydantic import BaseModel from pydantic import BaseModel
from haiku.rag.agents.research.models import Citation
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig from haiku.rag.config.models import AppConfig
from haiku.rag.store.models import SearchResult from haiku.rag.store.models import SearchResult
@ -13,26 +14,13 @@ MAX_QA_HISTORY = 50
AGUI_STATE_KEY = "haiku.rag.chat" AGUI_STATE_KEY = "haiku.rag.chat"
class CitationInfo(BaseModel):
"""Citation info for frontend display."""
index: int
document_id: str
chunk_id: str
document_uri: str
document_title: str | None = None
page_numbers: list[int] = []
headings: list[str] | None = None
content: str
class QAResponse(BaseModel): class QAResponse(BaseModel):
"""A Q&A pair from conversation history with citations.""" """A Q&A pair from conversation history with citations."""
question: str question: str
answer: str answer: str
confidence: float = 0.9 confidence: float = 0.9
citations: list[CitationInfo] = [] citations: list[Citation] = []
@property @property
def sources(self) -> list[str]: def sources(self) -> list[str]:
@ -57,7 +45,7 @@ class ChatSessionState(BaseModel):
"""State shared between frontend and agent via AG-UI.""" """State shared between frontend and agent via AG-UI."""
session_id: str = "" session_id: str = ""
citations: list[CitationInfo] = [] citations: list[Citation] = []
qa_history: list[QAResponse] = [] qa_history: list[QAResponse] = []
background_context: str | None = None background_context: str | None = None
session_context: SessionContext | None = None session_context: SessionContext | None = None
@ -106,7 +94,7 @@ class ChatDeps:
] ]
if "citations" in state_data: if "citations" in state_data:
self.session_state.citations = [ self.session_state.citations = [
CitationInfo(**c) if isinstance(c, dict) else c Citation(**c) if isinstance(c, dict) else c
for c in state_data.get("citations", []) for c in state_data.get("citations", [])
] ]
if "background_context" in state_data: if "background_context" in state_data:

View file

@ -1,5 +1,6 @@
from haiku.rag.agents.research.dependencies import ResearchContext, ResearchDependencies from haiku.rag.agents.research.dependencies import ResearchContext, ResearchDependencies
from haiku.rag.agents.research.models import ( from haiku.rag.agents.research.models import (
Citation,
EvaluationResult, EvaluationResult,
ResearchReport, ResearchReport,
SearchAnswer, SearchAnswer,

View file

@ -29,8 +29,17 @@ from haiku.rag.config.models import AppConfig
from haiku.rag.utils import build_prompt, get_model from haiku.rag.utils import build_prompt, get_model
def format_context_for_prompt(context: ResearchContext) -> str: def format_context_for_prompt(
"""Format the research context as XML for planning prompts.""" context: ResearchContext,
include_pending_questions: bool = True,
) -> str:
"""Format the research context as XML for prompts.
Args:
context: The research context to format.
include_pending_questions: Whether to include pending sub-questions.
Set to False for synthesis prompts where pending questions aren't relevant.
"""
context_data: dict[str, object] = {} context_data: dict[str, object] = {}
if context.background_context: if context.background_context:
@ -38,7 +47,7 @@ def format_context_for_prompt(context: ResearchContext) -> str:
context_data["question"] = context.original_question context_data["question"] = context.original_question
if context.sub_questions: if include_pending_questions and context.sub_questions:
context_data["pending_questions"] = context.sub_questions context_data["pending_questions"] = context.sub_questions
if context.qa_responses: if context.qa_responses:
@ -47,34 +56,7 @@ def format_context_for_prompt(context: ResearchContext) -> str:
"question": qa.query, "question": qa.query,
"answer": qa.answer, "answer": qa.answer,
"confidence": qa.confidence, "confidence": qa.confidence,
"source": qa.citations[0].document_title or qa.citations[0].document_uri "source": qa.primary_source,
if qa.citations
else None,
}
for qa in context.qa_responses
]
return format_as_xml(context_data, root_tag="context")
def format_conversational_context_for_prompt(context: ResearchContext) -> str:
"""Format context for synthesis prompts."""
context_data: dict[str, object] = {}
if context.background_context:
context_data["background"] = context.background_context
context_data["question"] = context.original_question
if context.qa_responses:
context_data["prior_answers"] = [
{
"question": qa.query,
"answer": qa.answer,
"confidence": qa.confidence,
"source": qa.citations[0].document_title or qa.citations[0].document_uri
if qa.citations
else None,
} }
for qa in context.qa_responses for qa in context.qa_responses
] ]
@ -503,7 +485,9 @@ def build_conversational_graph(
deps_type=ResearchDependencies, deps_type=ResearchDependencies,
) )
context_xml = format_conversational_context_for_prompt(state.context) context_xml = format_context_for_prompt(
state.context, include_pending_questions=False
)
prompt = f"Answer the question based on the gathered evidence.\n\n{context_xml}" prompt = f"Answer the question based on the gathered evidence.\n\n{context_xml}"
agent_deps = ResearchDependencies( agent_deps = ResearchDependencies(
client=deps.client, client=deps.client,

View file

@ -23,8 +23,13 @@ class ResearchPlan(BaseModel):
class Citation(BaseModel): class Citation(BaseModel):
"""Resolved citation with full metadata for display/visual grounding.""" """Resolved citation with full metadata for display/visual grounding.
Used by both research graph and chat agent. The optional index field
supports UI display ordering in chat contexts.
"""
index: int | None = None
document_id: str document_id: str
chunk_id: str chunk_id: str
document_uri: str document_uri: str
@ -59,6 +64,14 @@ class SearchAnswer(RawSearchAnswer):
description="Resolved citations with full metadata", description="Resolved citations with full metadata",
) )
@property
def primary_source(self) -> str | None:
"""Get primary source title from citations."""
if not self.citations:
return None
first = self.citations[0]
return first.document_title or first.document_uri
@classmethod @classmethod
def from_raw( def from_raw(
cls, cls,

View file

@ -21,8 +21,8 @@ from haiku.rag.agents.chat.state import (
AGUI_STATE_KEY, AGUI_STATE_KEY,
ChatDeps, ChatDeps,
ChatSessionState, ChatSessionState,
CitationInfo,
) )
from haiku.rag.agents.research.models import Citation
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
@ -104,7 +104,7 @@ class ChatApp(App):
self.session_state: ChatSessionState | None = None self.session_state: ChatSessionState | None = None
self._is_processing = False self._is_processing = False
self._tool_call_widgets: dict[str, Any] = {} self._tool_call_widgets: dict[str, Any] = {}
self._last_citations: list[CitationInfo] = [] self._last_citations: list[Citation] = []
self._selected_citation_idx: int | None = None self._selected_citation_idx: int | None = None
self._current_worker: Worker[None] | None = None self._current_worker: Worker[None] | None = None
self._message_history: list[ModelMessage] = [] self._message_history: list[ModelMessage] = []
@ -170,7 +170,7 @@ class ChatApp(App):
snapshot = getattr(meta_event, "snapshot", {}) snapshot = getattr(meta_event, "snapshot", {})
chat_state = snapshot.get(AGUI_STATE_KEY, snapshot) chat_state = snapshot.get(AGUI_STATE_KEY, snapshot)
self._last_citations = [ self._last_citations = [
CitationInfo(**c) for c in chat_state["citations"] Citation(**c) for c in chat_state["citations"]
] ]
async def _event_stream_handler( async def _event_stream_handler(

View file

@ -4,7 +4,7 @@ from textual.containers import Horizontal, VerticalScroll
from textual.message import Message from textual.message import Message
from textual.widgets import Collapsible, LoadingIndicator, Markdown, Static from textual.widgets import Collapsible, LoadingIndicator, Markdown, Static
from haiku.rag.agents.chat.state import CitationInfo from haiku.rag.agents.research.models import Citation
if TYPE_CHECKING: if TYPE_CHECKING:
from textual.app import ComposeResult from textual.app import ComposeResult
@ -94,7 +94,7 @@ class CitationWidget(Collapsible):
super().__init__() super().__init__()
self.citation_index = citation_index self.citation_index = citation_index
def __init__(self, citation: CitationInfo, **kwargs) -> None: def __init__(self, citation: Citation, **kwargs) -> None:
title = f"[{citation.index}] {citation.document_title or citation.document_uri}" title = f"[{citation.index}] {citation.document_title or citation.document_uri}"
if citation.page_numbers: if citation.page_numbers:
pages = ", ".join(map(str, citation.page_numbers[:3])) pages = ", ".join(map(str, citation.page_numbers[:3]))
@ -120,7 +120,8 @@ class CitationWidget(Collapsible):
def on_focus(self) -> None: def on_focus(self) -> None:
"""When focused, mark as selected.""" """When focused, mark as selected."""
self.post_message(self.Selected(self.citation.index - 1)) index = self.citation.index or 1
self.post_message(self.Selected(index - 1))
def on_key(self, event: "Key") -> None: def on_key(self, event: "Key") -> None:
"""Handle Enter to toggle expand/collapse.""" """Handle Enter to toggle expand/collapse."""
@ -326,7 +327,7 @@ class ChatHistory(VerticalScroll):
widget.mark_complete() widget.mark_complete()
widget.add_class("complete") widget.add_class("complete")
async def add_citations(self, citations: list[CitationInfo]) -> None: async def add_citations(self, citations: list[Citation]) -> None:
"""Add citations inline after a response.""" """Add citations inline after a response."""
if not citations: if not citations:
return return

View file

@ -6,12 +6,12 @@ from haiku.rag.agents.chat import (
AGUI_STATE_KEY, AGUI_STATE_KEY,
ChatDeps, ChatDeps,
ChatSessionState, ChatSessionState,
CitationInfo,
QAResponse, QAResponse,
SearchAgent, SearchAgent,
create_chat_agent, create_chat_agent,
) )
from haiku.rag.agents.chat.state import MAX_QA_HISTORY from haiku.rag.agents.chat.state import MAX_QA_HISTORY
from haiku.rag.agents.research.models import Citation
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config from haiku.rag.config import Config
@ -88,9 +88,9 @@ def test_chat_agent_has_dynamic_system_prompt():
assert "add_background_context" in func_names assert "add_background_context" in func_names
def test_citation_info(): def test_citation():
"""Test CitationInfo model.""" """Test Citation model."""
citation = CitationInfo( citation = Citation(
index=1, index=1,
document_id="doc-123", document_id="doc-123",
chunk_id="chunk-456", chunk_id="chunk-456",
@ -108,7 +108,7 @@ def test_citation_info():
def test_qa_response(): def test_qa_response():
"""Test QAResponse model.""" """Test QAResponse model."""
citation = CitationInfo( citation = Citation(
index=1, index=1,
document_id="doc-123", document_id="doc-123",
chunk_id="chunk-456", chunk_id="chunk-456",
@ -131,7 +131,7 @@ def test_qa_response():
def test_qa_response_sources_with_uri_fallback(): def test_qa_response_sources_with_uri_fallback():
"""Test QAResponse.sources falls back to URI when title is None.""" """Test QAResponse.sources falls back to URI when title is None."""
citation = CitationInfo( citation = Citation(
index=1, index=1,
document_id="doc-123", document_id="doc-123",
chunk_id="chunk-456", chunk_id="chunk-456",

View file

@ -4,10 +4,10 @@ from pathlib import Path
import pytest import pytest
from haiku.rag.agents.chat.state import ( from haiku.rag.agents.chat.state import (
CitationInfo,
QAResponse, QAResponse,
SessionContext, SessionContext,
) )
from haiku.rag.agents.research.models import Citation
from haiku.rag.config import Config from haiku.rag.config import Config
@ -88,7 +88,7 @@ class TestSummarizeSession:
answer="The API uses JWT tokens for authentication.", answer="The API uses JWT tokens for authentication.",
confidence=0.95, confidence=0.95,
citations=[ citations=[
CitationInfo( Citation(
index=1, index=1,
document_id="doc-1", document_id="doc-1",
chunk_id="chunk-1", chunk_id="chunk-1",
@ -121,7 +121,7 @@ class TestSummarizeSession:
answer="The API uses JWT tokens for authentication.", answer="The API uses JWT tokens for authentication.",
confidence=0.95, confidence=0.95,
citations=[ citations=[
CitationInfo( Citation(
index=1, index=1,
document_id="doc-1", document_id="doc-1",
chunk_id="chunk-1", chunk_id="chunk-1",
@ -136,7 +136,7 @@ class TestSummarizeSession:
answer="Rate limiting is set to 100 requests per minute.", answer="Rate limiting is set to 100 requests per minute.",
confidence=0.9, confidence=0.9,
citations=[ citations=[
CitationInfo( Citation(
index=1, index=1,
document_id="doc-2", document_id="doc-2",
chunk_id="chunk-2", chunk_id="chunk-2",
@ -151,7 +151,7 @@ class TestSummarizeSession:
answer="Use the /refresh endpoint with your refresh token.", answer="Use the /refresh endpoint with your refresh token.",
confidence=0.85, confidence=0.85,
citations=[ citations=[
CitationInfo( Citation(
index=1, index=1,
document_id="doc-1", document_id="doc-1",
chunk_id="chunk-3", chunk_id="chunk-3",

View file

@ -210,7 +210,7 @@ def test_chat_deps_state_setter_without_session_state():
def test_chat_deps_state_setter_with_citation_dicts(): def test_chat_deps_state_setter_with_citation_dicts():
"""Test ChatDeps.state setter converts citation dicts to CitationInfo.""" """Test ChatDeps.state setter converts citation dicts to Citation."""
from unittest.mock import MagicMock from unittest.mock import MagicMock
from haiku.rag.agents.chat.state import AGUI_STATE_KEY, ChatDeps, ChatSessionState from haiku.rag.agents.chat.state import AGUI_STATE_KEY, ChatDeps, ChatSessionState

View file

@ -0,0 +1,181 @@
from haiku.rag.agents.research.models import Citation, SearchAnswer
class TestCitation:
"""Tests for unified Citation class."""
def test_citation_without_index(self):
"""Test Citation can be created without index (research graph use case)."""
citation = Citation(
document_id="doc-1",
chunk_id="chunk-1",
document_uri="test.md",
document_title="Test Document",
page_numbers=[1, 2],
headings=["Introduction"],
content="Test content",
)
assert citation.document_id == "doc-1"
assert citation.chunk_id == "chunk-1"
assert citation.document_uri == "test.md"
assert citation.document_title == "Test Document"
assert citation.page_numbers == [1, 2]
assert citation.headings == ["Introduction"]
assert citation.content == "Test content"
assert citation.index is None
def test_citation_with_index(self):
"""Test Citation can be created with index (chat agent use case)."""
citation = Citation(
index=1,
document_id="doc-1",
chunk_id="chunk-1",
document_uri="test.md",
content="Test content",
)
assert citation.index == 1
assert citation.document_id == "doc-1"
def test_citation_index_defaults_to_none(self):
"""Test Citation index defaults to None."""
citation = Citation(
document_id="doc-1",
chunk_id="chunk-1",
document_uri="test.md",
content="Test content",
)
assert citation.index is None
def test_citation_serialization_includes_index_when_set(self):
"""Test Citation serialization includes index when set."""
citation = Citation(
index=2,
document_id="doc-1",
chunk_id="chunk-1",
document_uri="test.md",
content="Test content",
)
data = citation.model_dump()
assert data["index"] == 2
def test_citation_deserialization_from_dict_with_index(self):
"""Test Citation can be deserialized from dict with index (AG-UI state sync)."""
data = {
"index": 1,
"document_id": "doc-1",
"chunk_id": "chunk-1",
"document_uri": "test.md",
"document_title": "Test Doc",
"page_numbers": [1, 2],
"headings": ["Intro"],
"content": "Test content",
}
citation = Citation.model_validate(data)
assert citation.index == 1
assert citation.document_id == "doc-1"
class TestSearchAnswerPrimarySource:
"""Tests for SearchAnswer.primary_source property."""
def test_primary_source_returns_title_when_available(self):
"""Test primary_source returns first citation's title."""
answer = SearchAnswer(
query="test query",
answer="test answer",
citations=[
Citation(
document_id="doc-1",
chunk_id="chunk-1",
document_uri="test.md",
document_title="Test Document",
content="content",
),
],
)
assert answer.primary_source == "Test Document"
def test_primary_source_returns_uri_when_no_title(self):
"""Test primary_source returns URI when title is None."""
answer = SearchAnswer(
query="test query",
answer="test answer",
citations=[
Citation(
document_id="doc-1",
chunk_id="chunk-1",
document_uri="test.md",
document_title=None,
content="content",
),
],
)
assert answer.primary_source == "test.md"
def test_primary_source_returns_none_when_no_citations(self):
"""Test primary_source returns None when no citations."""
answer = SearchAnswer(
query="test query",
answer="test answer",
citations=[],
)
assert answer.primary_source is None
class TestFormatContextMerged:
"""Tests for merged format_context_for_prompt function."""
def test_format_context_includes_pending_questions_by_default(self):
"""Test format_context_for_prompt includes pending_questions by default."""
from haiku.rag.agents.research.dependencies import ResearchContext
from haiku.rag.agents.research.graph import format_context_for_prompt
context = ResearchContext(
original_question="What is X?",
sub_questions=["What is A?", "What is B?"],
)
result = format_context_for_prompt(context)
assert "<pending_questions>" in result
assert "What is A?" in result
assert "What is B?" in result
def test_format_context_excludes_pending_questions_when_flag_false(self):
"""Test format_context_for_prompt excludes pending_questions when flag is False."""
from haiku.rag.agents.research.dependencies import ResearchContext
from haiku.rag.agents.research.graph import format_context_for_prompt
context = ResearchContext(
original_question="What is X?",
sub_questions=["What is A?", "What is B?"],
)
result = format_context_for_prompt(context, include_pending_questions=False)
assert "<pending_questions>" not in result
assert "What is A?" not in result
def test_format_context_uses_primary_source_helper(self):
"""Test format_context_for_prompt uses primary_source from SearchAnswer."""
from haiku.rag.agents.research.dependencies import ResearchContext
from haiku.rag.agents.research.graph import format_context_for_prompt
context = ResearchContext(
original_question="What is X?",
)
# Add a QA response with citation
answer = SearchAnswer(
query="What is A?",
answer="A is...",
confidence=0.9,
citations=[
Citation(
document_id="doc-1",
chunk_id="chunk-1",
document_uri="test.md",
document_title="Test Document",
content="content",
),
],
)
context.add_qa_response(answer)
result = format_context_for_prompt(context)
assert "Test Document" in result

View file

@ -83,25 +83,25 @@ def test_format_context_for_prompt_excludes_background_when_none():
assert "<background>" not in result assert "<background>" not in result
def test_format_conversational_context_for_prompt_includes_background(): def test_format_context_for_prompt_without_pending_includes_background():
"""Test format_conversational_context_for_prompt includes background.""" """Test format_context_for_prompt with include_pending_questions=False includes background."""
from haiku.rag.agents.research.graph import format_conversational_context_for_prompt from haiku.rag.agents.research.graph import format_context_for_prompt
context = ResearchContext( context = ResearchContext(
original_question="What is X?", original_question="What is X?",
background_context="X is a concept in domain Y.", background_context="X is a concept in domain Y.",
) )
result = format_conversational_context_for_prompt(context) result = format_context_for_prompt(context, include_pending_questions=False)
assert "X is a concept in domain Y." in result assert "X is a concept in domain Y." in result
assert "<background>" in result assert "<background>" in result
def test_format_conversational_context_for_prompt_excludes_background_when_none(): def test_format_context_for_prompt_without_pending_excludes_background_when_none():
"""Test format_conversational_context_for_prompt excludes background when None.""" """Test format_context_for_prompt with include_pending_questions=False excludes background when None."""
from haiku.rag.agents.research.graph import format_conversational_context_for_prompt from haiku.rag.agents.research.graph import format_context_for_prompt
context = ResearchContext(original_question="What is X?") context = ResearchContext(original_question="What is X?")
result = format_conversational_context_for_prompt(context) result = format_context_for_prompt(context, include_pending_questions=False)
assert "<background>" not in result assert "<background>" not in result

View file

@ -173,7 +173,7 @@ async def test_chat_history_can_add_tool_calls(temp_db_path: Path):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_chat_history_can_add_citations(temp_db_path: Path): async def test_chat_history_can_add_citations(temp_db_path: Path):
"""Test that ChatHistory can display inline citations.""" """Test that ChatHistory can display inline citations."""
from haiku.rag.agents.chat.state import CitationInfo from haiku.rag.agents.research.models import Citation
from haiku.rag.chat.app import ChatApp from haiku.rag.chat.app import ChatApp
from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget
@ -188,7 +188,7 @@ async def test_chat_history_can_add_citations(temp_db_path: Path):
chat_history = app.query_one(ChatHistory) chat_history = app.query_one(ChatHistory)
test_citations = [ test_citations = [
CitationInfo( Citation(
index=1, index=1,
document_id="doc1", document_id="doc1",
chunk_id="chunk1", chunk_id="chunk1",
@ -198,7 +198,7 @@ async def test_chat_history_can_add_citations(temp_db_path: Path):
headings=["Section 1"], headings=["Section 1"],
content="This is some test content from doc 1", content="This is some test content from doc 1",
), ),
CitationInfo( Citation(
index=2, index=2,
document_id="doc2", document_id="doc2",
chunk_id="chunk2", chunk_id="chunk2",
@ -285,7 +285,7 @@ async def test_clear_chat_resets_session(temp_db_path: Path):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_citation_expand_collapse_with_enter(temp_db_path: Path): async def test_citation_expand_collapse_with_enter(temp_db_path: Path):
"""Test that pressing Enter on a focused citation toggles expand/collapse.""" """Test that pressing Enter on a focused citation toggles expand/collapse."""
from haiku.rag.agents.chat.state import CitationInfo from haiku.rag.agents.research.models import Citation
from haiku.rag.chat.app import ChatApp from haiku.rag.chat.app import ChatApp
from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget
@ -300,7 +300,7 @@ async def test_citation_expand_collapse_with_enter(temp_db_path: Path):
chat_history = app.query_one(ChatHistory) chat_history = app.query_one(ChatHistory)
# Add a citation # Add a citation
test_citation = CitationInfo( test_citation = Citation(
index=1, index=1,
document_id="doc1", document_id="doc1",
chunk_id="chunk1", chunk_id="chunk1",