diff --git a/homeassistant/components/assist_pipeline/models.py b/homeassistant/components/assist_pipeline/models.py index c41db1972e5f..f9246ccdcb9a 100644 --- a/homeassistant/components/assist_pipeline/models.py +++ b/homeassistant/components/assist_pipeline/models.py @@ -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, } diff --git a/homeassistant/components/assist_pipeline/pipeline.py b/homeassistant/components/assist_pipeline/pipeline.py index 235634a6be40..f630f6a01ef4 100644 --- a/homeassistant/components/assist_pipeline/pipeline.py +++ b/homeassistant/components/assist_pipeline/pipeline.py @@ -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 diff --git a/homeassistant/components/assist_pipeline/run.py b/homeassistant/components/assist_pipeline/run.py index f45245de59f2..74438f3a9e79 100644 --- a/homeassistant/components/assist_pipeline/run.py +++ b/homeassistant/components/assist_pipeline/run.py @@ -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) diff --git a/tests/components/assist_pipeline/test_pipeline.py b/tests/components/assist_pipeline/test_pipeline.py index c24cab61953a..742e403a045c 100644 --- a/tests/components/assist_pipeline/test_pipeline.py +++ b/tests/components/assist_pipeline/test_pipeline.py @@ -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} + ) diff --git a/tests/components/assist_pipeline/test_websocket.py b/tests/components/assist_pipeline/test_websocket.py index ed9c23bf3dbd..0367140ac31d 100644 --- a/tests/components/assist_pipeline/test_websocket.py +++ b/tests/components/assist_pipeline/test_websocket.py @@ -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] diff --git a/tests/components/cloud/test_stt.py b/tests/components/cloud/test_stt.py index 49df346d2b7f..48553688514b 100644 --- a/tests/components/cloud/test_stt.py +++ b/tests/components/cloud/test_stt.py @@ -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") diff --git a/tests/components/cloud/test_tts.py b/tests/components/cloud/test_tts.py index a95196d5136e..cd25426bc0cd 100644 --- a/tests/components/cloud/test_tts.py +++ b/tests/components/cloud/test_tts.py @@ -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(