|
| 1 | +from __future__ import annotations |
1 | 2 | from abc import ABC, abstractmethod |
2 | 3 | from collections import namedtuple, OrderedDict |
3 | 4 | from collections.abc import Iterable |
|
8 | 9 | import lz4.frame |
9 | 10 | from typing import Dict, List, Union, Any |
10 | 11 | import pyarrow |
| 12 | +from enum import Enum |
| 13 | +import copy |
11 | 14 |
|
12 | 15 | from databricks.sql import exc, OperationalError |
13 | 16 | from databricks.sql.cloudfetch.download_manager import ResultFileDownloadManager |
14 | 17 | from databricks.sql.thrift_api.TCLIService.ttypes import ( |
15 | 18 | TSparkArrowResultLink, |
16 | 19 | TSparkRowSetType, |
17 | 20 | TRowSet, |
| 21 | + TSparkParameter, |
| 22 | + TSparkParameterValue, |
18 | 23 | ) |
19 | 24 |
|
20 | 25 | BIT_MASKS = [1, 2, 4, 8, 16, 32, 64, 128] |
@@ -404,7 +409,7 @@ def convert_arrow_based_set_to_arrow_table(arrow_batches, lz4_compressed, schema |
404 | 409 |
|
405 | 410 |
|
406 | 411 | def convert_decimals_in_arrow_table(table, description) -> pyarrow.Table: |
407 | | - for (i, col) in enumerate(table.itercolumns()): |
| 412 | + for i, col in enumerate(table.itercolumns()): |
408 | 413 | if description[i][1] == "decimal": |
409 | 414 | decimal_col = col.to_pandas().apply( |
410 | 415 | lambda v: v if v is None else Decimal(v) |
@@ -470,3 +475,86 @@ def _create_arrow_array(t_col_value_wrapper, arrow_type): |
470 | 475 | result[i] = None |
471 | 476 |
|
472 | 477 | return pyarrow.array(result, type=arrow_type) |
| 478 | + |
| 479 | + |
| 480 | +class DbSqlType(Enum): |
| 481 | + STRING = "STRING" |
| 482 | + DATE = "DATE" |
| 483 | + TIMESTAMP = "TIMESTAMP" |
| 484 | + FLOAT = "FLOAT" |
| 485 | + DECIMAL = "DECIMAL" |
| 486 | + INTEGER = "INTEGER" |
| 487 | + BIGINT = "BIGINT" |
| 488 | + SMALLINT = "SMALLINT" |
| 489 | + TINYINT = "TINYINT" |
| 490 | + BOOLEAN = "BOOLEAN" |
| 491 | + INTERVAL_MONTH = "INTERVAL MONTH" |
| 492 | + INTERVAL_DAY = "INTERVAL DAY" |
| 493 | + |
| 494 | + |
| 495 | +class DbSqlParameter: |
| 496 | + name: str |
| 497 | + value: Any |
| 498 | + type: DbSqlType |
| 499 | + |
| 500 | + def __init__(self, name="", value=None, type=None): |
| 501 | + self.name = name |
| 502 | + self.value = value |
| 503 | + self.type = type |
| 504 | + |
| 505 | + def __eq__(self, other): |
| 506 | + return isinstance(other, self.__class__) and self.__dict__ == other.__dict__ |
| 507 | + |
| 508 | + |
| 509 | +def named_parameters_to_dbsqlparams_v1(parameters: Dict[str, str]): |
| 510 | + dbsqlparams = [] |
| 511 | + for name, parameter in parameters.items(): |
| 512 | + dbsqlparams.append(DbSqlParameter(name=name, value=parameter)) |
| 513 | + return dbsqlparams |
| 514 | + |
| 515 | + |
| 516 | +def named_parameters_to_dbsqlparams_v2(parameters: List[Any]): |
| 517 | + dbsqlparams = [] |
| 518 | + for parameter in parameters: |
| 519 | + if isinstance(parameter, DbSqlParameter): |
| 520 | + dbsqlparams.append(parameter) |
| 521 | + else: |
| 522 | + dbsqlparams.append(DbSqlParameter(value=parameter)) |
| 523 | + return dbsqlparams |
| 524 | + |
| 525 | + |
| 526 | +def infer_types(params: list[DbSqlParameter]): |
| 527 | + type_lookup_table = { |
| 528 | + str: DbSqlType.STRING, |
| 529 | + int: DbSqlType.INTEGER, |
| 530 | + float: DbSqlType.FLOAT, |
| 531 | + datetime.datetime: DbSqlType.TIMESTAMP, |
| 532 | + bool: DbSqlType.BOOLEAN, |
| 533 | + } |
| 534 | + newParams = copy.deepcopy(params) |
| 535 | + for param in newParams: |
| 536 | + if not param.type: |
| 537 | + if type(param.value) in type_lookup_table: |
| 538 | + param.type = type_lookup_table[type(param.value)] |
| 539 | + else: |
| 540 | + raise ValueError("Parameter type cannot be inferred") |
| 541 | + param.value = str(param.value) |
| 542 | + return newParams |
| 543 | + |
| 544 | + |
| 545 | +def named_parameters_to_tsparkparams(parameters: Union[List[Any], Dict[str, str]]): |
| 546 | + tspark_params = [] |
| 547 | + if isinstance(parameters, dict): |
| 548 | + dbsql_params = named_parameters_to_dbsqlparams_v1(parameters) |
| 549 | + else: |
| 550 | + dbsql_params = named_parameters_to_dbsqlparams_v2(parameters) |
| 551 | + inferred_type_parameters = infer_types(dbsql_params) |
| 552 | + for param in inferred_type_parameters: |
| 553 | + tspark_params.append( |
| 554 | + TSparkParameter( |
| 555 | + type=param.type.value, |
| 556 | + name=param.name, |
| 557 | + value=TSparkParameterValue(stringValue=param.value), |
| 558 | + ) |
| 559 | + ) |
| 560 | + return tspark_params |
0 commit comments