Skip to content

Commit 978adb3

Browse files
authored
Install sqlalchemy-spanner package into Google provider (#31925)
1 parent d5bf74c commit 978adb3

4 files changed

Lines changed: 83 additions & 5 deletions

File tree

airflow/providers/google/cloud/hooks/spanner.py

Lines changed: 46 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,28 +18,43 @@
1818
"""This module contains a Google Cloud Spanner Hook."""
1919
from __future__ import annotations
2020

21-
from typing import Callable, Sequence
21+
from typing import Callable, NamedTuple, Sequence
2222

2323
from google.api_core.exceptions import AlreadyExists, GoogleAPICallError
2424
from google.cloud.spanner_v1.client import Client
2525
from google.cloud.spanner_v1.database import Database
2626
from google.cloud.spanner_v1.instance import Instance
2727
from google.cloud.spanner_v1.transaction import Transaction
2828
from google.longrunning.operations_grpc_pb2 import Operation
29+
from sqlalchemy import create_engine
2930

3031
from airflow.exceptions import AirflowException
32+
from airflow.providers.common.sql.hooks.sql import DbApiHook
3133
from airflow.providers.google.common.consts import CLIENT_INFO
32-
from airflow.providers.google.common.hooks.base_google import GoogleBaseHook
34+
from airflow.providers.google.common.hooks.base_google import GoogleBaseHook, get_field
3335

3436

35-
class SpannerHook(GoogleBaseHook):
37+
class SpannerConnectionParams(NamedTuple):
38+
"""Information about Google Spanner connection parameters."""
39+
40+
project_id: str | None
41+
instance_id: str | None
42+
database_id: str | None
43+
44+
45+
class SpannerHook(GoogleBaseHook, DbApiHook):
3646
"""
3747
Hook for Google Cloud Spanner APIs.
3848
3949
All the methods in the hook where project_id is used must be called with
4050
keyword arguments rather than positional.
4151
"""
4252

53+
conn_name_attr = "gcp_conn_id"
54+
default_conn_name = "google_cloud_spanner_default"
55+
conn_type = "gcpspanner"
56+
hook_name = "Google Cloud Spanner"
57+
4358
def __init__(
4459
self,
4560
gcp_conn_id: str = "google_cloud_default",
@@ -70,6 +85,34 @@ def _get_client(self, project_id: str) -> Client:
7085
)
7186
return self._client
7287

88+
def _get_conn_params(self) -> SpannerConnectionParams:
89+
"""Extract spanner database connection parameters."""
90+
extras = self.get_connection(self.gcp_conn_id).extra_dejson
91+
project_id = get_field(extras, "project_id") or self.project_id
92+
instance_id = get_field(extras, "instance_id")
93+
database_id = get_field(extras, "database_id")
94+
return SpannerConnectionParams(project_id, instance_id, database_id)
95+
96+
def get_uri(self) -> str:
97+
"""Override DbApiHook get_uri method for get_sqlalchemy_engine()."""
98+
project_id, instance_id, database_id = self._get_conn_params()
99+
if not all([instance_id, database_id]):
100+
raise AirflowException("The instance_id or database_id were not specified")
101+
return f"spanner+spanner:///projects/{project_id}/instances/{instance_id}/databases/{database_id}"
102+
103+
def get_sqlalchemy_engine(self, engine_kwargs=None):
104+
"""
105+
Get an sqlalchemy_engine object.
106+
107+
:param engine_kwargs: Kwargs used in :func:`~sqlalchemy.create_engine`.
108+
:return: the created engine.
109+
"""
110+
if engine_kwargs is None:
111+
engine_kwargs = {}
112+
project_id, _, _ = self._get_conn_params()
113+
spanner_client = self._get_client(project_id=project_id)
114+
return create_engine(self.get_uri(), connect_args={"client": spanner_client}, **engine_kwargs)
115+
73116
@GoogleBaseHook.fallback_to_default_project_id
74117
def get_instance(
75118
self,

airflow/providers/google/provider.yaml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -126,6 +126,7 @@ dependencies:
126126
- proto-plus>=1.19.6
127127
- PyOpenSSL
128128
- sqlalchemy-bigquery>=1.2.1
129+
- sqlalchemy-spanner>=1.6.2
129130

130131
integrations:
131132
- integration-name: Google Analytics360

generated/provider_dependencies.json

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -440,7 +440,8 @@
440440
"pandas-gbq",
441441
"pandas>=0.17.1",
442442
"proto-plus>=1.19.6",
443-
"sqlalchemy-bigquery>=1.2.1"
443+
"sqlalchemy-bigquery>=1.2.1",
444+
"sqlalchemy-spanner>=1.6.2"
444445
],
445446
"cross-providers-deps": [
446447
"amazon",

tests/providers/google/cloud/hooks/test_spanner.py

Lines changed: 34 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,9 +18,10 @@
1818
from __future__ import annotations
1919

2020
from unittest import mock
21-
from unittest.mock import PropertyMock
21+
from unittest.mock import MagicMock, PropertyMock
2222

2323
import pytest
24+
import sqlalchemy
2425

2526
from airflow.providers.google.cloud.hooks.spanner import SpannerHook
2627
from airflow.providers.google.common.consts import CLIENT_INFO
@@ -33,6 +34,8 @@
3334
SPANNER_INSTANCE = "instance"
3435
SPANNER_CONFIGURATION = "configuration"
3536
SPANNER_DATABASE = "database-name"
37+
SPANNER_PROJECT_ID = "test_project_id"
38+
SPANNER_CONN_PARAMS = (SPANNER_PROJECT_ID, SPANNER_INSTANCE, SPANNER_DATABASE)
3639

3740

3841
class TestGcpSpannerHookDefaultProjectId:
@@ -431,6 +434,21 @@ def test_execute_dml_overridden_project_id(self, get_client):
431434
run_in_transaction_method.assert_called_once_with(mock.ANY)
432435
assert res is None
433436

437+
def test_get_uri(self):
438+
self.spanner_hook_default_project_id._get_conn_params = MagicMock(return_value=SPANNER_CONN_PARAMS)
439+
uri = self.spanner_hook_default_project_id.get_uri()
440+
assert (
441+
uri
442+
== f"spanner+spanner:///projects/{SPANNER_PROJECT_ID}/instances/{SPANNER_INSTANCE}/databases/{SPANNER_DATABASE}"
443+
)
444+
445+
@mock.patch("airflow.providers.google.cloud.hooks.spanner.SpannerHook._get_client")
446+
def test_get_sqlalchemy_engine(self, get_client):
447+
self.spanner_hook_default_project_id._get_conn_params = MagicMock(return_value=SPANNER_CONN_PARAMS)
448+
engine = self.spanner_hook_default_project_id.get_sqlalchemy_engine()
449+
assert isinstance(engine, sqlalchemy.engine.Engine)
450+
assert engine.name == "spanner+spanner"
451+
434452

435453
class TestGcpSpannerHookNoDefaultProjectID:
436454
def setup_method(self):
@@ -675,3 +693,18 @@ def test_execute_dml_overridden_project_id(self, get_client):
675693
database_method.assert_called_once_with(database_id="database-name")
676694
run_in_transaction_method.assert_called_once_with(mock.ANY)
677695
assert res is None
696+
697+
def test_get_uri(self):
698+
self.spanner_hook_no_default_project_id._get_conn_params = MagicMock(return_value=SPANNER_CONN_PARAMS)
699+
uri = self.spanner_hook_no_default_project_id.get_uri()
700+
assert (
701+
uri
702+
== f"spanner+spanner:///projects/{SPANNER_PROJECT_ID}/instances/{SPANNER_INSTANCE}/databases/{SPANNER_DATABASE}"
703+
)
704+
705+
@mock.patch("airflow.providers.google.cloud.hooks.spanner.SpannerHook._get_client")
706+
def test_get_sqlalchemy_engine(self, get_client):
707+
self.spanner_hook_no_default_project_id._get_conn_params = MagicMock(return_value=SPANNER_CONN_PARAMS)
708+
engine = self.spanner_hook_no_default_project_id.get_sqlalchemy_engine()
709+
assert isinstance(engine, sqlalchemy.engine.Engine)
710+
assert engine.name == "spanner+spanner"

0 commit comments

Comments
 (0)