Jaehyeon Kim
12/04/2024, 10:06 PMdef test_read_sql(self):
df = daft.read_sql(
sql="SELECT * FROM demo LIMIT 770",
conn=lambda: sqlalchemy.create_engine("sqlite:///develop.db").connect(),
partition_col="id",
num_partitions=7,
)
physical_plan_scheduler = df._builder.to_physical_plan_scheduler(
daft.context.get_context().daft_execution_config
)
physical_plan_dict = json.loads(physical_plan_scheduler.to_json_string())
sql_queries = []
for task in physical_plan_dict["TabularScan"]["scan_tasks"]:
sql_queries.append(task["file_format_config"]["Database"]["sql"])
print(sql_queries)
self.assertEqual(len(sql_queries), 7)
When I add it to a custom SQL client, however, it fails with the following error.
FAILED tests/sql_client_test.py::TestSqlClient::test_sql_client - ValueError: TypeError: cannot pickle 'builtins.LogicalPlanBuilder' object
def test_sql_client(self):
sql_client = SqlClient(
db_url="sqlite:///develop.db",
sql="SELECT * FROM demo LIMIT 770",
partition_col="id",
num_partitions=7,
)
sql_queries = sql_client.split_queries()
print(sql_queries)
self.assertEqual(len(sql_queries), 7)
Here is the sql client source.
import json
from typing import Dict, List, Optional
import sqlalchemy
from sqlalchemy.engine import Connection
import daft
from daft.dataframe import DataFrame
from daft.datatype import DataType
__all__ = ["SqlClient", "SqlClientError"]
class SqlClientError(Exception):
def __init__(self, message=None, code=None):
self.message = message
self.code = code
class SqlClient(object):
def __init__(
self,
db_url: str,
sql: str,
partition_col: Optional[str] = None,
num_partitions: Optional[int] = None,
disable_pushdowns_to_sql: bool = False,
infer_schema: bool = True,
infer_schema_length: int = 10,
schema: Optional[Dict[str, DataType]] = None,
**kwargs,
):
self.df = self.create_df(
sql,
lambda: self.create_conn(db_url, **kwargs),
partition_col,
num_partitions,
disable_pushdowns_to_sql,
infer_schema,
infer_schema_length,
schema,
)
def create_conn(self, db_url: str, **kwargs) -> Connection:
try:
return sqlalchemy.create_engine(db_url, **kwargs).connect()
except Exception as e:
code = e.code if hasattr(e, "code") else None
raise SqlClientError(str(e), code)
def create_df(
self,
sql: str,
conn: Connection,
partition_col: Optional[str] = None,
num_partitions: Optional[int] = None,
disable_pushdowns_to_sql: bool = False,
infer_schema: bool = True,
infer_schema_length: int = 10,
schema: Optional[Dict[str, DataType]] = None,
) -> DataFrame:
return daft.read_sql(
sql=sql,
conn=conn,
partition_col=partition_col,
num_partitions=num_partitions,
disable_pushdowns_to_sql=disable_pushdowns_to_sql,
infer_schema=infer_schema,
infer_schema_length=infer_schema_length,
schema=schema,
)
def split_queries(self) -> List[str]:
physical_plan_scheduler = self.df._builder.to_physical_plan_scheduler(
daft.context.get_context().daft_execution_config
)
physical_plan_dict = json.loads(physical_plan_scheduler.to_json_string())
sql_queries = []
for task in physical_plan_dict["TabularScan"]["scan_tasks"]:
sql_queries.append(task["file_format_config"]["Database"]["sql"])
return sql_queries
I guess it is due to Python's inability to pickle module object (https://stackoverflow.com/questions/2790828/python-cant-pickle-module-objects-error).
One of the answers of the above stackoverflow thread mentioned to use the dill package instead of pickle.
Is it the right option? If so, can you please guide me how to do so? (I believe it would be a limitation to building a python package that depends on daft...)Colin Ho
12/05/2024, 11:45 PMLogicalPlanBuilder actually cannot be pickled because it is not serializable. My guess is because you are storing the dataframe in self.df . If possible, could you try storing just the physical_plan_dict? Or store the arguments to making a dataframe instead, then when only build it when you need to. Lmk if this helpsJaehyeon Kim
12/06/2024, 12:07 AM