This is an automated email from the ASF dual-hosted git repository.
shahar1 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 446378da69a Move template-field validation/transformation out of init
for Amazon Bedrock operators (#70306)
446378da69a is described below
commit 446378da69aa58e3bbcc6cd4c794f9ad14f6f017
Author: Mustafa Avcu <[email protected]>
AuthorDate: Fri Jul 24 01:26:25 2026 -0400
Move template-field validation/transformation out of init for Amazon
Bedrock operators (#70306)
---
.../providers/amazon/aws/operators/bedrock.py | 16 +++++----
.../unit/amazon/aws/operators/test_bedrock.py | 40 ++++++++++++++++++++++
.../ci/prek/validate_operators_init_exemptions.txt | 2 --
3 files changed, 49 insertions(+), 9 deletions(-)
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 045f0f7192c..71fa6bf1dd1 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 @@ class
BedrockCreateKnowledgeBaseOperator(AwsBaseOperator[BedrockAgentHook]):
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 @@ class
BedrockCreateKnowledgeBaseOperator(AwsBaseOperator[BedrockAgentHook]):
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 @@ class
BedrockCreateKnowledgeBaseOperator(AwsBaseOperator[BedrockAgentHook]):
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 @@ class
BedrockRaGOperator(AwsBaseOperator[BedrockAgentRuntimeHook]):
):
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 @@ class
BedrockRaGOperator(AwsBaseOperator[BedrockAgentRuntimeHook]):
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 5085ea08453..cae3468a9d9 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 @@ from airflow.providers.amazon.aws.hooks.bedrock import (
BedrockAgentCoreControlHook,
BedrockAgentCoreHook,
BedrockAgentHook,
+ BedrockAgentRuntimeHook,
BedrockHook,
BedrockRuntimeHook,
)
@@ -621,6 +622,25 @@ class TestBedrockCreateKnowledgeBaseOperator:
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 @@ class TestBedrockRaGOperator:
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 22a82b96ff5..7ff6497cf73 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