codekingpro/portable-devtools
114k
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 