Gate combined-chat capability vision on the driving model
This commit is contained in:
parent
597808c56e
commit
6c5bc0aae1
4 changed files with 40 additions and 19 deletions
|
|
@ -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
|
from haiku.rag.config.models import AppConfig, ModelConfig
|
||||||
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,8 +129,14 @@ 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,
|
||||||
) -> 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
|
||||||
|
model the hosting agent actually runs. Defaults to ``config.analysis.model``
|
||||||
|
(falling back to ``config.qa.model``).
|
||||||
|
"""
|
||||||
if config is None:
|
if config is None:
|
||||||
from haiku.rag.config import get_config
|
from haiku.rag.config import get_config
|
||||||
|
|
||||||
|
|
@ -141,7 +147,7 @@ def create_capability(
|
||||||
state_type=AnalysisState,
|
state_type=AnalysisState,
|
||||||
state_namespace=STATE_NAMESPACE,
|
state_namespace=STATE_NAMESPACE,
|
||||||
instruction_text=instructions(),
|
instruction_text=instructions(),
|
||||||
model=config.analysis.model or config.qa.model,
|
model=model or config.analysis.model or config.qa.model,
|
||||||
tool_names=_TOOL_NAMES,
|
tool_names=_TOOL_NAMES,
|
||||||
request_limit=request_limit,
|
request_limit=request_limit,
|
||||||
id=_CAPABILITY_ID,
|
id=_CAPABILITY_ID,
|
||||||
|
|
|
||||||
|
|
@ -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
|
from haiku.rag.config.models import AppConfig, ModelConfig
|
||||||
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,8 +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,
|
||||||
) -> 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
|
||||||
|
model the hosting agent actually runs. Defaults to ``config.qa.model``.
|
||||||
|
"""
|
||||||
if config is None:
|
if config is None:
|
||||||
from haiku.rag.config import get_config
|
from haiku.rag.config import get_config
|
||||||
|
|
||||||
|
|
@ -83,7 +88,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=config.qa.model,
|
model=model or config.qa.model,
|
||||||
tool_names=_TOOL_NAMES,
|
tool_names=_TOOL_NAMES,
|
||||||
request_limit=request_limit,
|
request_limit=request_limit,
|
||||||
id=_CAPABILITY_ID,
|
id=_CAPABILITY_ID,
|
||||||
|
|
|
||||||
|
|
@ -38,12 +38,24 @@ def run_chat(
|
||||||
capability_list = []
|
capability_list = []
|
||||||
defer_loading = len(enabled) > 1
|
defer_loading = len(enabled) > 1
|
||||||
|
|
||||||
|
# One agent drives every attached capability, so a capability's
|
||||||
|
# image-attachment gate must track that single model: analysis.model only
|
||||||
|
# when analysis runs alone, otherwise qa.model. Passing it to every
|
||||||
|
# capability keeps their vision flag aligned with the model actually running.
|
||||||
|
if "rag" not in enabled and "analysis" in enabled:
|
||||||
|
driving_model = config.analysis.model or config.qa.model
|
||||||
|
else:
|
||||||
|
driving_model = config.qa.model
|
||||||
|
|
||||||
if "rag" in enabled:
|
if "rag" in enabled:
|
||||||
from haiku.rag.capabilities.rag import create_capability
|
from haiku.rag.capabilities.rag import create_capability
|
||||||
|
|
||||||
capability_list.append(
|
capability_list.append(
|
||||||
create_capability(
|
create_capability(
|
||||||
db_path=db_path, config=config, defer_loading=defer_loading
|
db_path=db_path,
|
||||||
|
config=config,
|
||||||
|
defer_loading=defer_loading,
|
||||||
|
model=driving_model,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -52,18 +64,13 @@ def run_chat(
|
||||||
|
|
||||||
capability_list.append(
|
capability_list.append(
|
||||||
create_capability(
|
create_capability(
|
||||||
db_path=db_path, config=config, defer_loading=defer_loading
|
db_path=db_path,
|
||||||
|
config=config,
|
||||||
|
defer_loading=defer_loading,
|
||||||
|
model=driving_model,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
# Drive with the analysis model when analysis is the only capability, so the
|
|
||||||
# running model matches the one the analysis capability configures (including
|
|
||||||
# its vision flag). RAG runs on the QA model.
|
|
||||||
if "rag" not in enabled and "analysis" in enabled:
|
|
||||||
driving_model = config.analysis.model or config.qa.model
|
|
||||||
else:
|
|
||||||
driving_model = config.qa.model
|
|
||||||
|
|
||||||
app = ChatApp(
|
app = ChatApp(
|
||||||
db_path,
|
db_path,
|
||||||
capabilities=capability_list,
|
capabilities=capability_list,
|
||||||
|
|
|
||||||
|
|
@ -54,10 +54,11 @@ def test_run_chat_defers_multiple_capabilities(temp_db_path: Path):
|
||||||
(["rag", "analysis"], "qa-model"),
|
(["rag", "analysis"], "qa-model"),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_run_chat_drives_analysis_only_with_analysis_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
|
||||||
):
|
):
|
||||||
"""Analysis-only chat runs on analysis.model; otherwise on qa.model."""
|
"""Analysis-only chat runs on analysis.model; otherwise on qa.model. Every
|
||||||
|
attached capability's vision gate (its ``model``) tracks that one model."""
|
||||||
from haiku.rag.config.models import AppConfig, ModelConfig
|
from haiku.rag.config.models import AppConfig, ModelConfig
|
||||||
|
|
||||||
config = AppConfig()
|
config = AppConfig()
|
||||||
|
|
@ -70,7 +71,7 @@ def test_run_chat_drives_analysis_only_with_analysis_model(
|
||||||
return "resolved-model"
|
return "resolved-model"
|
||||||
|
|
||||||
with (
|
with (
|
||||||
patch("haiku.rag.chat.app.ChatApp"),
|
patch("haiku.rag.chat.app.ChatApp") as mock_app,
|
||||||
patch("haiku.rag.config.get_config", return_value=config),
|
patch("haiku.rag.config.get_config", return_value=config),
|
||||||
patch("haiku.rag.utils.get_model", side_effect=fake_get_model),
|
patch("haiku.rag.utils.get_model", side_effect=fake_get_model),
|
||||||
):
|
):
|
||||||
|
|
@ -79,6 +80,8 @@ def test_run_chat_drives_analysis_only_with_analysis_model(
|
||||||
run_chat(db_path=temp_db_path, capabilities=enabled)
|
run_chat(db_path=temp_db_path, capabilities=enabled)
|
||||||
|
|
||||||
assert captured["name"] == expected_model
|
assert captured["name"] == expected_model
|
||||||
|
attached = mock_app.call_args.kwargs["capabilities"]
|
||||||
|
assert {capability.model.name for capability in attached} == {expected_model}
|
||||||
|
|
||||||
|
|
||||||
def _make_mock_client():
|
def _make_mock_client():
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue