Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
horizontal_shard.py479 linesDownload Raw Back to ext
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 
codekingpro/portable-devtools · Team Ai