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
2 changes: 2 additions & 0 deletions plugins/bigquery/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@ dependencies = [
"flyte[connector]",
"google-cloud-bigquery",
"google-cloud-bigquery-storage",
"pandas",
"pyarrow",
]

[dependency-groups]
Expand Down
7 changes: 6 additions & 1 deletion plugins/bigquery/src/flyteplugins/bigquery/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,12 @@ def run_query(date: str) -> DataFrame[dict]:
```
"""

from flyte.io.extend import DataFrameTransformerEngine

from flyteplugins.bigquery.connector import BigQueryConnector
from flyteplugins.bigquery.dataframe import BQToPandasDecodingHandler
from flyteplugins.bigquery.task import BigQueryConfig, BigQueryTask

__all__ = ["BigQueryConfig", "BigQueryConnector", "BigQueryTask"]
DataFrameTransformerEngine.register(BQToPandasDecodingHandler())

__all__ = ["BQToPandasDecodingHandler", "BigQueryConfig", "BigQueryConnector", "BigQueryTask"]
79 changes: 79 additions & 0 deletions plugins/bigquery/src/flyteplugins/bigquery/dataframe.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
import typing

from flyte.io.extend import DataFrameDecoder
from flyteidl2.core import literals_pb2
from google.cloud import bigquery_storage
from google.cloud.bigquery_storage_v1 import types

if typing.TYPE_CHECKING:
import pandas as pd
import pyarrow as pa
else:
from flyte._utils import lazy_module

pd = lazy_module("pandas")
pa = lazy_module("pyarrow")

BIGQUERY = "bq"


def _parse_bigquery_uri(uri: str) -> tuple[str, str, str]:
"""Parse bq://<project>:<dataset>.<table> into its components."""
if not uri.startswith("bq://"):
raise ValueError(f"Invalid BigQuery URI {uri!r}. Expected bq://<project>:<dataset>.<table>.")

try:
project_id, table_path = uri.removeprefix("bq://").split(":", 1)
dataset_id, table_id = table_path.split(".", 1)
except ValueError as exc:
raise ValueError(f"Invalid BigQuery URI {uri!r}. Expected bq://<project>:<dataset>.<table>.") from exc

if not project_id or not dataset_id or not table_id:
raise ValueError(f"Invalid BigQuery URI {uri!r}. Expected bq://<project>:<dataset>.<table>.")

return project_id, dataset_id, table_id


def _read_from_bq(
flyte_value: literals_pb2.StructuredDataset,
current_task_metadata: literals_pb2.StructuredDatasetMetadata,
) -> "pd.DataFrame":
project_id, dataset_id, table_id = _parse_bigquery_uri(flyte_value.uri)

read_options = None
structured_dataset_type = current_task_metadata.structured_dataset_type
if structured_dataset_type and structured_dataset_type.columns:
read_options = types.ReadSession.TableReadOptions(
selected_fields=[column.name for column in structured_dataset_type.columns]
)

table = f"projects/{project_id}/datasets/{dataset_id}/tables/{table_id}"
read_session = types.ReadSession(
table=table,
data_format=types.DataFormat.ARROW,
read_options=read_options,
)
client = bigquery_storage.BigQueryReadClient()
session = client.create_read_session(
parent=f"projects/{project_id}",
read_session=read_session,
)

frames = [page.to_dataframe() for stream in session.streams for page in client.read_rows(stream.name).rows().pages]
if frames:
return pd.concat(frames)

schema = pa.ipc.read_schema(pa.py_buffer(session.arrow_schema.serialized_schema))
return schema.empty_table().to_pandas()


class BQToPandasDecodingHandler(DataFrameDecoder):
def __init__(self):
super().__init__(pd.DataFrame, BIGQUERY, "")

async def decode(
self,
flyte_value: literals_pb2.StructuredDataset,
current_task_metadata: literals_pb2.StructuredDatasetMetadata,
) -> "pd.DataFrame":
return _read_from_bq(flyte_value, current_task_metadata)
87 changes: 87 additions & 0 deletions plugins/bigquery/tests/test_dataframe.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
from unittest.mock import MagicMock, patch

import pandas as pd
import pyarrow as pa
import pytest
from flyte.io.extend import DataFrameTransformerEngine
from flyteidl2.core import literals_pb2, types_pb2

from flyteplugins.bigquery.dataframe import (
BQToPandasDecodingHandler,
_parse_bigquery_uri,
_read_from_bq,
)


def _structured_dataset(uri: str = "bq://test-project:test_dataset.test_table"):
return literals_pb2.StructuredDataset(uri=uri)


def _metadata(*column_names: str):
return literals_pb2.StructuredDatasetMetadata(
structured_dataset_type=types_pb2.StructuredDatasetType(
columns=[types_pb2.StructuredDatasetType.DatasetColumn(name=name) for name in column_names]
)
)


def test_bigquery_decoder_is_registered():
decoder = DataFrameTransformerEngine.get_decoder(pd.DataFrame, "bq", "")
assert isinstance(decoder, BQToPandasDecodingHandler)


@pytest.mark.parametrize(
("uri", "expected"),
[
("bq://project:dataset.table", ("project", "dataset", "table")),
("bq://my-project:my_dataset.table$20260101", ("my-project", "my_dataset", "table$20260101")),
],
)
def test_parse_bigquery_uri(uri, expected):
assert _parse_bigquery_uri(uri) == expected


@pytest.mark.parametrize("uri", ["gs://bucket/file", "bq://project", "bq://:dataset.table", "bq://project:.table"])
def test_parse_bigquery_uri_rejects_invalid_uri(uri):
with pytest.raises(ValueError, match="Expected bq://"):
_parse_bigquery_uri(uri)


def test_read_from_bq():
first = pd.DataFrame({"name": ["Alice"], "age": [25]})
second = pd.DataFrame({"name": ["Bob"], "age": [30]})
pages = [MagicMock(), MagicMock()]
pages[0].to_dataframe.return_value = first
pages[1].to_dataframe.return_value = second

session = MagicMock()
session.streams = [MagicMock(name="stream-one")]
reader = MagicMock()
reader.rows.return_value.pages = pages

with patch("flyteplugins.bigquery.dataframe.bigquery_storage.BigQueryReadClient") as client_cls:
client = client_cls.return_value
client.create_read_session.return_value = session
client.read_rows.return_value = reader

result = _read_from_bq(_structured_dataset(), _metadata("name"))

pd.testing.assert_frame_equal(result.reset_index(drop=True), pd.concat([first, second], ignore_index=True))
request = client.create_read_session.call_args.kwargs
assert request["parent"] == "projects/test-project"
assert request["read_session"].table == "projects/test-project/datasets/test_dataset/tables/test_table"
assert list(request["read_session"].read_options.selected_fields) == ["name"]


def test_read_from_empty_bq_table():
schema = pa.schema([("name", pa.string()), ("age", pa.int64())])
session = MagicMock()
session.streams = []
session.arrow_schema.serialized_schema = schema.serialize().to_pybytes()

with patch("flyteplugins.bigquery.dataframe.bigquery_storage.BigQueryReadClient") as client_cls:
client_cls.return_value.create_read_session.return_value = session
result = _read_from_bq(_structured_dataset(), _metadata())

assert result.empty
assert list(result.columns) == ["name", "age"]
Loading
Loading