diff --git a/providers/amazon/src/airflow/providers/amazon/aws/operators/bedrock.py b/providers/amazon/src/airflow/providers/amazon/aws/operators/bedrock.py index 045f0f7192c69..71fa6bf1dd110 100644 --- a/providers/amazon/src/airflow/providers/amazon/aws/operators/bedrock.py +++ b/providers/amazon/src/airflow/providers/amazon/aws/operators/bedrock.py @@ -731,10 +731,6 @@ def __init__( self.storage_config = storage_config self.create_knowledge_base_kwargs = create_knowledge_base_kwargs or {} self.embedding_model_arn = embedding_model_arn - self.knowledge_base_config = { - "type": "VECTOR", - "vectorKnowledgeBaseConfiguration": {"embeddingModelArn": self.embedding_model_arn}, - } self.wait_for_indexing = wait_for_indexing self.indexing_error_retry_delay = indexing_error_retry_delay self.indexing_error_max_attempts = indexing_error_max_attempts @@ -754,6 +750,11 @@ def execute_complete(self, context: Context, event: dict[str, Any] | None = None return validated_event["knowledge_base_id"] def execute(self, context: Context) -> str: + knowledge_base_config = { + "type": "VECTOR", + "vectorKnowledgeBaseConfiguration": {"embeddingModelArn": self.embedding_model_arn}, + } + def _create_kb(): # This API call will return the following if the index has not completed, but there is no apparent # way to check the state of the index beforehand, so retry on index failure if set to do so. @@ -764,7 +765,7 @@ def _create_kb(): return self.hook.conn.create_knowledge_base( name=self.name, roleArn=self.role_arn, - knowledgeBaseConfiguration=self.knowledge_base_config, + knowledgeBaseConfiguration=knowledge_base_config, storageConfiguration=self.storage_config, **self.create_knowledge_base_kwargs, )["knowledgeBase"]["knowledgeBaseId"] @@ -1065,10 +1066,10 @@ def __init__( ): super().__init__(**kwargs) self.input = input + self.source_type = source_type + self.model_arn = model_arn self.prompt_template = prompt_template - self.source_type = source_type.upper() self.knowledge_base_id = knowledge_base_id - self.model_arn = model_arn self.vector_search_config = vector_search_config self.sources = sources self.rag_kwargs = rag_kwargs or {} @@ -1132,6 +1133,7 @@ def build_rag_config(self) -> dict[str, Any]: return result def execute(self, context: Context) -> Any: + self.source_type = self.source_type.upper() self.validate_inputs() result = self.hook.conn.retrieve_and_generate( diff --git a/providers/amazon/tests/unit/amazon/aws/operators/test_bedrock.py b/providers/amazon/tests/unit/amazon/aws/operators/test_bedrock.py index 5085ea08453f0..cae3468a9d9ff 100644 --- a/providers/amazon/tests/unit/amazon/aws/operators/test_bedrock.py +++ b/providers/amazon/tests/unit/amazon/aws/operators/test_bedrock.py @@ -31,6 +31,7 @@ BedrockAgentCoreControlHook, BedrockAgentCoreHook, BedrockAgentHook, + BedrockAgentRuntimeHook, BedrockHook, BedrockRuntimeHook, ) @@ -621,6 +622,25 @@ def test_returns_id(self, mock_conn): assert result == self.KNOWLEDGE_BASE_ID + def test_knowledge_base_config_uses_rendered_embedding_model_arn(self, mock_conn): + """The knowledgeBaseConfiguration must be built from embedding_model_arn as it + stands at execute() time, since template rendering happens after __init__.""" + self.operator.wait_for_completion = False + rendered_arn = "arn:aws:bedrock:us-east-1::foundation-model/rendered-model" + self.operator.embedding_model_arn = rendered_arn + + self.operator.execute({}) + + mock_conn.create_knowledge_base.assert_called_once_with( + name=self.KNOWLEDGE_BASE_ID, + roleArn="role-arn", + knowledgeBaseConfiguration={ + "type": "VECTOR", + "vectorKnowledgeBaseConfiguration": {"embeddingModelArn": rendered_arn}, + }, + storageConfiguration=self.operator.storage_config, + ) + def test_template_fields(self): validate_template_fields(self.operator) @@ -941,6 +961,26 @@ def test_input_validation( with pytest.raises(AttributeError): op.validate_inputs() + @mock.patch.object(BedrockAgentRuntimeHook, "conn", new_callable=mock.PropertyMock) + def test_source_type_normalized_in_execute_not_init(self, mock_conn): + """source_type upper-casing must happen in execute(), after templating, not in __init__.""" + mock_client = mock.MagicMock() + mock_client.retrieve_and_generate.return_value = {"output": {"text": "answer"}, "citations": []} + mock_conn.return_value = mock_client + + op = BedrockRaGOperator( + task_id="test_rag", + input="some text prompt", + source_type="knowledge_base", + model_arn=self.MODEL_ARN, + knowledge_base_id=self.KNOWLEDGE_BASE_ID, + ) + assert op.source_type == "knowledge_base" + + op.execute({}) + + assert op.source_type == "KNOWLEDGE_BASE" + @pytest.mark.parametrize( "prompt_template", [ diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index 22a82b96ff594..7ff6497cf73f0 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -7,8 +7,6 @@ # execute()) MUST remove its entry in the same PR — the hook fails on stale entries. # Burn-down tracked at https://github.com/apache/airflow/issues/70296 providers/amazon/src/airflow/providers/amazon/aws/operators/appflow.py::AppflowBaseOperator -providers/amazon/src/airflow/providers/amazon/aws/operators/bedrock.py::BedrockCreateKnowledgeBaseOperator -providers/amazon/src/airflow/providers/amazon/aws/operators/bedrock.py::BedrockRaGOperator providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py::DmsModifyTaskOperator providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py::DmsStartReplicationOperator providers/amazon/src/airflow/providers/amazon/aws/operators/ecs.py::EcsRunTaskOperator