diff --git a/providers/cohere/src/airflow/providers/cohere/operators/embedding.py b/providers/cohere/src/airflow/providers/cohere/operators/embedding.py index 62270e7570857..312f9237c5517 100644 --- a/providers/cohere/src/airflow/providers/cohere/operators/embedding.py +++ b/providers/cohere/src/airflow/providers/cohere/operators/embedding.py @@ -68,8 +68,6 @@ def __init__( **kwargs: Any, ): super().__init__(**kwargs) - if isinstance(input_text, str): - input_text = [input_text] self.conn_id = conn_id self.input_text = input_text self.timeout = timeout @@ -88,7 +86,10 @@ def hook(self) -> CohereHook: def execute(self, context: Context) -> list[list[float]]: """Embed texts using Cohere embed services.""" - embedding_response = self.hook.create_embeddings(self.input_text) + # input_text is a template field, so normalize a single string into a list here + # (after rendering) rather than in __init__, where it would act on the un-rendered value. + input_text = [self.input_text] if isinstance(self.input_text, str) else self.input_text + embedding_response = self.hook.create_embeddings(input_text) # NOTE: Return type `EmbedByTypeResponseEmbeddings` was removed temporarily due to limitations # in XCom serialization/deserialization of complex types like Cohere embeddings and Pydantic models. # diff --git a/providers/cohere/tests/unit/cohere/operators/test_embedding.py b/providers/cohere/tests/unit/cohere/operators/test_embedding.py index 1b6bf810e93d7..be971bde88399 100644 --- a/providers/cohere/tests/unit/cohere/operators/test_embedding.py +++ b/providers/cohere/tests/unit/cohere/operators/test_embedding.py @@ -56,3 +56,29 @@ def test_cohere_embedding_operator(cohere_client, get_connection): val = op.execute(context={}) cohere_client.assert_called_once_with(api_key=api_key, base_url=base_url, timeout=timeout) assert val == embedded_obj + + +@patch("airflow.providers.cohere.operators.embedding.CohereEmbeddingOperator.hook") +def test_single_string_input_text_normalized_at_execute(mock_hook): + """ + A single-string ``input_text`` is wrapped into a list in ``execute()`` (after templating), + not in ``__init__``. ``input_text`` is a template field, so wrapping it in the constructor + would operate on the un-rendered Jinja expression. + """ + op = CohereEmbeddingOperator(task_id="embed", conn_id="some_conn", input_text="single text") + + # __init__ keeps the value verbatim; the list normalization is deferred to execute(). + assert op.input_text == "single text" + + op.execute(context={}) + mock_hook.create_embeddings.assert_called_once_with(["single text"]) + + +@patch("airflow.providers.cohere.operators.embedding.CohereEmbeddingOperator.hook") +def test_list_input_text_passed_through_at_execute(mock_hook): + """A list ``input_text`` is passed to the hook unchanged.""" + texts = ["first", "second"] + op = CohereEmbeddingOperator(task_id="embed", conn_id="some_conn", input_text=texts) + + op.execute(context={}) + mock_hook.create_embeddings.assert_called_once_with(texts) diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index db9e1ad7ce87d..7f5f7bc618cd1 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -27,7 +27,6 @@ providers/apache/hive/src/airflow/providers/apache/hive/sensors/named_hive_parti providers/apache/kafka/src/airflow/providers/apache/kafka/operators/produce.py::ProduceToTopicOperator providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/kueue.py::KubernetesInstallKueueOperator providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/pod.py::KubernetesPodOperator -providers/cohere/src/airflow/providers/cohere/operators/embedding.py::CohereEmbeddingOperator providers/databricks/src/airflow/providers/databricks/operators/databricks_repos.py::DatabricksReposCreateOperator providers/databricks/src/airflow/providers/databricks/operators/databricks_repos.py::DatabricksReposDeleteOperator providers/databricks/src/airflow/providers/databricks/operators/databricks_repos.py::DatabricksReposUpdateOperator