This is an automated email from the ASF dual-hosted git repository.

dabla 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 fd0f7e263bf Add auth_protocol support to SambaHook for Kerberos 
authentication #29590 (#64643)
fd0f7e263bf is described below

commit fd0f7e263bfa78c57c136d84bc677b4a2098b9fd
Author: Haseeb Malik <[email protected]>
AuthorDate: Sat Aug 8 02:40:02 2026 -0400

    Add auth_protocol support to SambaHook for Kerberos authentication #29590 
(#64643)
---
 providers/samba/docs/index.rst                     |  11 +-
 providers/samba/provider.yaml                      |  14 +++
 providers/samba/pyproject.toml                     |   4 +
 .../airflow/providers/samba/get_provider_info.py   |  11 +-
 .../src/airflow/providers/samba/hooks/samba.py     |  55 ++++++++--
 .../samba/tests/unit/samba/hooks/test_samba.py     | 117 ++++++++++++++++++++-
 uv.lock                                            |  13 ++-
 7 files changed, 210 insertions(+), 15 deletions(-)

diff --git a/providers/samba/docs/index.rst b/providers/samba/docs/index.rst
index 63ff1e5fb6f..35c48f21651 100644
--- a/providers/samba/docs/index.rst
+++ b/providers/samba/docs/index.rst
@@ -134,11 +134,12 @@ Install them when installing from PyPI. For example:
     pip install apache-airflow-providers-samba[google]
 
 
-==========  ===================================
-Extra       Dependencies
-==========  ===================================
-``google``  ``apache-airflow-providers-google``
-==========  ===================================
+============  ==========================================
+Extra         Dependencies
+============  ==========================================
+``google``    ``apache-airflow-providers-google``
+``kerberos``  ``krb5``, ``smbprotocol[kerberos]>=1.5.0``
+============  ==========================================
 
 Downloading official packages
 -----------------------------
diff --git a/providers/samba/provider.yaml b/providers/samba/provider.yaml
index 17cbc0e0243..564051413e2 100644
--- a/providers/samba/provider.yaml
+++ b/providers/samba/provider.yaml
@@ -98,6 +98,20 @@ connection-types:
         description: >-
           The share OS type (`posix` or `windows`).
           Used to determine the formatting of file and folder paths.
+      auth_protocol:
+        label: Auth Protocol
+        schema:
+          type:
+            - string
+            - 'null'
+          default: negotiate
+          enum:
+            - negotiate
+            - ntlm
+            - kerberos
+        description: >-
+          Must be one of: negotiate, ntlm, kerberos. With `kerberos`, the 
system's
+          ticket cache is used and login/password are optional.
     ui-field-behaviour:
       hidden-fields: []
       relabeling:
diff --git a/providers/samba/pyproject.toml b/providers/samba/pyproject.toml
index 268199698f8..d5105dd606d 100644
--- a/providers/samba/pyproject.toml
+++ b/providers/samba/pyproject.toml
@@ -70,6 +70,10 @@ dependencies = [
 "google" = [
     "apache-airflow-providers-google"
 ]
+"kerberos" = [
+    "krb5",
+    "smbprotocol[kerberos]>=1.5.0",
+]
 
 [dependency-groups]
 dev = [
diff --git a/providers/samba/src/airflow/providers/samba/get_provider_info.py 
b/providers/samba/src/airflow/providers/samba/get_provider_info.py
index fa6f98566f7..4f57b7a9ebc 100644
--- a/providers/samba/src/airflow/providers/samba/get_provider_info.py
+++ b/providers/samba/src/airflow/providers/samba/get_provider_info.py
@@ -53,7 +53,16 @@ def get_provider_info():
                         "label": "Share Type",
                         "schema": {"type": ["string", "null"], "default": 
"posix"},
                         "description": "The share OS type (`posix` or 
`windows`). Used to determine the formatting of file and folder paths.",
-                    }
+                    },
+                    "auth_protocol": {
+                        "label": "Auth Protocol",
+                        "schema": {
+                            "type": ["string", "null"],
+                            "default": "negotiate",
+                            "enum": ["negotiate", "ntlm", "kerberos"],
+                        },
+                        "description": "Must be one of: negotiate, ntlm, 
kerberos. With `kerberos`, the system's ticket cache is used and login/password 
are optional.",
+                    },
                 },
                 "ui-field-behaviour": {"hidden-fields": [], "relabeling": 
{"schema": "Share"}},
             }
diff --git a/providers/samba/src/airflow/providers/samba/hooks/samba.py 
b/providers/samba/src/airflow/providers/samba/hooks/samba.py
index c36f0b06dc9..1e6bf5a9cde 100644
--- a/providers/samba/src/airflow/providers/samba/hooks/samba.py
+++ b/providers/samba/src/airflow/providers/samba/hooks/samba.py
@@ -24,7 +24,7 @@ from typing import TYPE_CHECKING, Any, Literal
 
 import smbclient
 
-from airflow.providers.common.compat.sdk import BaseHook
+from airflow.providers.common.compat.sdk import 
AirflowOptionalProviderFeatureException, BaseHook
 
 if TYPE_CHECKING:
     import smbprotocol.connection
@@ -43,6 +43,8 @@ class SambaHook(BaseHook):
         the connection is used in its place.
     :param share_type:
         An optional share type name. If this is unset then it will assume a 
posix share type.
+    :param auth_protocol:
+        An optional authentication protocol. If this is unset then it defaults 
to negotiate.
     """
 
     conn_name_attr = "samba_conn_id"
@@ -50,22 +52,50 @@ class SambaHook(BaseHook):
     conn_type = "samba"
     hook_name = "Samba"
 
+    VALID_AUTH_PROTOCOLS = {"negotiate", "ntlm", "kerberos"}
+
     def __init__(
         self,
         samba_conn_id: str = default_conn_name,
         share: str | None = None,
         share_type: Literal["posix", "windows"] | None = None,
+        auth_protocol: Literal["negotiate", "ntlm", "kerberos"] | None = None,
     ) -> None:
         super().__init__()
         conn = self.get_connection(samba_conn_id)
+        extra = conn.extra_dejson
+
+        legacy_auth = extra.get("auth")
+        legacy_auth_protocol = (
+            legacy_auth
+            if isinstance(legacy_auth, str) and legacy_auth in 
self.VALID_AUTH_PROTOCOLS
+            else "negotiate"
+        )
+        self._auth_protocol: str = auth_protocol or extra.get("auth_protocol", 
legacy_auth_protocol)
+        if self._auth_protocol not in self.VALID_AUTH_PROTOCOLS:
+            raise ValueError(
+                f"Invalid auth_protocol '{self._auth_protocol}'. "
+                f"Must be one of {sorted(self.VALID_AUTH_PROTOCOLS)}."
+            )
 
-        if not conn.login:
+        uses_kerberos = self._auth_protocol == "kerberos"
+
+        if uses_kerberos:
+            try:
+                import krb5  # noqa: F401
+            except ImportError:
+                raise AirflowOptionalProviderFeatureException(
+                    "Kerberos authentication requires the 'krb5' package. "
+                    "Install it with: pip install 
'apache-airflow-providers-samba[kerberos]'"
+                )
+
+        if not conn.login and not uses_kerberos:
             self.log.info("Login not provided")
 
-        if not conn.password:
+        if not conn.password and not uses_kerberos:
             self.log.info("Password not provided")
 
-        self._share_type = share_type or conn.extra_dejson.get("share_type", 
"posix")
+        self._share_type = share_type or extra.get("share_type", "posix")
         if self._share_type not in {"posix", "windows"}:
             self._share_type = "posix"
             self.log.warning(
@@ -77,11 +107,12 @@ class SambaHook(BaseHook):
         self._host = conn.host
         self._share = share or conn.schema
         self._connection_cache = connection_cache
-        self._conn_kwargs = {
+        self._conn_kwargs: dict[str, Any] = {
             "username": conn.login,
             "password": conn.password,
             "port": conn.port or 445,
             "connection_cache": connection_cache,
+            "auth_protocol": self._auth_protocol,
         }
 
     def __enter__(self):
@@ -325,12 +356,22 @@ class SambaHook(BaseHook):
     def get_connection_form_widgets(cls) -> dict[str, Any]:
         """Return connection widgets to add to connection form."""
         from flask_babel import lazy_gettext
-        from wtforms import StringField
+        from wtforms import SelectField, StringField
 
         return {
             "share_type": StringField(
                 label=lazy_gettext("Share Type"),
                 description="The share OS type (`posix` or `windows`). Used to 
determine the formatting of file and folder paths.",
                 default="posix",
-            )
+            ),
+            "auth_protocol": SelectField(
+                label=lazy_gettext("Auth Protocol"),
+                description=(
+                    "Authentication protocol: `negotiate` (auto-select, 
default), "
+                    "`ntlm`, or `kerberos`. When using `kerberos`, the 
system's "
+                    "Kerberos ticket cache is used and username/password are 
optional."
+                ),
+                choices=["negotiate", "ntlm", "kerberos"],
+                default="negotiate",
+            ),
         }
diff --git a/providers/samba/tests/unit/samba/hooks/test_samba.py 
b/providers/samba/tests/unit/samba/hooks/test_samba.py
index 44d79c88303..06a5e3fb147 100644
--- a/providers/samba/tests/unit/samba/hooks/test_samba.py
+++ b/providers/samba/tests/unit/samba/hooks/test_samba.py
@@ -23,7 +23,10 @@ from unittest import mock
 import pytest
 
 from airflow.models import Connection
-from airflow.providers.common.compat.sdk import AirflowNotFoundException
+from airflow.providers.common.compat.sdk import (
+    AirflowNotFoundException,
+    AirflowOptionalProviderFeatureException,
+)
 from airflow.providers.samba.hooks.samba import SambaHook
 
 try:
@@ -63,6 +66,7 @@ class TestSambaHook:
                 "password": CONNECTION.password,
                 "port": 445,
                 "connection_cache": {},
+                "auth_protocol": "negotiate",
             }
             cache = kwargs.get("connection_cache")
             mock_connection = mock.Mock()
@@ -117,6 +121,7 @@ class TestSambaHook:
             "username": CONNECTION.login,
             "password": CONNECTION.password,
             "port": 445,
+            "auth_protocol": "negotiate",
         }
         with mock.patch("smbclient." + name) as p:
             kwargs = {}
@@ -191,6 +196,116 @@ class TestSambaHook:
         hook = SambaHook("samba_default", share_type=path_type)
         assert hook._join_path(path) == full_path
 
+    @mock.patch.dict("sys.modules", {"krb5": mock.MagicMock()})
+    @mock.patch("smbclient.register_session")
+    @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection")
+    def test_kerberos_auth_via_extra(self, get_conn_mock, register_session):
+        """Test that auth_protocol='kerberos' from extra is passed to 
smbclient."""
+        connection = Connection(
+            host="kerb-host.example.com",
+            schema="share",
+            extra='{"auth_protocol": "kerberos"}',
+        )
+        get_conn_mock.return_value = connection
+        register_session.return_value = None
+        with SambaHook("samba_default"):
+            _, kwargs = tuple(register_session.call_args_list[0])
+            assert kwargs["auth_protocol"] == "kerberos"
+            assert kwargs["username"] is None
+            assert kwargs["password"] is None
+
+    @mock.patch.dict("sys.modules", {"krb5": mock.MagicMock()})
+    @mock.patch("smbclient.register_session")
+    @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection")
+    def test_kerberos_auth_via_legacy_auth_key(self, get_conn_mock, 
register_session):
+        """Test backward compat: extra {"auth": "kerberos"} is recognized."""
+        connection = Connection(
+            host="kerb-host.example.com",
+            schema="share",
+            extra='{"auth": "kerberos"}',
+        )
+        get_conn_mock.return_value = connection
+        register_session.return_value = None
+        with SambaHook("samba_default"):
+            _, kwargs = tuple(register_session.call_args_list[0])
+            assert kwargs["auth_protocol"] == "kerberos"
+
+    @mock.patch.dict("sys.modules", {"krb5": mock.MagicMock()})
+    @mock.patch("smbclient.register_session")
+    @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection")
+    def test_kerberos_auth_via_constructor(self, get_conn_mock, 
register_session):
+        """Test that constructor auth_protocol overrides extra."""
+        connection = Connection(
+            host="kerb-host.example.com",
+            schema="share",
+            login="user",
+            password="pass",
+        )
+        get_conn_mock.return_value = connection
+        register_session.return_value = None
+        with SambaHook("samba_default", auth_protocol="kerberos"):
+            _, kwargs = tuple(register_session.call_args_list[0])
+            assert kwargs["auth_protocol"] == "kerberos"
+
+    @mock.patch("smbclient.register_session")
+    @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection")
+    def test_legacy_auth_key_ignored_when_invalid(self, get_conn_mock, 
register_session):
+        """Test that extra {"auth": "basic"} is ignored and defaults to 
negotiate."""
+        connection = Connection(
+            host="host",
+            schema="share",
+            login="user",
+            password="pass",
+            extra='{"auth": "basic"}',
+        )
+        get_conn_mock.return_value = connection
+        register_session.return_value = None
+        with SambaHook("samba_default"):
+            _, kwargs = tuple(register_session.call_args_list[0])
+            assert kwargs["auth_protocol"] == "negotiate"
+
+    @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection")
+    def test_invalid_auth_protocol_raises(self, get_conn_mock):
+        """Test that an invalid auth_protocol raises ValueError."""
+        connection = Connection(
+            host="host",
+            schema="share",
+            extra='{"auth_protocol": "invalid"}',
+        )
+        get_conn_mock.return_value = connection
+        with pytest.raises(ValueError, match="Invalid auth_protocol 
'invalid'"):
+            SambaHook("samba_default")
+
+    @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection")
+    def test_kerberos_without_dependency_raises(self, get_conn_mock):
+        """Test that kerberos auth without krb5 installed raises 
AirflowOptionalProviderFeatureException."""
+        connection = Connection(
+            host="kerb-host.example.com",
+            schema="share",
+            extra='{"auth_protocol": "kerberos"}',
+        )
+        get_conn_mock.return_value = connection
+        with mock.patch.dict("sys.modules", {"krb5": None}):
+            with pytest.raises(AirflowOptionalProviderFeatureException, 
match="krb5"):
+                SambaHook("samba_default")
+
+    @mock.patch("smbclient.register_session")
+    @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection")
+    def test_ntlm_auth(self, get_conn_mock, register_session):
+        """Test that auth_protocol='ntlm' is passed correctly."""
+        connection = Connection(
+            host="host",
+            schema="share",
+            login="user",
+            password="pass",
+            extra='{"auth_protocol": "ntlm"}',
+        )
+        get_conn_mock.return_value = connection
+        register_session.return_value = None
+        with SambaHook("samba_default"):
+            _, kwargs = tuple(register_session.call_args_list[0])
+            assert kwargs["auth_protocol"] == "ntlm"
+
     @mock.patch("airflow.providers.samba.hooks.samba.smbclient.open_file", 
return_value=mock.Mock())
     @mock.patch(f"{BASEHOOK_PATCH_PATH}.get_connection")
     def test_open_file(self, get_conn_mock, open_file_mock):
diff --git a/uv.lock b/uv.lock
index 17216a814ac..76b7408a20f 100644
--- a/uv.lock
+++ b/uv.lock
@@ -7491,6 +7491,10 @@ dependencies = [
 google = [
     { name = "apache-airflow-providers-google" },
 ]
+kerberos = [
+    { name = "krb5" },
+    { name = "smbprotocol", extra = ["kerberos"] },
+]
 
 [package.dev-dependencies]
 dev = [
@@ -7509,9 +7513,11 @@ requires-dist = [
     { name = "apache-airflow", editable = "." },
     { name = "apache-airflow-providers-common-compat", editable = 
"providers/common/compat" },
     { name = "apache-airflow-providers-google", marker = "extra == 'google'", 
editable = "providers/google" },
+    { name = "krb5", marker = "extra == 'kerberos'" },
     { name = "smbprotocol", specifier = ">=1.5.0" },
+    { name = "smbprotocol", extras = ["kerberos"], marker = "extra == 
'kerberos'", specifier = ">=1.5.0" },
 ]
-provides-extras = ["google"]
+provides-extras = ["google", "kerberos"]
 
 [package.metadata.requires-dev]
 dev = [
@@ -22199,6 +22205,11 @@ wheels = [
     { url = 
"https://files.pythonhosted.org/packages/76/d3/250e13e7b3473a0f441f4d680d511ad30ce432d0180115dec9840d88fe89/smbprotocol-1.17.0-py3-none-any.whl";,
 hash = 
"sha256:bd1abff5417f5af83ca516a64ab8e5acece3dbdcf58d4e5e23e47f5165a77349", size 
= 128208, upload-time = "2026-07-07T03:23:19.178Z" },
 ]
 
+[package.optional-dependencies]
+kerberos = [
+    { name = "pyspnego", extra = ["kerberos"] },
+]
+
 [[package]]
 name = "smmap"
 version = "5.0.3"

Reply via email to