diff --git a/sdk/python/feast/infra/offline_stores/contrib/trino_offline_store/trino.py b/sdk/python/feast/infra/offline_stores/contrib/trino_offline_store/trino.py index c9d4119f94f..3e0ef87e908 100644 --- a/sdk/python/feast/infra/offline_stores/contrib/trino_offline_store/trino.py +++ b/sdk/python/feast/infra/offline_stores/contrib/trino_offline_store/trino.py @@ -127,7 +127,11 @@ def to_trino_auth(self): model_cls = CLASSES_BY_AUTH_TYPE[auth_type]["auth_model"] model = model_cls(**self.config) - return trino_auth_cls(**model.model_dump()) + kwargs = { + field: value.get_secret_value() if isinstance(value, SecretStr) else value + for field, value in model.model_dump().items() + } + return trino_auth_cls(**kwargs) class TrinoOfflineStoreConfig(FeastConfigBaseModel): diff --git a/sdk/python/tests/unit/infra/offline_stores/contrib/trino_offline_store/test_trino_auth.py b/sdk/python/tests/unit/infra/offline_stores/contrib/trino_offline_store/test_trino_auth.py new file mode 100644 index 00000000000..02a49e8924f --- /dev/null +++ b/sdk/python/tests/unit/infra/offline_stores/contrib/trino_offline_store/test_trino_auth.py @@ -0,0 +1,65 @@ +from pydantic import SecretStr +from trino.auth import BasicAuthentication, JWTAuthentication, OAuth2Authentication + +from feast.infra.offline_stores.contrib.trino_offline_store.trino import ( + CLASSES_BY_AUTH_TYPE, + AuthConfig, + FeastConfigBaseModel, +) + + +def test_jwt_auth_produces_plain_str_token(): + auth = AuthConfig(type="jwt", config={"token": "my-secret-token"}) + + trino_auth = auth.to_trino_auth() + + assert isinstance(trino_auth, JWTAuthentication) + assert trino_auth.token == "my-secret-token" + assert isinstance(trino_auth.token, str) + + +def test_oauth2_auth_unchanged(): + auth = AuthConfig(type="oauth2", config=None) + + trino_auth = auth.to_trino_auth() + + assert isinstance(trino_auth, OAuth2Authentication) + + +def test_basic_auth_with_plain_fields_unaffected(): + auth = AuthConfig(type="basic", config={"username": "alice", "password": "hunter2"}) + + trino_auth = auth.to_trino_auth() + + assert isinstance(trino_auth, BasicAuthentication) + assert trino_auth._username == "alice" + assert trino_auth._password == "hunter2" + + +class _MixedAuthModel(FeastConfigBaseModel): + username: str + token: SecretStr + + +class _MixedAuth: + def __init__(self, username: str, token: str): + self.username = username + self.token = token + + +def test_to_trino_auth_unwraps_only_secret_fields_in_mixed_model(monkeypatch): + monkeypatch.setitem( + CLASSES_BY_AUTH_TYPE, + "jwt", + {"auth_model": _MixedAuthModel, "trino_auth": _MixedAuth}, + ) + auth = AuthConfig( + type="jwt", config={"username": "alice", "token": "my-secret-token"} + ) + + trino_auth = auth.to_trino_auth() + + assert trino_auth.username == "alice" + assert trino_auth.token == "my-secret-token" + assert isinstance(trino_auth.username, str) + assert isinstance(trino_auth.token, str)