Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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.
Expand All @@ -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"]
Expand Down Expand Up @@ -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 {}
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
BedrockAgentCoreControlHook,
BedrockAgentCoreHook,
BedrockAgentHook,
BedrockAgentRuntimeHook,
BedrockHook,
BedrockRuntimeHook,
)
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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",
[
Expand Down
2 changes: 0 additions & 2 deletions scripts/ci/prek/validate_operators_init_exemptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down