Add user_id to assist pipelines (#182750)

This commit is contained in:
Felix Schneider
2026-09-27 10:00:50 +02:00
committed by GitHub
parent a9196e0470
commit 14a2a0f486
7 changed files with 232 additions and 7 deletions
@@ -56,6 +56,8 @@ class Pipeline:
wake_word_id: str | None
prefer_local_intents: bool = False
user_id: str | None = None
id: str = field(default_factory=ulid_util.ulid_now)
@classmethod
@@ -79,6 +81,7 @@ class Pipeline:
wake_word_entity=data["wake_word_entity"],
wake_word_id=data["wake_word_id"],
prefer_local_intents=data.get("prefer_local_intents", False),
user_id=data.get("user_id"),
)
def to_json(self) -> dict[str, Any]:
@@ -97,6 +100,7 @@ class Pipeline:
"wake_word_entity": self.wake_word_entity,
"wake_word_id": self.wake_word_id,
"prefer_local_intents": self.prefer_local_intents,
"user_id": self.user_id,
}
@@ -81,9 +81,20 @@ PIPELINE_FIELDS: VolDictType = {
probatio.Required("wake_word_id"): probatio.Any(str, None),
probatio.Optional("prefer_local_intents"): bool,
probatio.Optional("acknowledge_media_id"): str,
probatio.Optional("user_id"): probatio.Any(str, None),
}
async def _async_validate_user(hass: HomeAssistant, data: dict[str, Any]) -> None:
"""Validate that the user a pipeline acts as exists and can act."""
user_id: str | None = data.get("user_id")
if user_id is None:
return
user = await hass.auth.async_get_user(user_id)
if user is None or not user.is_active:
raise probatio.Invalid(f"Unknown user {user_id}")
@callback
def _async_resolve_default_pipeline_settings(
hass: HomeAssistant,
@@ -291,6 +302,7 @@ async def async_update_pipeline(
wake_word_entity: str | UndefinedType | None = UNDEFINED,
wake_word_id: str | UndefinedType | None = UNDEFINED,
prefer_local_intents: bool | UndefinedType = UNDEFINED,
user_id: str | UndefinedType | None = UNDEFINED,
) -> None:
"""Update a pipeline."""
pipeline_data = hass.data[KEY_ASSIST_PIPELINE]
@@ -315,6 +327,7 @@ async def async_update_pipeline(
("wake_word_entity", wake_word_entity),
("wake_word_id", wake_word_id),
("prefer_local_intents", prefer_local_intents),
("user_id", user_id),
)
if val is not UNDEFINED
}
@@ -361,6 +374,7 @@ class PipelineStorageCollection(
async def _process_create_data(self, data: dict) -> dict:
"""Validate the config is valid."""
validated_data: dict = validate_language(data)
await _async_validate_user(self.hass, validated_data)
return validated_data
@callback
@@ -373,6 +387,7 @@ class PipelineStorageCollection(
async def _update_data(self, item: Pipeline, update_data: dict) -> Pipeline:
"""Return a new updated item."""
update_data = validate_language(update_data)
await _async_validate_user(self.hass, update_data)
return Pipeline(id=item.id, **update_data)
@override
@@ -237,10 +237,43 @@ class PipelineRun:
self._unregister()
raise
async def _async_apply_pipeline_user(self) -> None:
"""Give the run the identity its pipeline acts as.
A run started by a satellite has no user of its own, so anything acting on
behalf of the caller cannot tell whose request it is: the pipeline lends it
one.
A run started by a user keeps their own identity, unless that user is an
administrator: they may act as anyone anyway, which also makes the setting
testable from the pipeline debug page.
"""
user_id: str | None = self.pipeline.user_id
if user_id is None or user_id == self.context.user_id:
return
if (caller_id := self.context.user_id) is not None:
caller = await self.hass.auth.async_get_user(caller_id)
if caller is None or not caller.is_admin:
return
# the stored user can have been deactivated or removed since, and this path
# builds a context outside the authentication boundary, so a stale identity is
# not handed the policies it no longer has
user = await self.hass.auth.async_get_user(user_id)
if user is None or not user.is_active:
return
self.context = Context(
user_id=user_id,
parent_id=self.context.parent_id,
id=self.context.id,
)
async def async_execute(
self, pipeline_input: PipelineInput, *, validate: bool = False
) -> None:
"""Run the pipeline processor with the Home Assistant lifecycle."""
await self._async_apply_pipeline_user()
request = pipeline_input.create_processor_request()
validation_error: PipelineError | None = None
self._set_request_identity(request)
@@ -1,7 +1,7 @@
"""Websocket tests for Voice Assistant integration."""
from collections.abc import AsyncGenerator, Generator
from dataclasses import FrozenInstanceError
from dataclasses import FrozenInstanceError, replace
from pathlib import Path
from typing import Any
from unittest.mock import ANY, AsyncMock, Mock, patch
@@ -68,7 +68,7 @@ from .conftest import (
make_10ms_chunk,
)
from tests.common import MockConfigEntry, async_mock_service, flush_store
from tests.common import MockConfigEntry, MockUser, async_mock_service, flush_store
from tests.typing import ClientSessionGenerator, WebSocketGenerator
@@ -561,7 +561,7 @@ async def test_default_pipeline_unsupported_tts_language(
async def test_update_pipeline(
hass: HomeAssistant, hass_storage: dict[str, Any]
hass: HomeAssistant, hass_storage: dict[str, Any], hass_admin_user: MockUser
) -> None:
"""Test async_update_pipeline."""
assert await async_setup_component(hass, DOMAIN, {})
@@ -600,6 +600,7 @@ async def test_update_pipeline(
tts_voice="test_voice",
wake_word_entity="wake_work.test_1",
wake_word_id="wake_word_id_1",
user_id=hass_admin_user.id,
)
pipelines = async_get_pipelines(hass)
@@ -619,6 +620,7 @@ async def test_update_pipeline(
tts_voice="test_voice",
wake_word_entity="wake_work.test_1",
wake_word_id="wake_word_id_1",
user_id=hass_admin_user.id,
)
]
assert len(hass_storage[STORAGE_KEY]["data"]["items"]) == 1
@@ -636,6 +638,7 @@ async def test_update_pipeline(
"wake_word_entity": "wake_work.test_1",
"wake_word_id": "wake_word_id_1",
"prefer_local_intents": False,
"user_id": hass_admin_user.id,
}
await async_update_pipeline(
@@ -663,6 +666,7 @@ async def test_update_pipeline(
tts_voice="test_voice",
wake_word_entity="wake_work.test_1",
wake_word_id="wake_word_id_1",
user_id=hass_admin_user.id,
)
]
assert len(hass_storage[STORAGE_KEY]["data"]["items"]) == 1
@@ -680,6 +684,7 @@ async def test_update_pipeline(
"wake_word_entity": "wake_work.test_1",
"wake_word_id": "wake_word_id_1",
"prefer_local_intents": False,
"user_id": hass_admin_user.id,
}
@@ -2494,3 +2499,155 @@ async def test_pipeline_error_before_tts_does_not_leak_result_stream(
)
assert len(hass.data[tts.DATA_TTS_MANAGER].token_to_stream) == 0
async def _run_with_pipeline_user(
hass: HomeAssistant,
pipeline_user_id: str,
context: Context,
mock_chat_session: chat_session.ChatSession,
) -> Context:
"""Run a pipeline that acts as a user and return the context it ran with."""
pipeline = replace(
assist_pipeline.pipeline.async_get_pipeline(hass), user_id=pipeline_user_id
)
processor = Mock(
response_audio=None,
supports_streaming_response=False,
async_validate=AsyncMock(),
async_execute=AsyncMock(),
invalidate=Mock(),
cleanup=Mock(),
)
with patch(
"homeassistant.components.assist_pipeline.run._create_pipeline_processor",
return_value=processor,
):
pipeline_input = assist_pipeline.pipeline.PipelineInput(
intent_input="test input",
session=mock_chat_session,
run=assist_pipeline.pipeline.PipelineRun(
hass,
context=context,
pipeline=pipeline,
start_stage=assist_pipeline.PipelineStage.INTENT,
end_stage=assist_pipeline.PipelineStage.INTENT,
event_callback=lambda event: None,
),
)
await pipeline_input.execute()
return pipeline_input.run.context
async def test_pipeline_user_id_sets_context_user(
hass: HomeAssistant,
init_components: None,
hass_admin_user: MockUser,
mock_chat_session: chat_session.ChatSession,
) -> None:
"""Test a pipeline lends its user to a run that has none."""
context = Context()
run_context = await _run_with_pipeline_user(
hass, hass_admin_user.id, context, mock_chat_session
)
assert run_context.user_id == hass_admin_user.id
assert run_context.id == context.id
async def test_pipeline_user_id_keeps_non_admin_caller(
hass: HomeAssistant,
init_components: None,
hass_admin_user: MockUser,
hass_read_only_user: MockUser,
mock_chat_session: chat_session.ChatSession,
) -> None:
"""Test a pipeline never hands a caller permissions they do not have."""
run_context = await _run_with_pipeline_user(
hass,
hass_admin_user.id,
Context(user_id=hass_read_only_user.id),
mock_chat_session,
)
assert run_context.user_id == hass_read_only_user.id
async def test_pipeline_user_id_applies_for_admin_caller(
hass: HomeAssistant,
init_components: None,
hass_admin_user: MockUser,
hass_read_only_user: MockUser,
mock_chat_session: chat_session.ChatSession,
) -> None:
"""Test an administrator runs as the user the pipeline acts as."""
run_context = await _run_with_pipeline_user(
hass,
hass_read_only_user.id,
Context(user_id=hass_admin_user.id),
mock_chat_session,
)
assert run_context.user_id == hass_read_only_user.id
async def test_pipeline_user_id_deactivated_user(
hass: HomeAssistant,
init_components: None,
hass_read_only_user: MockUser,
mock_chat_session: chat_session.ChatSession,
) -> None:
"""Test a run does not act as a user that was deactivated after being set."""
await hass.auth.async_deactivate_user(hass_read_only_user)
run_context = await _run_with_pipeline_user(
hass, hass_read_only_user.id, Context(), mock_chat_session
)
assert run_context.user_id is None
async def test_pipeline_unknown_user_id(
hass: HomeAssistant,
init_components: None,
hass_admin_user: MockUser,
) -> None:
"""Test a pipeline cannot be stored with a user that does not exist."""
pipeline_store = hass.data[
assist_pipeline.pipeline.KEY_ASSIST_PIPELINE
].pipeline_store
settings = {
"name": "Test",
"language": "en-US",
"conversation_engine": "test agent",
"conversation_language": "en-US",
"tts_engine": "test tts",
"tts_language": "en-US",
"tts_voice": "test voice",
"stt_engine": "test stt",
"stt_language": "en-US",
"wake_word_entity": None,
"wake_word_id": None,
}
with pytest.raises(probatio.Invalid):
await pipeline_store.async_create_item(settings | {"user_id": "does-not-exist"})
pipeline = await pipeline_store.async_create_item(
settings | {"user_id": hass_admin_user.id}
)
assert pipeline.user_id == hass_admin_user.id
with pytest.raises(probatio.Invalid):
await pipeline_store.async_update_item(
pipeline.id, settings | {"user_id": "does-not-exist"}
)
# a deactivated user is an identity a run must not be given either
await hass.auth.async_deactivate_user(hass_admin_user)
with pytest.raises(probatio.Invalid):
await pipeline_store.async_create_item(
settings | {"user_id": hass_admin_user.id}
)
@@ -953,6 +953,7 @@ async def test_add_pipeline(
"wake_word_entity": "wakeword_entity_1",
"wake_word_id": "wakeword_id_1",
"prefer_local_intents": True,
"user_id": None,
}
assert len(pipeline_store.data) == 2
@@ -1159,6 +1160,7 @@ async def test_get_pipeline(
"wake_word_entity": None,
"wake_word_id": None,
"prefer_local_intents": False,
"user_id": None,
}
# Get conversation agent as pipeline
@@ -1185,6 +1187,7 @@ async def test_get_pipeline(
"wake_word_entity": None,
"wake_word_id": None,
"prefer_local_intents": False,
"user_id": None,
}
await client.send_json_auto_id(
@@ -1215,6 +1218,7 @@ async def test_get_pipeline(
"wake_word_entity": "wakeword_entity_1",
"wake_word_id": "wakeword_id_1",
"prefer_local_intents": False,
"user_id": None,
}
)
msg = await client.receive_json()
@@ -1244,6 +1248,7 @@ async def test_get_pipeline(
"wake_word_entity": "wakeword_entity_1",
"wake_word_id": "wakeword_id_1",
"prefer_local_intents": False,
"user_id": None,
}
@@ -1272,6 +1277,7 @@ async def test_list_pipelines(
"wake_word_entity": None,
"wake_word_id": None,
"prefer_local_intents": False,
"user_id": None,
}
],
"preferred_pipeline": ANY,
@@ -1364,6 +1370,7 @@ async def test_update_pipeline(
"wake_word_entity": "new_wakeword_entity",
"wake_word_id": "new_wakeword_id",
"prefer_local_intents": False,
"user_id": None,
}
assert len(pipeline_store.data) == 2
@@ -1416,6 +1423,7 @@ async def test_update_pipeline(
"wake_word_entity": None,
"wake_word_id": None,
"prefer_local_intents": False,
"user_id": None,
}
pipeline = pipeline_store.data[pipeline_id]
+6 -2
View File
@@ -157,8 +157,12 @@ async def test_migrating_pipelines(
)
assert hass_storage[STORAGE_KEY]["data"]["items"][0]["wake_word_entity"] is None
assert hass_storage[STORAGE_KEY]["data"]["items"][0]["wake_word_id"] is None
assert hass_storage[STORAGE_KEY]["data"]["items"][1] == PIPELINE_DATA["items"][1]
assert hass_storage[STORAGE_KEY]["data"]["items"][2] == PIPELINE_DATA["items"][2]
assert hass_storage[STORAGE_KEY]["data"]["items"][1] == PIPELINE_DATA["items"][
1
] | {"user_id": None}
assert hass_storage[STORAGE_KEY]["data"]["items"][2] == PIPELINE_DATA["items"][
2
] | {"user_id": None}
@pytest.fixture(name="setup_stt")
+6 -2
View File
@@ -478,8 +478,12 @@ async def test_migrating_pipelines(
)
assert hass_storage[STORAGE_KEY]["data"]["items"][0]["wake_word_entity"] is None
assert hass_storage[STORAGE_KEY]["data"]["items"][0]["wake_word_id"] is None
assert hass_storage[STORAGE_KEY]["data"]["items"][1] == PIPELINE_DATA["items"][1]
assert hass_storage[STORAGE_KEY]["data"]["items"][2] == PIPELINE_DATA["items"][2]
assert hass_storage[STORAGE_KEY]["data"]["items"][1] == PIPELINE_DATA["items"][
1
] | {"user_id": None}
assert hass_storage[STORAGE_KEY]["data"]["items"][2] == PIPELINE_DATA["items"][
2
] | {"user_id": None}
@pytest.mark.parametrize(