codekingpro/portable-devtools
114k
1# ext/horizontal_shard.py
2# Copyright (C) 2005-2026 the SQLAlchemy authors and contributors
3# <see AUTHORS file>
4#
5# This module is part of SQLAlchemy and is released under
6# the MIT License: https://www.opensource.org/licenses/mit-license.php
7
8"""Horizontal sharding support.
9
10Defines a rudimental 'horizontal sharding' system which allows a Session to
11distribute queries and persistence operations across multiple databases.
12
13For a usage example, see the :ref:`examples_sharding` example included in
14the source distribution.
15
16.. deepalchemy:: The horizontal sharding extension is an advanced feature,
17 involving a complex statement -> database interaction as well as
18 use of semi-public APIs for non-trivial cases. Simpler approaches to
19 referring to multiple database "shards", most commonly using a distinct
20 :class:`_orm.Session` per "shard", should always be considered first
21 before using this more complex and less-production-tested system.
22
23
24
25"""
26from __future__ import annotations
27
28from typing import Any
29from typing import Callable
30from typing import Dict
31from typing import Iterable
32from typing import Optional
33from typing import Tuple
34from typing import Type
35from typing import TYPE_CHECKING
36from typing import TypeVar
37from typing import Union
38
39from .. import event
40from .. import exc
41from .. import inspect
42from .. import util
43from ..orm import PassiveFlag
44from ..orm._typing import OrmExecuteOptionsParameter
45from ..orm.interfaces import ORMOption
46from ..orm.mapper import Mapper
47from ..orm.query import Query
48from ..orm.session import _BindArguments
49from ..orm.session import _PKIdentityArgument
50from ..orm.session import Session
51from ..util.typing import Protocol
52from ..util.typing import Self
53
54if TYPE_CHECKING:
55 from ..engine.base import Connection
56 from ..engine.base import Engine
57 from ..engine.base import OptionEngine
58 from ..engine.result import IteratorResult
59 from ..engine.result import Result
60 from ..orm import LoaderCallableStatus
61 from ..orm._typing import _O
62 from ..orm.bulk_persistence import BulkUDCompileState
63 from ..orm.context import QueryContext
64 from ..orm.session import _EntityBindKey
65 from ..orm.session import _SessionBind
66 from ..orm.session import ORMExecuteState
67 from ..orm.state import InstanceState
68 from ..sql import Executable
69 from ..sql._typing import _TP
70 from ..sql.elements import ClauseElement
71
72__all__ = ["ShardedSession", "ShardedQuery"]
73
74_T = TypeVar("_T", bound=Any)
75
76
77ShardIdentifier = str
78
79
80class ShardChooser(Protocol):
81 def __call__(
82 self,
83 mapper: Optional[Mapper[_T]],
84 instance: Any,
85 clause: Optional[ClauseElement],
86 ) -> Any: ...
87
88
89class IdentityChooser(Protocol):
90 def __call__(
91 self,
92 mapper: Mapper[_T],
93 primary_key: _PKIdentityArgument,
94 *,
95 lazy_loaded_from: Optional[InstanceState[Any]],
96 execution_options: OrmExecuteOptionsParameter,
97 bind_arguments: _BindArguments,
98 **kw: Any,
99 ) -> Any: ...
100
101
102class ShardedQuery(Query[_T]):
103 """Query class used with :class:`.ShardedSession`.
104
105 .. legacy:: The :class:`.ShardedQuery` is a subclass of the legacy
106 :class:`.Query` class. The :class:`.ShardedSession` now supports
107 2.0 style execution via the :meth:`.ShardedSession.execute` method.
108
109 """
110
111 def __init__(self, *args: Any, **kwargs: Any) -> None:
112 super().__init__(*args, **kwargs)
113 assert isinstance(self.session, ShardedSession)
114
115 self.identity_chooser = self.session.identity_chooser
116 self.execute_chooser = self.session.execute_chooser
117 self._shard_id = None
118
119 def set_shard(self, shard_id: ShardIdentifier) -> Self:
120 """Return a new query, limited to a single shard ID.
121
122 All subsequent operations with the returned query will
123 be against the single shard regardless of other state.
124
125 The shard_id can be passed for a 2.0 style execution to the
126 bind_arguments dictionary of :meth:`.Session.execute`::
127
128 results = session.execute(stmt, bind_arguments={"shard_id": "my_shard"})
129
130 """ # noqa: E501
131 return self.execution_options(_sa_shard_id=shard_id)
132
133
134class ShardedSession(Session):
135 shard_chooser: ShardChooser
136 identity_chooser: IdentityChooser
137 execute_chooser: Callable[[ORMExecuteState], Iterable[Any]]
138
139 def __init__(
140 self,
141 shard_chooser: ShardChooser,
142 identity_chooser: Optional[IdentityChooser] = None,
143 execute_chooser: Optional[
144 Callable[[ORMExecuteState], Iterable[Any]]
145 ] = None,
146 shards: Optional[Dict[str, Any]] = None,
147 query_cls: Type[Query[_T]] = ShardedQuery,
148 *,
149 id_chooser: Optional[
150 Callable[[Query[_T], Iterable[_T]], Iterable[Any]]
151 ] = None,
152 query_chooser: Optional[Callable[[Executable], Iterable[Any]]] = None,
153 **kwargs: Any,
154 ) -> None:
155 """Construct a ShardedSession.
156
157 :param shard_chooser: A callable which, passed a Mapper, a mapped
158 instance, and possibly a SQL clause, returns a shard ID. This id
159 may be based off of the attributes present within the object, or on
160 some round-robin scheme. If the scheme is based on a selection, it
161 should set whatever state on the instance to mark it in the future as
162 participating in that shard.
163
164 :param identity_chooser: A callable, passed a Mapper and primary key
165 argument, which should return a list of shard ids where this
166 primary key might reside.
167
168 .. versionchanged:: 2.0 The ``identity_chooser`` parameter
169 supersedes the ``id_chooser`` parameter.
170
171 :param execute_chooser: For a given :class:`.ORMExecuteState`,
172 returns the list of shard_ids
173 where the query should be issued. Results from all shards returned
174 will be combined together into a single listing.
175
176 .. versionchanged:: 1.4 The ``execute_chooser`` parameter
177 supersedes the ``query_chooser`` parameter.
178
179 :param shards: A dictionary of string shard names
180 to :class:`~sqlalchemy.engine.Engine` objects.
181
182 """
183 super().__init__(query_cls=query_cls, **kwargs)
184
185 event.listen(
186 self, "do_orm_execute", execute_and_instances, retval=True
187 )
188 self.shard_chooser = shard_chooser
189
190 if id_chooser:
191 _id_chooser = id_chooser
192 util.warn_deprecated(
193 "The ``id_chooser`` parameter is deprecated; "
194 "please use ``identity_chooser``.",
195 "2.0",
196 )
197
198 def _legacy_identity_chooser(
199 mapper: Mapper[_T],
200 primary_key: _PKIdentityArgument,
201 *,
202 lazy_loaded_from: Optional[InstanceState[Any]],
203 execution_options: OrmExecuteOptionsParameter,
204 bind_arguments: _BindArguments,
205 **kw: Any,
206 ) -> Any:
207 q = self.query(mapper)
208 if lazy_loaded_from:
209 q = q._set_lazyload_from(lazy_loaded_from)
210 return _id_chooser(q, primary_key)
211
212 self.identity_chooser = _legacy_identity_chooser
213 elif identity_chooser:
214 self.identity_chooser = identity_chooser
215 else:
216 raise exc.ArgumentError(
217 "identity_chooser or id_chooser is required"
218 )
219
220 if query_chooser:
221 _query_chooser = query_chooser
222 util.warn_deprecated(
223 "The ``query_chooser`` parameter is deprecated; "
224 "please use ``execute_chooser``.",
225 "1.4",
226 )
227 if execute_chooser:
228 raise exc.ArgumentError(
229 "Can't pass query_chooser and execute_chooser "
230 "at the same time."
231 )
232
233 def _default_execute_chooser(
234 orm_context: ORMExecuteState,
235 ) -> Iterable[Any]:
236 return _query_chooser(orm_context.statement)
237
238 if execute_chooser is None:
239 execute_chooser = _default_execute_chooser
240
241 if execute_chooser is None:
242 raise exc.ArgumentError(
243 "execute_chooser or query_chooser is required"
244 )
245 self.execute_chooser = execute_chooser
246 self.__shards: Dict[ShardIdentifier, _SessionBind] = {}
247 if shards is not None:
248 for k in shards:
249 self.bind_shard(k, shards[k])
250
251 def _identity_lookup(
252 self,
253 mapper: Mapper[_O],
254 primary_key_identity: Union[Any, Tuple[Any, ...]],
255 identity_token: Optional[Any] = None,
256 passive: PassiveFlag = PassiveFlag.PASSIVE_OFF,
257 lazy_loaded_from: Optional[InstanceState[Any]] = None,
258 execution_options: OrmExecuteOptionsParameter = util.EMPTY_DICT,
259 bind_arguments: Optional[_BindArguments] = None,
260 **kw: Any,
261 ) -> Union[Optional[_O], LoaderCallableStatus]:
262 """override the default :meth:`.Session._identity_lookup` method so
263 that we search for a given non-token primary key identity across all
264 possible identity tokens (e.g. shard ids).
265
266 .. versionchanged:: 1.4 Moved :meth:`.Session._identity_lookup` from
267 the :class:`_query.Query` object to the :class:`.Session`.
268
269 """
270
271 if identity_token is not None:
272 obj = super()._identity_lookup(
273 mapper,
274 primary_key_identity,
275 identity_token=identity_token,
276 **kw,
277 )
278
279 return obj
280 else:
281 for shard_id in self.identity_chooser(
282 mapper,
283 primary_key_identity,
284 lazy_loaded_from=lazy_loaded_from,
285 execution_options=execution_options,
286 bind_arguments=dict(bind_arguments) if bind_arguments else {},
287 ):
288 obj2 = super()._identity_lookup(
289 mapper,
290 primary_key_identity,
291 identity_token=shard_id,
292 lazy_loaded_from=lazy_loaded_from,
293 **kw,
294 )
295 if obj2 is not None:
296 return obj2
297
298 return None
299
300 def _choose_shard_and_assign(
301 self,
302 mapper: Optional[_EntityBindKey[_O]],
303 instance: Any,
304 **kw: Any,
305 ) -> Any:
306 if instance is not None:
307 state = inspect(instance)
308 if state.key:
309 token = state.key[2]
310 assert token is not None
311 return token
312 elif state.identity_token:
313 return state.identity_token
314
315 assert isinstance(mapper, Mapper)
316 shard_id = self.shard_chooser(mapper, instance, **kw)
317 if instance is not None:
318 state.identity_token = shard_id
319 return shard_id
320
321 def connection_callable(
322 self,
323 mapper: Optional[Mapper[_T]] = None,
324 instance: Optional[Any] = None,
325 shard_id: Optional[ShardIdentifier] = None,
326 **kw: Any,
327 ) -> Connection:
328 """Provide a :class:`_engine.Connection` to use in the unit of work
329 flush process.
330
331 """
332
333 if shard_id is None:
334 shard_id = self._choose_shard_and_assign(mapper, instance)
335
336 if self.in_transaction():
337 trans = self.get_transaction()
338 assert trans is not None
339 return trans.connection(mapper, shard_id=shard_id)
340 else:
341 bind = self.get_bind(
342 mapper=mapper, shard_id=shard_id, instance=instance
343 )
344
345 if isinstance(bind, Engine):
346 return bind.connect(**kw)
347 else:
348 assert isinstance(bind, Connection)
349 return bind
350
351 def get_bind(
352 self,
353 mapper: Optional[_EntityBindKey[_O]] = None,
354 *,
355 shard_id: Optional[ShardIdentifier] = None,
356 instance: Optional[Any] = None,
357 clause: Optional[ClauseElement] = None,
358 **kw: Any,
359 ) -> _SessionBind:
360 if shard_id is None:
361 shard_id = self._choose_shard_and_assign(
362 mapper, instance=instance, clause=clause
363 )
364 assert shard_id is not None
365 return self.__shards[shard_id]
366
367 def bind_shard(
368 self, shard_id: ShardIdentifier, bind: Union[Engine, OptionEngine]
369 ) -> None:
370 self.__shards[shard_id] = bind
371
372
373class set_shard_id(ORMOption):
374 """a loader option for statements to apply a specific shard id to the
375 primary query as well as for additional relationship and column
376 loaders.
377
378 The :class:`_horizontal.set_shard_id` option may be applied using
379 the :meth:`_sql.Executable.options` method of any executable statement::
380
381 stmt = (
382 select(MyObject)
383 .where(MyObject.name == "some name")
384 .options(set_shard_id("shard1"))
385 )
386
387 Above, the statement when invoked will limit to the "shard1" shard
388 identifier for the primary query as well as for all relationship and
389 column loading strategies, including eager loaders such as
390 :func:`_orm.selectinload`, deferred column loaders like :func:`_orm.defer`,
391 and the lazy relationship loader :func:`_orm.lazyload`.
392
393 In this way, the :class:`_horizontal.set_shard_id` option has much wider
394 scope than using the "shard_id" argument within the
395 :paramref:`_orm.Session.execute.bind_arguments` dictionary.
396
397
398 .. versionadded:: 2.0.0
399
400 """
401
402 __slots__ = ("shard_id", "propagate_to_loaders")
403
404 def __init__(
405 self, shard_id: ShardIdentifier, propagate_to_loaders: bool = True
406 ):
407 """Construct a :class:`_horizontal.set_shard_id` option.
408
409 :param shard_id: shard identifier
410 :param propagate_to_loaders: if left at its default of ``True``, the
411 shard option will take place for lazy loaders such as
412 :func:`_orm.lazyload` and :func:`_orm.defer`; if False, the option
413 will not be propagated to loaded objects. Note that :func:`_orm.defer`
414 always limits to the shard_id of the parent row in any case, so the
415 parameter only has a net effect on the behavior of the
416 :func:`_orm.lazyload` strategy.
417
418 """
419 self.shard_id = shard_id
420 self.propagate_to_loaders = propagate_to_loaders
421
422
423def execute_and_instances(
424 orm_context: ORMExecuteState,
425) -> Union[Result[_T], IteratorResult[_TP]]:
426 active_options: Union[
427 None,
428 QueryContext.default_load_options,
429 Type[QueryContext.default_load_options],
430 BulkUDCompileState.default_update_options,
431 Type[BulkUDCompileState.default_update_options],
432 ]
433
434 if orm_context.is_select:
435 active_options = orm_context.load_options
436
437 elif orm_context.is_update or orm_context.is_delete:
438 active_options = orm_context.update_delete_options
439 else:
440 active_options = None
441
442 session = orm_context.session
443 assert isinstance(session, ShardedSession)
444
445 def iter_for_shard(
446 shard_id: ShardIdentifier,
447 ) -> Union[Result[_T], IteratorResult[_TP]]:
448 bind_arguments = dict(orm_context.bind_arguments)
449 bind_arguments["shard_id"] = shard_id
450
451 orm_context.update_execution_options(identity_token=shard_id)
452 return orm_context.invoke_statement(bind_arguments=bind_arguments)
453
454 for orm_opt in orm_context._non_compile_orm_options:
455 # TODO: if we had an ORMOption that gets applied at ORM statement
456 # execution time, that would allow this to be more generalized.
457 # for now just iterate and look for our options
458 if isinstance(orm_opt, set_shard_id):
459 shard_id = orm_opt.shard_id
460 break
461 else:
462 if active_options and active_options._identity_token is not None:
463 shard_id = active_options._identity_token
464 elif "_sa_shard_id" in orm_context.execution_options:
465 shard_id = orm_context.execution_options["_sa_shard_id"]
466 elif "shard_id" in orm_context.bind_arguments:
467 shard_id = orm_context.bind_arguments["shard_id"]
468 else:
469 shard_id = None
470
471 if shard_id is not None:
472 return iter_for_shard(shard_id)
473 else:
474 partial = []
475 for shard_id in session.execute_chooser(orm_context):
476 result_ = iter_for_shard(shard_id)
477 partial.append(result_)
478 return partial[0].merge(*partial[1:])
479 