duckdb-sqlalchemy 0.19.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,237 @@
1
+ from __future__ import annotations
2
+
3
+ import hashlib
4
+ from itertools import cycle
5
+ from typing import (
6
+ Any,
7
+ Dict,
8
+ Mapping,
9
+ MutableMapping,
10
+ Optional,
11
+ Sequence,
12
+ Tuple,
13
+ Type,
14
+ Union,
15
+ )
16
+ from urllib.parse import urlencode
17
+
18
+ import sqlalchemy
19
+ from sqlalchemy import create_engine
20
+ from sqlalchemy.engine import URL as SAURL
21
+ from sqlalchemy.engine.url import make_url as sa_make_url
22
+ from sqlalchemy.pool import Pool, QueuePool
23
+
24
+ MOTHERDUCK_PATH_QUERY_KEYS = {
25
+ "user",
26
+ "session_hint",
27
+ "attach_mode",
28
+ "access_mode",
29
+ "dbinstance_inactivity_ttl",
30
+ "motherduck_dbinstance_inactivity_ttl",
31
+ "saas_mode",
32
+ "cache_buster",
33
+ }
34
+
35
+ DIALECT_QUERY_KEYS = {"duckdb_sqlalchemy_pool", "pool"}
36
+
37
+
38
+ def _normalize_path_query_aliases(path_query: Dict[str, Any]) -> Dict[str, Any]:
39
+ if "motherduck_dbinstance_inactivity_ttl" in path_query:
40
+ if "dbinstance_inactivity_ttl" not in path_query:
41
+ path_query["dbinstance_inactivity_ttl"] = path_query[
42
+ "motherduck_dbinstance_inactivity_ttl"
43
+ ]
44
+ path_query.pop("motherduck_dbinstance_inactivity_ttl", None)
45
+ return path_query
46
+
47
+
48
+ def split_url_query(query: Dict[str, Any]) -> Tuple[Dict[str, Any], Dict[str, Any]]:
49
+ path_query: Dict[str, Any] = {}
50
+ url_config: Dict[str, Any] = {}
51
+ for key, value in query.items():
52
+ if key in DIALECT_QUERY_KEYS:
53
+ continue
54
+ if key in MOTHERDUCK_PATH_QUERY_KEYS:
55
+ path_query[key] = value
56
+ else:
57
+ url_config[key] = value
58
+ return _normalize_path_query_aliases(path_query), url_config
59
+
60
+
61
+ def extract_path_query_from_config(config: Dict[str, Any]) -> Dict[str, Any]:
62
+ path_query: Dict[str, Any] = {}
63
+ for key in list(config):
64
+ if key in MOTHERDUCK_PATH_QUERY_KEYS:
65
+ path_query[key] = config.pop(key)
66
+ return _normalize_path_query_aliases(path_query)
67
+
68
+
69
+ def append_query_to_database(
70
+ database: Optional[str], query: Dict[str, Any]
71
+ ) -> Optional[str]:
72
+ if not query:
73
+ return database
74
+ query_string = urlencode(query, doseq=True)
75
+ if database is None:
76
+ return f"?{query_string}"
77
+ separator = "&" if "?" in database else "?"
78
+ return f"{database}{separator}{query_string}"
79
+
80
+
81
+ def _stringify_query_value(value: Any) -> str:
82
+ if isinstance(value, bool):
83
+ return "true" if value else "false"
84
+ return str(value)
85
+
86
+
87
+ def _coerce_query_value(value: Any) -> Any:
88
+ if value is None:
89
+ return None
90
+ if isinstance(value, Sequence) and not isinstance(value, (str, bytes)):
91
+ return tuple(_stringify_query_value(v) for v in value)
92
+ return _stringify_query_value(value)
93
+
94
+
95
+ def _coerce_query_mapping(mapping: Mapping[str, Any]) -> Dict[str, Any]:
96
+ return {
97
+ key: value
98
+ for key, value in ((k, _coerce_query_value(v)) for k, v in mapping.items())
99
+ if value is not None
100
+ }
101
+
102
+
103
+ def MotherDuckURL(
104
+ *,
105
+ database: str,
106
+ query: Optional[Mapping[str, Any]] = None,
107
+ path_query: Optional[Mapping[str, Any]] = None,
108
+ **kwargs: Any,
109
+ ) -> SAURL:
110
+ """
111
+ Build a SQLAlchemy URL for MotherDuck, ensuring routing/cache parameters
112
+ live in the database string.
113
+ """
114
+
115
+ path_params: Dict[str, Any] = dict(path_query or {})
116
+ config_params: Dict[str, Any] = dict(query or {})
117
+
118
+ for key, value in kwargs.items():
119
+ if key in MOTHERDUCK_PATH_QUERY_KEYS:
120
+ path_params[key] = value
121
+ else:
122
+ config_params[key] = value
123
+
124
+ path_params = _normalize_path_query_aliases(_coerce_query_mapping(path_params))
125
+ config_params = _coerce_query_mapping(config_params)
126
+
127
+ database_with_query = append_query_to_database(database, path_params)
128
+ return SAURL.create("duckdb", database=database_with_query, query=config_params)
129
+
130
+
131
+ def stable_session_hint(
132
+ value: Union[str, int],
133
+ *,
134
+ salt: Optional[str] = None,
135
+ length: int = 16,
136
+ ) -> str:
137
+ if length <= 0:
138
+ raise ValueError("length must be positive")
139
+ payload = str(value)
140
+ if salt:
141
+ payload = f"{salt}:{payload}"
142
+ digest = hashlib.sha256(payload.encode("utf-8")).hexdigest()
143
+ return digest[:length]
144
+
145
+
146
+ def create_motherduck_engine(
147
+ *,
148
+ database: str,
149
+ query: Optional[Mapping[str, Any]] = None,
150
+ connect_args: Optional[Mapping[str, Any]] = None,
151
+ performance: bool = False,
152
+ poolclass: Optional[Type[Pool]] = None,
153
+ pool_pre_ping: Optional[bool] = None,
154
+ pool_recycle: Optional[int] = None,
155
+ **path_params: Any,
156
+ ) -> sqlalchemy.engine.Engine:
157
+ url = MotherDuckURL(database=database, query=query, **path_params)
158
+
159
+ engine_kwargs: Dict[str, Any] = {}
160
+ if poolclass is not None:
161
+ engine_kwargs["poolclass"] = poolclass
162
+ if pool_pre_ping is not None:
163
+ engine_kwargs["pool_pre_ping"] = pool_pre_ping
164
+ if pool_recycle is not None:
165
+ engine_kwargs["pool_recycle"] = pool_recycle
166
+
167
+ if performance:
168
+ engine_kwargs.setdefault("poolclass", QueuePool)
169
+ engine_kwargs.setdefault("pool_pre_ping", True)
170
+ engine_kwargs.setdefault("pool_recycle", 23 * 3600)
171
+
172
+ return create_engine(url, connect_args=dict(connect_args or {}), **engine_kwargs)
173
+
174
+
175
+ def _normalize_path_item(path: Union[str, SAURL]) -> SAURL:
176
+ if isinstance(path, SAURL):
177
+ return path
178
+ if path.startswith("duckdb://"):
179
+ return sa_make_url(path)
180
+ return SAURL.create("duckdb", database=path)
181
+
182
+
183
+ def _merge_connect_args(
184
+ base: MutableMapping[str, Any], extra: Mapping[str, Any]
185
+ ) -> Dict[str, Any]:
186
+ merged = dict(base)
187
+ if not extra:
188
+ return merged
189
+ extra = dict(extra)
190
+ if "config" in extra:
191
+ merged["config"] = {**merged.get("config", {}), **extra.pop("config")}
192
+ if "url_config" in extra:
193
+ merged["url_config"] = {
194
+ **merged.get("url_config", {}),
195
+ **extra.pop("url_config"),
196
+ }
197
+ merged.update(extra)
198
+ return merged
199
+
200
+
201
+ def _copy_connect_params(params: Mapping[str, Any]) -> Dict[str, Any]:
202
+ copied = dict(params)
203
+ if "config" in copied:
204
+ copied["config"] = dict(copied["config"])
205
+ if "url_config" in copied:
206
+ copied["url_config"] = dict(copied["url_config"])
207
+ return copied
208
+
209
+
210
+ def create_engine_from_paths(
211
+ paths: Sequence[Union[str, SAURL]],
212
+ *,
213
+ connect_args: Optional[Mapping[str, Any]] = None,
214
+ **engine_kwargs: Any,
215
+ ) -> sqlalchemy.engine.Engine:
216
+ if not paths:
217
+ raise ValueError("paths must not be empty")
218
+
219
+ urls = [_normalize_path_item(path) for path in paths]
220
+ if len({url.drivername for url in urls}) != 1:
221
+ raise ValueError("all paths must use the same drivername")
222
+
223
+ from . import Dialect # avoid import cycle
224
+
225
+ dialect = Dialect()
226
+ connect_params = []
227
+ for url in urls:
228
+ _, params = dialect.create_connect_args(url)
229
+ connect_params.append(_merge_connect_args(params, connect_args or {}))
230
+
231
+ params_cycle = cycle(connect_params)
232
+
233
+ def creator() -> Any:
234
+ params = _copy_connect_params(next(params_cycle))
235
+ return dialect.connect(**params)
236
+
237
+ return create_engine(urls[0], creator=creator, **engine_kwargs)
@@ -0,0 +1,39 @@
1
+ from typing import Any, Iterable, Optional
2
+
3
+ from sqlalchemy import func
4
+
5
+ __all__ = ["table_function", "read_parquet", "read_csv", "read_csv_auto"]
6
+
7
+
8
+ def table_function(
9
+ name: str,
10
+ *args: Any,
11
+ columns: Optional[Iterable[str]] = None,
12
+ **kwargs: Any,
13
+ ) -> Any:
14
+ fn = getattr(func, name)(*args, **kwargs)
15
+ if columns:
16
+ if hasattr(fn, "table_valued"):
17
+ return fn.table_valued(*columns)
18
+ raise NotImplementedError(
19
+ "table_valued requires SQLAlchemy >= 1.4 to name columns"
20
+ )
21
+ return fn
22
+
23
+
24
+ def read_parquet(
25
+ path: str, *, columns: Optional[Iterable[str]] = None, **kwargs: Any
26
+ ) -> Any:
27
+ return table_function("read_parquet", path, columns=columns, **kwargs)
28
+
29
+
30
+ def read_csv(
31
+ path: str, *, columns: Optional[Iterable[str]] = None, **kwargs: Any
32
+ ) -> Any:
33
+ return table_function("read_csv", path, columns=columns, **kwargs)
34
+
35
+
36
+ def read_csv_auto(
37
+ path: str, *, columns: Optional[Iterable[str]] = None, **kwargs: Any
38
+ ) -> Any:
39
+ return table_function("read_csv_auto", path, columns=columns, **kwargs)
File without changes
@@ -0,0 +1,5 @@
1
+ from sqlalchemy.testing.requirements import SuiteRequirements
2
+
3
+
4
+ class Requirements(SuiteRequirements):
5
+ pass
File without changes
@@ -0,0 +1,55 @@
1
+ import warnings
2
+ from functools import wraps
3
+ from typing import Any, Callable, Generator, TypeVar
4
+
5
+ from pytest import fixture, raises
6
+ from sqlalchemy import create_engine
7
+ from sqlalchemy.dialects import registry # type: ignore
8
+ from sqlalchemy.engine import Dialect, Engine
9
+ from sqlalchemy.engine.base import Connection
10
+ from sqlalchemy.orm import Session, sessionmaker
11
+ from typing_extensions import ParamSpec
12
+
13
+ warnings.filterwarnings(
14
+ "ignore",
15
+ "distutils Version classes are deprecated. Use packaging.version instead.",
16
+ DeprecationWarning,
17
+ )
18
+ P = ParamSpec("P")
19
+
20
+ FuncT = TypeVar("FuncT", bound=Callable[..., Any])
21
+
22
+
23
+ @fixture
24
+ def engine() -> Engine:
25
+ registry.register("duckdb", "duckdb_sqlalchemy", "Dialect")
26
+
27
+ return create_engine("duckdb:///:memory:")
28
+
29
+
30
+ @fixture
31
+ def conn(engine: Engine) -> Generator[Connection, None, None]:
32
+ with engine.connect() as conn:
33
+ yield conn
34
+
35
+
36
+ @fixture()
37
+ def dialect(engine: Engine) -> Dialect:
38
+ return engine.dialect
39
+
40
+
41
+ @fixture
42
+ def session(engine: Engine) -> Session:
43
+ return sessionmaker(bind=engine)()
44
+
45
+
46
+ def raises_msg(msg: str) -> Callable[[Callable[P, None]], Callable[P, None]]:
47
+ def decorator(func: Callable[P, None]) -> Callable[P, None]:
48
+ @wraps(func)
49
+ def wrapped_test(*args: P.args, **kwargs: P.kwargs) -> None:
50
+ with raises(RuntimeError, match=msg):
51
+ func(*args, **kwargs)
52
+
53
+ return wrapped_test
54
+
55
+ return decorator
@@ -0,0 +1,3 @@
1
+ CREATE TABLE test_table (
2
+ duration INTERVAL
3
+ )
@@ -0,0 +1,11 @@
1
+ import os
2
+
3
+ import pytest
4
+
5
+ if not os.getenv("DUCKDB_SQLA_SUITE"):
6
+ pytest.skip(
7
+ "SQLAlchemy suite tests are disabled. Set DUCKDB_SQLA_SUITE=1 to enable.",
8
+ allow_module_level=True,
9
+ )
10
+
11
+ pytest_plugins = "sqlalchemy.testing.plugin.pytestplugin"
@@ -0,0 +1 @@
1
+ from sqlalchemy.testing.suite import * # noqa: F401,F403