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
36 changes: 36 additions & 0 deletions providers/trino/src/airflow/providers/trino/hooks/trino.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
from collections.abc import Iterable, Mapping
from pathlib import Path
from typing import TYPE_CHECKING, Any, TypeVar
from urllib.parse import quote_plus, urlencode

import trino
from trino.exceptions import DatabaseError
Expand Down Expand Up @@ -322,3 +323,38 @@ def get_openlineage_database_dialect(self, _):
def get_openlineage_default_schema(self):
"""Return Trino default schema."""
return trino.constants.DEFAULT_SCHEMA

def get_uri(self) -> str:
"""Return the Trino URI for the connection."""
conn = self.connection
uri = "trino://"

auth_part = ""
if conn.login:
auth_part = quote_plus(conn.login)
if conn.password:
auth_part = f"{auth_part}:{quote_plus(conn.password)}"
auth_part = f"{auth_part}@"

host_part = conn.host or "localhost"
if conn.port:
host_part = f"{host_part}:{conn.port}"

schema_part = ""
if conn.schema:
schema_part = f"/{quote_plus(conn.schema)}"
extra_schema = conn.extra_dejson.get("schema")
if extra_schema:
schema_part = f"{schema_part}/{quote_plus(extra_schema)}"
Comment thread
jason810496 marked this conversation as resolved.

uri = f"{uri}{auth_part}{host_part}{schema_part}"

extra = conn.extra_dejson.copy()
if "schema" in extra:
extra.pop("schema")

query_params = {k: str(v) for k, v in extra.items() if v is not None}
if query_params:
uri = f"{uri}?{urlencode(query_params)}"
Comment thread
jason810496 marked this conversation as resolved.

return uri
52 changes: 52 additions & 0 deletions providers/trino/tests/unit/trino/hooks/test_trino.py
Original file line number Diff line number Diff line change
Expand Up @@ -444,3 +444,55 @@ def get_first(self, *_):
},
)
]


@pytest.mark.parametrize(
"conn_params, expected_uri",
[
(
{"login": "user", "password": "pass", "host": "localhost", "port": 8080, "schema": "hive"},
"trino://user:pass@localhost:8080/hive",
),
(
{
"login": "user",
"password": "pass",
"host": "localhost",
"port": 8080,
"schema": "hive",
"extra": json.dumps({"schema": "sales"}),
},
"trino://user:pass@localhost:8080/hive/sales",
),
(
{"login": "user@example.com", "password": "p@ss:word", "host": "localhost", "schema": "hive"},
"trino://user%40example.com:p%40ss%3Aword@localhost/hive",
),
(
{"host": "localhost", "port": 8080, "schema": "hive"},
"trino://localhost:8080/hive",
),
(
{
"login": "user",
"host": "host.example.com",
"schema": "hive",
"extra": json.dumps({"param1": "value1", "param2": "value2"}),
},
"trino://user@host.example.com/hive?param1=value1&param2=value2",
),
],
ids=[
"basic-connection",
"with-extra-schema",
"special-chars",
"no-credentials",
"extra-params",
],
)
def test_get_uri(conn_params, expected_uri):
"""Test TrinoHook.get_uri properly formats connection URIs."""
with patch(HOOK_GET_CONNECTION) as mock_get_connection:
mock_get_connection.return_value = Connection(**conn_params)
hook = TrinoHook()
assert hook.get_uri() == expected_uri