Merge pull request #341 from ggozad/fix/skill-document-filtering
Fix: apply document filtering to list_documents & analyze tool
This commit is contained in:
commit
0a4cd31168
7 changed files with 83 additions and 8 deletions
|
|
@ -6,6 +6,11 @@
|
||||||
- **S3/Object storage support**: Connect to LanceDB on S3, GCS, Azure Blob, or HDFS via `lancedb.uri` and `storage_options` config. Supports S3-compatible stores with custom endpoints.
|
- **S3/Object storage support**: Connect to LanceDB on S3, GCS, Azure Blob, or HDFS via `lancedb.uri` and `storage_options` config. Supports S3-compatible stores with custom endpoints.
|
||||||
- **Remote skill generation**: `create-skill` now supports remote databases — omit `--db` and provide `--config-file` to generate skills that connect to object storage at runtime instead of bundling the database.
|
- **Remote skill generation**: `create-skill` now supports remote databases — omit `--db` and provide `--config-file` to generate skills that connect to object storage at runtime instead of bundling the database.
|
||||||
|
|
||||||
|
### Fixed
|
||||||
|
|
||||||
|
- **Skill `list_documents` ignores `document_filter`**: `list_documents` tool now respects `state.document_filter`, consistent with `search`, `ask`, and `research`
|
||||||
|
- **Skill `analyze` ignores `document_filter`**: `analyze` tool now uses `state.document_filter` (combined with any explicit `filter` parameter). Added `document_filter` field to `RLMState`
|
||||||
|
|
||||||
## [0.38.0] - 2026-04-07
|
## [0.38.0] - 2026-04-07
|
||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|
|
||||||
|
|
@ -41,7 +41,7 @@ class RAGState(BaseModel):
|
||||||
|
|
||||||
- **citations** — Accumulated citations from `ask` calls, with sequential indexing across calls.
|
- **citations** — Accumulated citations from `ask` calls, with sequential indexing across calls.
|
||||||
- **qa_history** — Questions and answers from `ask` calls. Prior Q&A is used as context for follow-up questions when embeddings are similar.
|
- **qa_history** — Questions and answers from `ask` calls. Prior Q&A is used as context for follow-up questions when embeddings are similar.
|
||||||
- **document_filter** — SQL WHERE clause applied to `search`, `ask`, and `research` calls. Set this to scope queries to specific documents.
|
- **document_filter** — SQL WHERE clause applied to `search`, `list_documents`, `ask`, and `research` calls. Set this to scope queries to specific documents.
|
||||||
- **searches** — Search results keyed by query string.
|
- **searches** — Search results keyed by query string.
|
||||||
- **documents** — Documents seen via `list_documents` or `get_document` (deduplicated by ID).
|
- **documents** — Documents seen via `list_documents` or `get_document` (deduplicated by ID).
|
||||||
- **reports** — Research reports from `research` calls.
|
- **reports** — Research reports from `research` calls.
|
||||||
|
|
|
||||||
|
|
@ -33,6 +33,7 @@ The skill manages an `RLMState` under the `"rlm"` namespace:
|
||||||
|
|
||||||
```python
|
```python
|
||||||
class RLMState(BaseModel):
|
class RLMState(BaseModel):
|
||||||
|
document_filter: str | None = None
|
||||||
analyses: list[AnalysisEntry] = []
|
analyses: list[AnalysisEntry] = []
|
||||||
|
|
||||||
class AnalysisEntry(BaseModel):
|
class AnalysisEntry(BaseModel):
|
||||||
|
|
@ -41,7 +42,8 @@ class AnalysisEntry(BaseModel):
|
||||||
program: str | None = None
|
program: str | None = None
|
||||||
```
|
```
|
||||||
|
|
||||||
Each `analyze` call appends an `AnalysisEntry` with the question, answer, and executed program.
|
- **document_filter** — SQL WHERE clause applied to `analyze` calls (combined with any explicit `filter` parameter). Set this to scope analysis to specific documents.
|
||||||
|
- **analyses** — Each `analyze` call appends an `AnalysisEntry` with the question, answer, and executed program.
|
||||||
|
|
||||||
## Usage with RAG Skill
|
## Usage with RAG Skill
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -8,6 +8,7 @@ from haiku.rag.agents.research.models import Citation
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.store.models.chunk import SearchResult
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
from haiku.rag.tools.document import DocumentInfo
|
from haiku.rag.tools.document import DocumentInfo
|
||||||
|
from haiku.rag.tools.filters import combine_filters
|
||||||
from haiku.rag.tools.qa import QAHistoryEntry
|
from haiku.rag.tools.qa import QAHistoryEntry
|
||||||
from haiku.skills.state import SkillRunDeps
|
from haiku.skills.state import SkillRunDeps
|
||||||
|
|
||||||
|
|
@ -90,11 +91,12 @@ async def skill_list_documents(
|
||||||
config: AppConfig,
|
config: AppConfig,
|
||||||
limit: int | None = None,
|
limit: int | None = None,
|
||||||
offset: int | None = None,
|
offset: int | None = None,
|
||||||
|
filter: str | None = None,
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
||||||
async with HaikuRAG(db_path, config=config, read_only=True) as rag:
|
async with HaikuRAG(db_path, config=config, read_only=True) as rag:
|
||||||
documents = await rag.list_documents(limit, offset)
|
documents = await rag.list_documents(limit, offset, filter=filter)
|
||||||
return [
|
return [
|
||||||
{
|
{
|
||||||
"id": doc.id,
|
"id": doc.id,
|
||||||
|
|
@ -350,8 +352,14 @@ def create_skill_tools(
|
||||||
limit: Maximum number of documents to return.
|
limit: Maximum number of documents to return.
|
||||||
offset: Number of documents to skip.
|
offset: Number of documents to skip.
|
||||||
"""
|
"""
|
||||||
result = await skill_list_documents(db_path, config, limit, offset)
|
|
||||||
state = _get_state(ctx, state_type)
|
state = _get_state(ctx, state_type)
|
||||||
|
result = await skill_list_documents(
|
||||||
|
db_path,
|
||||||
|
config,
|
||||||
|
limit,
|
||||||
|
offset,
|
||||||
|
filter=state.document_filter if state else None,
|
||||||
|
)
|
||||||
if state:
|
if state:
|
||||||
update_documents_state(state.documents, result)
|
update_documents_state(state.documents, result)
|
||||||
return result
|
return result
|
||||||
|
|
@ -469,11 +477,12 @@ def create_skill_tools(
|
||||||
document: Optional document ID or title to pre-load for analysis.
|
document: Optional document ID or title to pre-load for analysis.
|
||||||
filter: Optional SQL WHERE clause to filter documents.
|
filter: Optional SQL WHERE clause to filter documents.
|
||||||
"""
|
"""
|
||||||
output, answer, program = await skill_analyze(
|
|
||||||
db_path, config, question, document=document, filter=filter
|
|
||||||
)
|
|
||||||
|
|
||||||
state = _get_state(ctx, state_type)
|
state = _get_state(ctx, state_type)
|
||||||
|
state_filter = state.document_filter if state else None
|
||||||
|
effective_filter = combine_filters(state_filter, filter)
|
||||||
|
output, answer, program = await skill_analyze(
|
||||||
|
db_path, config, question, document=document, filter=effective_filter
|
||||||
|
)
|
||||||
if state:
|
if state:
|
||||||
state.analyses.append(
|
state.analyses.append(
|
||||||
AnalysisEntry(
|
AnalysisEntry(
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ from haiku.skills.parser import parse_skill_md
|
||||||
|
|
||||||
|
|
||||||
class RLMState(BaseModel):
|
class RLMState(BaseModel):
|
||||||
|
document_filter: str | None = None
|
||||||
analyses: list[AnalysisEntry] = []
|
analyses: list[AnalysisEntry] = []
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -276,6 +276,17 @@ class TestListDocumentsTool:
|
||||||
assert isinstance(state.documents[0], DocumentInfo)
|
assert isinstance(state.documents[0], DocumentInfo)
|
||||||
assert state.documents[0].id is not None
|
assert state.documents[0].id is not None
|
||||||
|
|
||||||
|
async def test_list_documents_applies_document_filter_from_state(self, rag_db):
|
||||||
|
from haiku.rag.skills.rag import RAGState, create_skill
|
||||||
|
|
||||||
|
skill = create_skill(db_path=rag_db)
|
||||||
|
list_docs = _get_tool(skill, "list_documents")
|
||||||
|
state = RAGState(document_filter="title = 'AI Overview'")
|
||||||
|
ctx = _make_ctx(state)
|
||||||
|
results = await list_docs(ctx)
|
||||||
|
assert len(results) == 1
|
||||||
|
assert results[0]["title"] == "AI Overview"
|
||||||
|
|
||||||
async def test_list_documents_no_duplicates_in_state(self, rag_db):
|
async def test_list_documents_no_duplicates_in_state(self, rag_db):
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
from haiku.rag.skills.rag import RAGState, create_skill
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -155,6 +155,53 @@ class TestAnalyzeTool:
|
||||||
assert state.analyses[0].answer == "42"
|
assert state.analyses[0].answer == "42"
|
||||||
assert state.analyses[0].program == "print(42)"
|
assert state.analyses[0].program == "print(42)"
|
||||||
|
|
||||||
|
async def test_analyze_applies_document_filter_from_state(
|
||||||
|
self, rag_db, monkeypatch
|
||||||
|
):
|
||||||
|
from haiku.rag.skills.rlm import RLMState, create_skill
|
||||||
|
|
||||||
|
captured_kwargs = {}
|
||||||
|
|
||||||
|
async def mock_rlm(self, question, **kwargs):
|
||||||
|
captured_kwargs.update(kwargs)
|
||||||
|
return RLMResult(answer="42", program="print(42)")
|
||||||
|
|
||||||
|
monkeypatch.setattr(HaikuRAG, "rlm", mock_rlm)
|
||||||
|
|
||||||
|
skill = create_skill(db_path=rag_db)
|
||||||
|
analyze = _get_tool(skill, "analyze")
|
||||||
|
state = RLMState(document_filter="title = 'AI Overview'")
|
||||||
|
ctx = _make_ctx(state)
|
||||||
|
await analyze(ctx, question="How many documents?")
|
||||||
|
assert captured_kwargs.get("filter") == "title = 'AI Overview'"
|
||||||
|
|
||||||
|
async def test_analyze_combines_state_filter_with_explicit_filter(
|
||||||
|
self, rag_db, monkeypatch
|
||||||
|
):
|
||||||
|
from haiku.rag.skills.rlm import RLMState, create_skill
|
||||||
|
|
||||||
|
captured_kwargs = {}
|
||||||
|
|
||||||
|
async def mock_rlm(self, question, **kwargs):
|
||||||
|
captured_kwargs.update(kwargs)
|
||||||
|
return RLMResult(answer="Result", program="code()")
|
||||||
|
|
||||||
|
monkeypatch.setattr(HaikuRAG, "rlm", mock_rlm)
|
||||||
|
|
||||||
|
skill = create_skill(db_path=rag_db)
|
||||||
|
analyze = _get_tool(skill, "analyze")
|
||||||
|
state = RLMState(document_filter="title = 'AI Overview'")
|
||||||
|
ctx = _make_ctx(state)
|
||||||
|
await analyze(
|
||||||
|
ctx,
|
||||||
|
question="Count pages",
|
||||||
|
filter="uri LIKE '%test%'",
|
||||||
|
)
|
||||||
|
result_filter = captured_kwargs["filter"]
|
||||||
|
assert isinstance(result_filter, str)
|
||||||
|
assert "title = 'AI Overview'" in result_filter
|
||||||
|
assert "uri LIKE '%test%'" in result_filter
|
||||||
|
|
||||||
async def test_analyze_with_document_and_filter(self, rag_db, monkeypatch):
|
async def test_analyze_with_document_and_filter(self, rag_db, monkeypatch):
|
||||||
from haiku.rag.skills.rlm import create_skill
|
from haiku.rag.skills.rlm import create_skill
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue