Make the capability image-attachment gate a vision bool

This commit is contained in:
Yiorgis Gozadinos 2026-07-24 16:00:01 +03:00
parent 6c5bc0aae1
commit f87fabe556
No known key found for this signature in database
5 changed files with 30 additions and 25 deletions

View file

@ -22,7 +22,7 @@ from pydantic_ai.toolsets import AgentToolset
from haiku.rag.capabilities._tools import CodeExecutionEntry, search_corpus from haiku.rag.capabilities._tools import CodeExecutionEntry, search_corpus
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig, ModelConfig 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.store.models.citation import Citation, resolve_citations from haiku.rag.store.models.citation import Citation, resolve_citations
from haiku.rag.tools.search import build_binary_parts_from_results from haiku.rag.tools.search import build_binary_parts_from_results
@ -87,7 +87,7 @@ class RAGCapabilityBase[StateT: BaseModel](AbstractCapability[Any]):
state_type: type[StateT] state_type: type[StateT]
state_namespace: str state_namespace: str
instruction_text: str instruction_text: str
model: ModelConfig vision: bool
tool_names: frozenset[str] tool_names: frozenset[str]
request_limit: int | None = None request_limit: int | None = None
state: StateT | None = field(default=None, repr=False) state: StateT | None = field(default=None, repr=False)
@ -220,7 +220,7 @@ class RAGCapabilityBase[StateT: BaseModel](AbstractCapability[Any]):
) )
state = cast(Any, self.state) state = cast(Any, self.state)
state.searches[query] = results state.searches[query] = results
if self.model.vision and (parts := build_binary_parts_from_results(results)): if self.vision and (parts := build_binary_parts_from_results(results)):
return ToolReturn(return_value=formatted, content=parts) return ToolReturn(return_value=formatted, content=parts)
return formatted return formatted

View file

@ -13,7 +13,7 @@ from haiku.rag.capabilities._base import (
RAGCapabilityBase, RAGCapabilityBase,
resolve_db_path, resolve_db_path,
) )
from haiku.rag.config.models import AppConfig, ModelConfig from haiku.rag.config.models import AppConfig
from haiku.rag.sandbox import AnalysisContext, Sandbox from haiku.rag.sandbox import AnalysisContext, Sandbox
from haiku.rag.store.models.chunk import SearchResult from haiku.rag.store.models.chunk import SearchResult
from haiku.rag.store.models.citation import Citation from haiku.rag.store.models.citation import Citation
@ -129,25 +129,27 @@ def create_capability(
*, *,
defer_loading: bool = True, defer_loading: bool = True,
request_limit: int | None = 30, request_limit: int | None = 30,
model: ModelConfig | None = None, vision: bool | None = None,
) -> AnalysisCapability: ) -> AnalysisCapability:
"""Create a native Pydantic AI analysis capability. """Create a native Pydantic AI analysis capability.
``model`` sets the capability's image-attachment gate and should be the ``vision`` gates whether picture chunks are attached to search results as
model the hosting agent actually runs. Defaults to ``config.analysis.model`` images, and should reflect the model the hosting agent actually runs.
(falling back to ``config.qa.model``). Defaults to ``config.analysis.model.vision`` (falling back to
``config.qa.model.vision``).
""" """
if config is None: if config is None:
from haiku.rag.config import get_config from haiku.rag.config import get_config
config = get_config() config = get_config()
analysis_model = config.analysis.model or config.qa.model
return AnalysisCapability( return AnalysisCapability(
db_path=resolve_db_path(db_path, config), db_path=resolve_db_path(db_path, config),
config=config, config=config,
state_type=AnalysisState, state_type=AnalysisState,
state_namespace=STATE_NAMESPACE, state_namespace=STATE_NAMESPACE,
instruction_text=instructions(), instruction_text=instructions(),
model=model or config.analysis.model or config.qa.model, vision=analysis_model.vision if vision is None else vision,
tool_names=_TOOL_NAMES, tool_names=_TOOL_NAMES,
request_limit=request_limit, request_limit=request_limit,
id=_CAPABILITY_ID, id=_CAPABILITY_ID,

View file

@ -12,7 +12,7 @@ from haiku.rag.capabilities._base import (
RAGCapabilityBase, RAGCapabilityBase,
resolve_db_path, resolve_db_path,
) )
from haiku.rag.config.models import AppConfig, ModelConfig 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.store.models.citation import Citation from haiku.rag.store.models.citation import Citation
@ -71,12 +71,13 @@ def create_capability(
*, *,
defer_loading: bool = True, defer_loading: bool = True,
request_limit: int | None = 20, request_limit: int | None = 20,
model: ModelConfig | None = None, vision: bool | None = None,
) -> RAGCapability: ) -> RAGCapability:
"""Create a native Pydantic AI RAG capability. """Create a native Pydantic AI RAG capability.
``model`` sets the capability's image-attachment gate and should be the ``vision`` gates whether picture chunks are attached to search results as
model the hosting agent actually runs. Defaults to ``config.qa.model``. images, and should reflect the model the hosting agent actually runs.
Defaults to ``config.qa.model.vision``.
""" """
if config is None: if config is None:
from haiku.rag.config import get_config from haiku.rag.config import get_config
@ -88,7 +89,7 @@ def create_capability(
state_type=RAGState, state_type=RAGState,
state_namespace=STATE_NAMESPACE, state_namespace=STATE_NAMESPACE,
instruction_text=instructions(), instruction_text=instructions(),
model=model or config.qa.model, vision=config.qa.model.vision if vision is None else vision,
tool_names=_TOOL_NAMES, tool_names=_TOOL_NAMES,
request_limit=request_limit, request_limit=request_limit,
id=_CAPABILITY_ID, id=_CAPABILITY_ID,

View file

@ -55,7 +55,7 @@ def run_chat(
db_path=db_path, db_path=db_path,
config=config, config=config,
defer_loading=defer_loading, defer_loading=defer_loading,
model=driving_model, vision=driving_model.vision,
) )
) )
@ -67,7 +67,7 @@ def run_chat(
db_path=db_path, db_path=db_path,
config=config, config=config,
defer_loading=defer_loading, defer_loading=defer_loading,
model=driving_model, vision=driving_model.vision,
) )
) )

View file

@ -47,23 +47,25 @@ def test_run_chat_defers_multiple_capabilities(temp_db_path: Path):
@pytest.mark.parametrize( @pytest.mark.parametrize(
("enabled", "expected_model"), ("enabled", "expected_model", "expected_vision"),
[ [
(["analysis"], "analysis-model"), (["analysis"], "analysis-model", False),
(["rag"], "qa-model"), (["rag"], "qa-model", True),
(["rag", "analysis"], "qa-model"), (["rag", "analysis"], "qa-model", True),
], ],
) )
def test_run_chat_gates_capability_vision_on_driving_model( def test_run_chat_gates_capability_vision_on_driving_model(
temp_db_path: Path, enabled, expected_model temp_db_path: Path, enabled, expected_model, expected_vision
): ):
"""Analysis-only chat runs on analysis.model; otherwise on qa.model. Every """Analysis-only chat runs on analysis.model; otherwise on qa.model. Every
attached capability's vision gate (its ``model``) tracks that one model.""" attached capability's vision gate tracks that one driving model."""
from haiku.rag.config.models import AppConfig, ModelConfig from haiku.rag.config.models import AppConfig, ModelConfig
config = AppConfig() config = AppConfig()
config.qa.model = ModelConfig(provider="openai", name="qa-model") config.qa.model = ModelConfig(provider="openai", name="qa-model", vision=True)
config.analysis.model = ModelConfig(provider="openai", name="analysis-model") config.analysis.model = ModelConfig(
provider="openai", name="analysis-model", vision=False
)
captured: dict[str, str] = {} captured: dict[str, str] = {}
def fake_get_model(model_config, _config): def fake_get_model(model_config, _config):
@ -81,7 +83,7 @@ def test_run_chat_gates_capability_vision_on_driving_model(
assert captured["name"] == expected_model assert captured["name"] == expected_model
attached = mock_app.call_args.kwargs["capabilities"] attached = mock_app.call_args.kwargs["capabilities"]
assert {capability.model.name for capability in attached} == {expected_model} assert {capability.vision for capability in attached} == {expected_vision}
def _make_mock_client(): def _make_mock_client():