Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
migrations.py277 linesDownload Raw Back to db
1import sys
2from typing import Sequence
3from typing_extensions import TypedDict, NotRequired
4from importlib_resources.abc import Traversable
5import re
6import hashlib
7from chromadb.db.base import SqlDB, Cursor
8from abc import abstractmethod
9from chromadb.config import System, Settings
10from chromadb.telemetry.opentelemetry import (
11    OpenTelemetryClient,
12    OpenTelemetryGranularity,
13    trace_method,
14)
15
16
17class MigrationFile(TypedDict):
18    path: NotRequired[Traversable]
19    dir: str
20    filename: str
21    version: int
22    scope: str
23
24
25class Migration(MigrationFile):
26    hash: str
27    sql: str
28
29
30class UninitializedMigrationsError(Exception):
31    def __init__(self) -> None:
32        super().__init__("Migrations have not been initialized")
33
34
35class UnappliedMigrationsError(Exception):
36    def __init__(self, dir: str, version: int):
37        self.dir = dir
38        self.version = version
39        super().__init__(
40            f"Unapplied migrations in {dir}, starting with version {version}"
41        )
42
43
44class InconsistentVersionError(Exception):
45    def __init__(self, dir: str, db_version: int, source_version: int):
46        super().__init__(
47            f"Inconsistent migration versions in {dir}:"
48            + f"db version was {db_version}, source version was {source_version}."
49            + " Has the migration sequence been modified since being applied to the DB?"
50        )
51
52
53class InconsistentHashError(Exception):
54    def __init__(self, path: str, db_hash: str, source_hash: str):
55        super().__init__(
56            f"Inconsistent hashes in {path}:"
57            + f"db hash was {db_hash}, source has was {source_hash}."
58            + " Was the migration file modified after being applied to the DB?"
59        )
60
61
62class InvalidHashError(Exception):
63    def __init__(self, alg: str):
64        super().__init__(f"Invalid hash algorithm specified: {alg}")
65
66
67class InvalidMigrationFilename(Exception):
68    pass
69
70
71class MigratableDB(SqlDB):
72    """Simple base class for databases which support basic migrations.
73
74    Migrations are SQL files stored as package resources and accessed via
75    importlib_resources.
76
77    All migrations in the same directory are assumed to be dependent on previous
78    migrations in the same directory, where "previous" is defined on lexographical
79    ordering of filenames.
80
81    Migrations have a ascending numeric version number and a hash of the file contents.
82    When migrations are applied, the hashes of previous migrations are checked to ensure
83    that the database is consistent with the source repository. If they are not, an
84    error is thrown and no migrations will be applied.
85
86    Migration files must follow the naming convention:
87    <version>.<description>.<scope>.sql, where <version> is a 5-digit zero-padded
88    integer, <description> is a short textual description, and <scope> is a short string
89    identifying the database implementation.
90    """
91
92    _settings: Settings
93
94    def __init__(self, system: System) -> None:
95        self._settings = system.settings
96        self._opentelemetry_client = system.require(OpenTelemetryClient)
97        super().__init__(system)
98
99    @staticmethod
100    @abstractmethod
101    def migration_scope() -> str:
102        """The database implementation to use for migrations (e.g, sqlite, pgsql)"""
103        pass
104
105    @abstractmethod
106    def migration_dirs(self) -> Sequence[Traversable]:
107        """Directories containing the migration sequences that should be applied to this
108        DB."""
109        pass
110
111    @abstractmethod
112    def setup_migrations(self) -> None:
113        """Idempotently creates the migrations table"""
114        pass
115
116    @abstractmethod
117    def migrations_initialized(self) -> bool:
118        """Return true if the migrations table exists"""
119        pass
120
121    @abstractmethod
122    def db_migrations(self, dir: Traversable) -> Sequence[Migration]:
123        """Return a list of all migrations already applied to this database, from the
124        given source directory, in ascending order."""
125        pass
126
127    @abstractmethod
128    def apply_migration(self, cur: Cursor, migration: Migration) -> None:
129        """Apply a single migration to the database"""
130        pass
131
132    def initialize_migrations(self) -> None:
133        """Initialize migrations for this DB"""
134        migrate = self._settings.require("migrations")
135
136        if migrate == "validate":
137            self.validate_migrations()
138
139        if migrate == "apply":
140            self.apply_migrations()
141
142    @trace_method("MigratableDB.validate_migrations", OpenTelemetryGranularity.ALL)
143    def validate_migrations(self) -> None:
144        """Validate all migrations and throw an exception if there are any unapplied
145        migrations in the source repo."""
146        if not self.migrations_initialized():
147            raise UninitializedMigrationsError()
148        for dir in self.migration_dirs():
149            db_migrations = self.db_migrations(dir)
150            source_migrations = find_migrations(
151                dir,
152                self.migration_scope(),
153                self._settings.require("migrations_hash_algorithm"),
154            )
155            unapplied_migrations = verify_migration_sequence(
156                db_migrations, source_migrations
157            )
158            if len(unapplied_migrations) > 0:
159                version = unapplied_migrations[0]["version"]
160                raise UnappliedMigrationsError(dir=dir.name, version=version)
161
162    @trace_method("MigratableDB.apply_migrations", OpenTelemetryGranularity.ALL)
163    def apply_migrations(self) -> None:
164        """Validate existing migrations, and apply all new ones."""
165        self.setup_migrations()
166        for dir in self.migration_dirs():
167            db_migrations = self.db_migrations(dir)
168            source_migrations = find_migrations(
169                dir,
170                self.migration_scope(),
171                self._settings.require("migrations_hash_algorithm"),
172            )
173            unapplied_migrations = verify_migration_sequence(
174                db_migrations, source_migrations
175            )
176            with self.tx() as cur:
177                for migration in unapplied_migrations:
178                    self.apply_migration(cur, migration)
179
180
181# Format is <version>-<name>.<scope>.sql
182# e.g, 00001-users.sqlite.sql
183filename_regex = re.compile(r"(\d+)-(.+)\.(.+)\.sql")
184
185
186def _parse_migration_filename(
187    dir: str, filename: str, path: Traversable
188) -> MigrationFile:
189    """Parse a migration filename into a MigrationFile object"""
190    match = filename_regex.match(filename)
191    if match is None:
192        raise InvalidMigrationFilename("Invalid migration filename: " + filename)
193    version, _, scope = match.groups()
194    return {
195        "path": path,
196        "dir": dir,
197        "filename": filename,
198        "version": int(version),
199        "scope": scope,
200    }
201
202
203def verify_migration_sequence(
204    db_migrations: Sequence[Migration],
205    source_migrations: Sequence[Migration],
206) -> Sequence[Migration]:
207    """Given a list of migrations already applied to a database, and a list of
208    migrations from the source code, validate that the applied migrations are correct
209    and match the expected migrations.
210
211    Throws an exception if any migrations are missing, out of order, or if the source
212    hash does not match.
213
214    Returns a list of all unapplied migrations, or an empty list if all migrations are
215    applied and the database is up to date."""
216
217    for db_migration, source_migration in zip(db_migrations, source_migrations):
218        if db_migration["version"] != source_migration["version"]:
219            raise InconsistentVersionError(
220                dir=db_migration["dir"],
221                db_version=db_migration["version"],
222                source_version=source_migration["version"],
223            )
224
225        if db_migration["hash"] != source_migration["hash"]:
226            raise InconsistentHashError(
227                path=db_migration["dir"] + "/" + db_migration["filename"],
228                db_hash=db_migration["hash"],
229                source_hash=source_migration["hash"],
230            )
231
232    return source_migrations[len(db_migrations) :]
233
234
235def find_migrations(
236    dir: Traversable, scope: str, hash_alg: str = "md5"
237) -> Sequence[Migration]:
238    """Return a list of all migration present in the given directory, in ascending
239    order. Filter by scope."""
240    files = [
241        _parse_migration_filename(dir.name, t.name, t)
242        for t in dir.iterdir()
243        if t.name.endswith(".sql")
244    ]
245    files = list(filter(lambda f: f["scope"] == scope, files))
246    files = sorted(files, key=lambda f: f["version"])
247    return [_read_migration_file(f, hash_alg) for f in files]
248
249
250def _read_migration_file(file: MigrationFile, hash_alg: str) -> Migration:
251    """Read a migration file"""
252    if "path" not in file or not file["path"].is_file():
253        raise FileNotFoundError(
254            f"No migration file found for dir {file['dir']} with filename {file['filename']} and scope {file['scope']} at version {file['version']}"
255        )
256    sql = file["path"].read_text()
257
258    if hash_alg == "md5":
259        hash = (
260            hashlib.md5(sql.encode("utf-8"), usedforsecurity=False).hexdigest()
261            if sys.version_info >= (3, 9)
262            else hashlib.md5(sql.encode("utf-8")).hexdigest()
263        )
264    elif hash_alg == "sha256":
265        hash = hashlib.sha256(sql.encode("utf-8")).hexdigest()
266    else:
267        raise InvalidHashError(alg=hash_alg)
268
269    return {
270        "hash": hash,
271        "sql": sql,
272        "dir": file["dir"],
273        "filename": file["filename"],
274        "version": file["version"],
275        "scope": file["scope"],
276    }
277 
codekingpro/portable-devtools · Team Ai