Removed unused accumulator states
This commit is contained in:
parent
cf4083d2d7
commit
a592a04031
11 changed files with 5 additions and 216 deletions
|
|
@ -155,7 +155,7 @@ class ChatDeps:
|
|||
if self.session_id:
|
||||
cached = get_cached_session_context(self.session_id)
|
||||
if cached and cached.summary:
|
||||
qa_session_state.session_context = cached.render_markdown()
|
||||
qa_session_state.session_context = cached.summary
|
||||
|
||||
# Handle initial_context -> session_context for first message
|
||||
# Only applies if session_context is still empty after restoring and cache check
|
||||
|
|
|
|||
|
|
@ -181,7 +181,7 @@ async def _update_context_background(
|
|||
# Update the QASessionState with the new context
|
||||
cached = get_cached_session_context(session_id)
|
||||
if cached and cached.summary:
|
||||
qa_session_state.session_context = cached.render_markdown()
|
||||
qa_session_state.session_context = cached.summary
|
||||
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -17,10 +17,6 @@ class SessionContext(BaseModel):
|
|||
summary: str = ""
|
||||
last_updated: datetime | None = None
|
||||
|
||||
def render_markdown(self) -> str:
|
||||
"""Render context for injection into research graph."""
|
||||
return self.summary
|
||||
|
||||
|
||||
class ChatSessionState(BaseModel):
|
||||
"""State shared between frontend and agent via AG-UI."""
|
||||
|
|
|
|||
|
|
@ -1,14 +1,8 @@
|
|||
from haiku.rag.tools.analysis import (
|
||||
ANALYSIS_NAMESPACE,
|
||||
AnalysisState,
|
||||
create_analysis_toolset,
|
||||
)
|
||||
from haiku.rag.tools.analysis import create_analysis_toolset
|
||||
from haiku.rag.tools.context import ToolContext
|
||||
from haiku.rag.tools.document import (
|
||||
DOCUMENT_NAMESPACE,
|
||||
DocumentInfo,
|
||||
DocumentListResponse,
|
||||
DocumentState,
|
||||
create_document_toolset,
|
||||
find_document,
|
||||
)
|
||||
|
|
@ -20,11 +14,9 @@ from haiku.rag.tools.filters import (
|
|||
)
|
||||
from haiku.rag.tools.models import AnalysisResult, QAResult
|
||||
from haiku.rag.tools.qa import (
|
||||
QA_NAMESPACE,
|
||||
QA_SESSION_NAMESPACE,
|
||||
QAHistoryEntry,
|
||||
QASessionState,
|
||||
QAState,
|
||||
create_qa_toolset,
|
||||
)
|
||||
from haiku.rag.tools.search import SEARCH_NAMESPACE, SearchState, create_search_toolset
|
||||
|
|
@ -46,20 +38,14 @@ __all__ = [
|
|||
"SEARCH_NAMESPACE",
|
||||
"SearchState",
|
||||
"create_search_toolset",
|
||||
"DOCUMENT_NAMESPACE",
|
||||
"DocumentInfo",
|
||||
"DocumentListResponse",
|
||||
"DocumentState",
|
||||
"create_document_toolset",
|
||||
"find_document",
|
||||
"QA_NAMESPACE",
|
||||
"QA_SESSION_NAMESPACE",
|
||||
"QAState",
|
||||
"QASessionState",
|
||||
"QAHistoryEntry",
|
||||
"create_qa_toolset",
|
||||
"ANALYSIS_NAMESPACE",
|
||||
"AnalysisState",
|
||||
"create_analysis_toolset",
|
||||
"SESSION_NAMESPACE",
|
||||
"SessionState",
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
from pydantic import BaseModel
|
||||
from pydantic_ai import FunctionToolset
|
||||
|
||||
from haiku.rag.agents.rlm.agent import create_rlm_agent
|
||||
|
|
@ -14,17 +13,6 @@ from haiku.rag.tools.filters import (
|
|||
)
|
||||
from haiku.rag.tools.models import AnalysisResult
|
||||
|
||||
ANALYSIS_NAMESPACE = "haiku.rag.analysis"
|
||||
|
||||
|
||||
class AnalysisState(BaseModel):
|
||||
"""State for analysis toolset.
|
||||
|
||||
Tracks programs produced across tool invocations.
|
||||
"""
|
||||
|
||||
programs: list[str] = []
|
||||
|
||||
|
||||
def create_analysis_toolset(
|
||||
client: HaikuRAG,
|
||||
|
|
@ -39,7 +27,6 @@ def create_analysis_toolset(
|
|||
client: HaikuRAG client for document operations.
|
||||
config: Application configuration.
|
||||
context: Optional ToolContext for state accumulation.
|
||||
If provided, code executions are tracked in AnalysisState.
|
||||
If SessionState is registered, it will be used for dynamic
|
||||
document filtering.
|
||||
base_filter: Optional base SQL WHERE clause applied to searches.
|
||||
|
|
@ -48,10 +35,6 @@ def create_analysis_toolset(
|
|||
Returns:
|
||||
FunctionToolset with an analyze tool.
|
||||
"""
|
||||
# Get or create analysis state if context provided
|
||||
state: AnalysisState | None = None
|
||||
if context is not None:
|
||||
state = context.get_or_create(ANALYSIS_NAMESPACE, AnalysisState)
|
||||
|
||||
async def analyze(
|
||||
task: str,
|
||||
|
|
@ -94,8 +77,6 @@ def create_analysis_toolset(
|
|||
result = await rlm_agent.run(task, deps=deps)
|
||||
|
||||
program = result.output.program
|
||||
if state is not None and program:
|
||||
state.programs.append(program)
|
||||
|
||||
return AnalysisResult(
|
||||
answer=result.output.answer,
|
||||
|
|
|
|||
|
|
@ -7,8 +7,6 @@ from haiku.rag.tools.context import ToolContext
|
|||
from haiku.rag.tools.filters import get_session_filter
|
||||
from haiku.rag.utils import get_model
|
||||
|
||||
DOCUMENT_NAMESPACE = "haiku.rag.document"
|
||||
|
||||
DOCUMENT_SUMMARY_PROMPT = """Generate a summary of the document content provided below.
|
||||
|
||||
Start with a one-paragraph overview, then list the main topics covered, and highlight any key findings or conclusions.
|
||||
|
|
@ -40,15 +38,6 @@ class DocumentListResponse(BaseModel):
|
|||
total_documents: int
|
||||
|
||||
|
||||
class DocumentState(BaseModel):
|
||||
"""State for document toolset.
|
||||
|
||||
Tracks documents accessed during tool invocations.
|
||||
"""
|
||||
|
||||
accessed_documents: list[DocumentInfo] = []
|
||||
|
||||
|
||||
async def find_document(client: HaikuRAG, query: str):
|
||||
"""Find a document by exact URI, partial URI, or partial title match."""
|
||||
# Try exact URI match first
|
||||
|
|
@ -91,7 +80,6 @@ def create_document_toolset(
|
|||
client: HaikuRAG client for document operations.
|
||||
config: Application configuration (used for summarization LLM).
|
||||
context: Optional ToolContext for state tracking.
|
||||
If provided, accessed documents are tracked in DocumentState.
|
||||
If SessionState is registered, it will be used for dynamic
|
||||
document filtering.
|
||||
base_filter: Optional base SQL WHERE clause applied to list operations.
|
||||
|
|
@ -99,10 +87,6 @@ def create_document_toolset(
|
|||
Returns:
|
||||
FunctionToolset with list_documents, get_document, summarize_document tools.
|
||||
"""
|
||||
# Get or create document state if context provided
|
||||
state: DocumentState | None = None
|
||||
if context is not None:
|
||||
state = context.get_or_create(DOCUMENT_NAMESPACE, DocumentState)
|
||||
|
||||
async def list_documents(page: int = 1) -> DocumentListResponse:
|
||||
"""List available documents in the knowledge base.
|
||||
|
|
@ -152,16 +136,6 @@ def create_document_toolset(
|
|||
if doc is None:
|
||||
return f"Document not found: {query}"
|
||||
|
||||
# Track accessed document in state
|
||||
if state is not None:
|
||||
state.accessed_documents.append(
|
||||
DocumentInfo(
|
||||
title=doc.title or "Untitled",
|
||||
uri=doc.uri or "",
|
||||
created=doc.created_at.strftime("%Y-%m-%d"),
|
||||
)
|
||||
)
|
||||
|
||||
return (
|
||||
f"**{doc.title or 'Untitled'}**\n\n"
|
||||
f"- ID: {doc.id}\n"
|
||||
|
|
@ -184,16 +158,6 @@ def create_document_toolset(
|
|||
if doc is None:
|
||||
return f"Document not found: {query}"
|
||||
|
||||
# Track accessed document in state
|
||||
if state is not None:
|
||||
state.accessed_documents.append(
|
||||
DocumentInfo(
|
||||
title=doc.title or "Untitled",
|
||||
uri=doc.uri or "",
|
||||
created=doc.created_at.strftime("%Y-%m-%d"),
|
||||
)
|
||||
)
|
||||
|
||||
# Use LLM to generate summary
|
||||
summary_model = get_model(config.qa.model, config)
|
||||
summary_agent: Agent[None, str] = Agent(
|
||||
|
|
|
|||
|
|
@ -29,8 +29,6 @@ from haiku.rag.tools.session import (
|
|||
compute_combined_state_delta,
|
||||
)
|
||||
|
||||
QA_NAMESPACE = "haiku.rag.qa"
|
||||
|
||||
PRIOR_ANSWER_RELEVANCE_THRESHOLD = 0.7
|
||||
|
||||
|
||||
|
|
@ -77,15 +75,6 @@ from haiku.rag.agents.chat.state import _rebuild_models # noqa: E402
|
|||
_rebuild_models(QAHistoryEntry)
|
||||
|
||||
|
||||
class QAState(BaseModel):
|
||||
"""State for QA toolset.
|
||||
|
||||
Tracks Q&A history across tool invocations.
|
||||
"""
|
||||
|
||||
history: list[QAResult] = []
|
||||
|
||||
|
||||
class QASessionState(BaseModel):
|
||||
"""Extended session state for QA with embedding cache."""
|
||||
|
||||
|
|
@ -115,7 +104,6 @@ def create_qa_toolset(
|
|||
client: HaikuRAG client for search operations.
|
||||
config: Application configuration.
|
||||
context: Optional ToolContext for state accumulation.
|
||||
If provided, Q&A results are accumulated in QAState.
|
||||
If SessionState is registered, it will be used for dynamic
|
||||
document filtering and citation indexing.
|
||||
base_filter: Optional base SQL WHERE clause applied to searches.
|
||||
|
|
@ -128,10 +116,6 @@ def create_qa_toolset(
|
|||
Returns:
|
||||
FunctionToolset with an ask tool.
|
||||
"""
|
||||
# Get or create QA state if context provided
|
||||
state: QAState | None = None
|
||||
if context is not None:
|
||||
state = context.get_or_create(QA_NAMESPACE, QAState)
|
||||
|
||||
async def ask(
|
||||
question: str,
|
||||
|
|
@ -285,10 +269,6 @@ def create_qa_toolset(
|
|||
citations=citations,
|
||||
)
|
||||
|
||||
# Accumulate in QA state if context provided
|
||||
if state is not None:
|
||||
state.history.append(qa_result)
|
||||
|
||||
# Update session state with citations
|
||||
if session_state is not None:
|
||||
session_state.citations = citations
|
||||
|
|
|
|||
|
|
@ -33,17 +33,6 @@ class TestSessionContext:
|
|||
assert ctx.summary == "User discussed authentication patterns."
|
||||
assert ctx.last_updated == now
|
||||
|
||||
def test_render_markdown_empty(self):
|
||||
"""Test render_markdown returns empty string when no summary."""
|
||||
ctx = SessionContext()
|
||||
assert ctx.render_markdown() == ""
|
||||
|
||||
def test_render_markdown_with_summary(self):
|
||||
"""Test render_markdown returns the summary directly."""
|
||||
summary = "## Key Facts\n- Authentication uses JWT\n- Rate limit is 100/min"
|
||||
ctx = SessionContext(summary=summary)
|
||||
assert ctx.render_markdown() == summary
|
||||
|
||||
def test_session_context_serialization_roundtrip(self):
|
||||
"""Test SessionContext serializes and deserializes correctly."""
|
||||
now = datetime.now()
|
||||
|
|
|
|||
|
|
@ -1,30 +1,6 @@
|
|||
import pytest
|
||||
|
||||
from haiku.rag.tools import ToolContext
|
||||
from haiku.rag.tools.analysis import (
|
||||
ANALYSIS_NAMESPACE,
|
||||
AnalysisState,
|
||||
create_analysis_toolset,
|
||||
)
|
||||
|
||||
|
||||
class TestAnalysisState:
|
||||
"""Tests for AnalysisState model."""
|
||||
|
||||
def test_analysis_state_defaults(self):
|
||||
"""AnalysisState initializes with empty programs."""
|
||||
state = AnalysisState()
|
||||
assert state.programs == []
|
||||
|
||||
def test_analysis_state_serialization(self):
|
||||
"""AnalysisState serializes and deserializes correctly."""
|
||||
state = AnalysisState()
|
||||
state.programs.append("print('hello')")
|
||||
|
||||
data = state.model_dump()
|
||||
restored = AnalysisState.model_validate(data)
|
||||
assert len(restored.programs) == 1
|
||||
assert restored.programs[0] == "print('hello')"
|
||||
from haiku.rag.tools.analysis import create_analysis_toolset
|
||||
|
||||
|
||||
class TestAnalysisToolset:
|
||||
|
|
@ -44,15 +20,6 @@ class TestAnalysisToolset:
|
|||
toolset = create_analysis_toolset(analysis_client, analysis_config)
|
||||
assert "analyze" in toolset.tools
|
||||
|
||||
def test_analysis_toolset_registers_state(self, analysis_client, analysis_config):
|
||||
"""Toolset registers AnalysisState under ANALYSIS_NAMESPACE."""
|
||||
context = ToolContext()
|
||||
create_analysis_toolset(analysis_client, analysis_config, context=context)
|
||||
|
||||
state = context.get(ANALYSIS_NAMESPACE)
|
||||
assert state is not None
|
||||
assert isinstance(state, AnalysisState)
|
||||
|
||||
def test_analysis_toolset_custom_tool_name(self, analysis_client, analysis_config):
|
||||
"""Toolset supports custom tool name."""
|
||||
toolset = create_analysis_toolset(
|
||||
|
|
|
|||
|
|
@ -1,11 +1,8 @@
|
|||
import pytest
|
||||
|
||||
from haiku.rag.tools import ToolContext
|
||||
from haiku.rag.tools.document import (
|
||||
DOCUMENT_NAMESPACE,
|
||||
DocumentInfo,
|
||||
DocumentListResponse,
|
||||
DocumentState,
|
||||
create_document_toolset,
|
||||
)
|
||||
|
||||
|
|
@ -36,11 +33,6 @@ class TestDocumentModels:
|
|||
assert response.total_pages == 3
|
||||
assert response.total_documents == 125
|
||||
|
||||
def test_document_state_defaults(self):
|
||||
"""DocumentState initializes with empty accessed list."""
|
||||
state = DocumentState()
|
||||
assert state.accessed_documents == []
|
||||
|
||||
|
||||
@pytest.mark.vcr()
|
||||
class TestDocumentToolset:
|
||||
|
|
@ -63,15 +55,6 @@ class TestDocumentToolset:
|
|||
assert "get_document" in toolset.tools
|
||||
assert "summarize_document" in toolset.tools
|
||||
|
||||
def test_document_toolset_registers_state(self, doc_client, doc_config):
|
||||
"""Toolset registers DocumentState under DOCUMENT_NAMESPACE."""
|
||||
context = ToolContext()
|
||||
create_document_toolset(doc_client, doc_config, context=context)
|
||||
|
||||
state = context.get(DOCUMENT_NAMESPACE)
|
||||
assert state is not None
|
||||
assert isinstance(state, DocumentState)
|
||||
|
||||
|
||||
@pytest.mark.vcr()
|
||||
class TestDocumentToolExecution:
|
||||
|
|
@ -135,20 +118,6 @@ class TestDocumentToolExecution:
|
|||
|
||||
assert "Document not found" in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_document_tracks_in_state(self, doc_client, doc_config):
|
||||
"""get_document tracks accessed documents in state."""
|
||||
context = ToolContext()
|
||||
toolset = create_document_toolset(doc_client, doc_config, context=context)
|
||||
|
||||
get_tool = toolset.tools["get_document"]
|
||||
await get_tool.function("Python Guide")
|
||||
|
||||
state = context.get(DOCUMENT_NAMESPACE)
|
||||
assert isinstance(state, DocumentState)
|
||||
assert len(state.accessed_documents) == 1
|
||||
assert state.accessed_documents[0].title == "Python Guide"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_documents_with_base_filter(self, doc_client, doc_config):
|
||||
"""list_documents respects base_filter."""
|
||||
|
|
|
|||
|
|
@ -1,40 +1,6 @@
|
|||
import pytest
|
||||
|
||||
from haiku.rag.tools import QAResult, ToolContext
|
||||
from haiku.rag.tools.qa import QA_NAMESPACE, QAState, create_qa_toolset
|
||||
|
||||
|
||||
class TestQAState:
|
||||
"""Tests for QAState model."""
|
||||
|
||||
def test_qa_state_defaults(self):
|
||||
"""QAState initializes with empty history."""
|
||||
state = QAState()
|
||||
assert state.history == []
|
||||
|
||||
def test_qa_state_add_result(self):
|
||||
"""Can add QAResult to history."""
|
||||
state = QAState()
|
||||
result = QAResult(question="What is Python?", answer="A programming language.")
|
||||
state.history.append(result)
|
||||
assert len(state.history) == 1
|
||||
assert state.history[0].question == "What is Python?"
|
||||
|
||||
def test_qa_state_serialization(self):
|
||||
"""QAState serializes and deserializes correctly."""
|
||||
state = QAState()
|
||||
state.history.append(
|
||||
QAResult(
|
||||
question="Test?",
|
||||
answer="Answer.",
|
||||
confidence=0.95,
|
||||
)
|
||||
)
|
||||
|
||||
data = state.model_dump()
|
||||
restored = QAState.model_validate(data)
|
||||
assert len(restored.history) == 1
|
||||
assert restored.history[0].confidence == 0.95
|
||||
from haiku.rag.tools.qa import create_qa_toolset
|
||||
|
||||
|
||||
class TestQAToolset:
|
||||
|
|
@ -54,15 +20,6 @@ class TestQAToolset:
|
|||
toolset = create_qa_toolset(qa_client_simple, qa_config)
|
||||
assert "ask" in toolset.tools
|
||||
|
||||
def test_qa_toolset_registers_state(self, qa_client_simple, qa_config):
|
||||
"""Toolset registers QAState under QA_NAMESPACE."""
|
||||
context = ToolContext()
|
||||
create_qa_toolset(qa_client_simple, qa_config, context=context)
|
||||
|
||||
state = context.get(QA_NAMESPACE)
|
||||
assert state is not None
|
||||
assert isinstance(state, QAState)
|
||||
|
||||
def test_qa_toolset_custom_tool_name(self, qa_client_simple, qa_config):
|
||||
"""Toolset supports custom tool name."""
|
||||
toolset = create_qa_toolset(
|
||||
|
|
|
|||
Loading…
Reference in a new issue