pydantic-table 0.3.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,4 @@
1
+ from pydantic_table.__version__ import __version__
2
+
3
+ from pydantic_table.table_model.field import ColumnField as ColumnField
4
+ from pydantic_table.table_model.model import TableModel as TableModel
@@ -0,0 +1,3 @@
1
+ from importlib.metadata import version
2
+
3
+ __version__ = version("pydantic-table")
File without changes
@@ -0,0 +1,162 @@
1
+ import json
2
+ import os
3
+ from pathlib import Path
4
+ from typing import Any, Type
5
+
6
+ import inspect
7
+
8
+ from alembic import op
9
+
10
+ from pydantic_table.alembic.exceptions import ArchiveException
11
+ from pydantic_table.logger import logg
12
+ from pydantic_table.table_model.field import ColumnFieldInfo
13
+ from pydantic_table.table_model.model import TableModel
14
+
15
+ import sqlalchemy as sa
16
+ import pydantic_table.sqlalchemy as sap
17
+
18
+ SAVE_DATA = False
19
+
20
+
21
+ class Archive:
22
+ def __init__(self) -> None:
23
+ self._revision_filepath = self._get_revision_filepath()
24
+
25
+ versions = self._revision_filepath.parent
26
+ self._dir = versions / ".archive"
27
+
28
+ if not self._dir.exists():
29
+ os.makedirs(self._dir)
30
+
31
+ revision_filename = Path(self._revision_filepath).stem
32
+ self._archive_file = self._dir / f"{revision_filename}.json"
33
+ logg.debug(f"archive file: {self._archive_file}")
34
+
35
+ @property
36
+ def file_exists(self) -> bool:
37
+ return self._archive_file.exists()
38
+
39
+ def archive_table_model(self, sa_table: sa.Table):
40
+ """
41
+ Archive table schema.
42
+ """
43
+ # TODO: check if content exists, not if file exists
44
+ if self.file_exists:
45
+ logg.debug("-> already exists")
46
+ return
47
+
48
+ table_dict = {}
49
+ for name, column in sa_table.c.items():
50
+ table_dict[name] = self._build_column_info_dict(column)
51
+
52
+ archive_dict = {sa_table.name: table_dict}
53
+ self._save_archive_dict(archive_dict)
54
+
55
+ def read_column_fields(self, table_name: str) -> dict[str, ColumnFieldInfo]:
56
+ """
57
+ Read table model from archive.
58
+
59
+ Return map of column field infos for each column name (i.e. TableModel.column_fields() content)
60
+ """
61
+ column_fields = self._check_and_load_file_content(content=table_name)
62
+ ret = {
63
+ column_name: self._build_column_info(column_info)
64
+ for column_name, column_info in column_fields.items()
65
+ }
66
+ return ret
67
+
68
+ # TODO: multiple columns
69
+ def archive_column_info(self, table: Type[TableModel], column_name: str):
70
+ """
71
+ Archive column information.
72
+
73
+ Extract column field information from table (not table model, which does not have it anymore).
74
+ Extract path to the migration file (assumes drop_column is called from a migration)
75
+ Create archive directory if does not exist.
76
+ Serialize field information to be saved in archive JSON.
77
+ Get dict data from sa.Table to archive.
78
+ """
79
+
80
+ # TODO: check if content exists, not if file exists
81
+ if self.file_exists:
82
+ logg.debug("-> already exists")
83
+ return
84
+
85
+ engine = op.get_bind()
86
+
87
+ tb = sap.Table(table, autoload_with=engine)
88
+ sa_column = tb.c[column_name]
89
+ archive_dict: dict[str, dict[str, Any]] = {
90
+ column_name: {"info": self._build_column_info_dict(sa_column)}
91
+ }
92
+
93
+ if SAVE_DATA:
94
+ result = engine.execute(tb.select())
95
+ data_dict = [dict(row) for row in result.mappings()]
96
+ archive_dict[column_name]["data"] = data_dict
97
+
98
+ self._save_archive_dict(archive_dict)
99
+
100
+ def read_column_info(self, column: str) -> ColumnFieldInfo:
101
+ """
102
+ Get column schema from archive
103
+ """
104
+ column_dict = self._check_and_load_file_content(content=column)
105
+
106
+ if not "info" in column_dict:
107
+ raise ArchiveException(
108
+ f"Field 'info' not found for column {column} in archive file {self._archive_file}!"
109
+ )
110
+
111
+ ret = self._build_column_info(column_dict["info"])
112
+ return ret
113
+
114
+ def _build_column_info_dict(self, column: sa.Column) -> dict[str, Any]:
115
+ """
116
+ Build column field info dict based on sqlalchemy column.
117
+ """
118
+ column_info = sap.ColumnFieldInfo(column)
119
+ ret = column_info.as_dict()
120
+ ret["annotation"] = ret["annotation"].__name__
121
+ return ret
122
+
123
+ def _build_column_info(self, column_info: dict) -> ColumnFieldInfo:
124
+ info_dict = column_info.copy()
125
+
126
+ info_dict["annotation"] = eval(info_dict["annotation"])
127
+ ret = ColumnFieldInfo(**info_dict)
128
+ return ret
129
+
130
+ def _save_archive_dict(self, dct: dict[str, Any]):
131
+ logg.debug(f"archive dict: {dct}")
132
+
133
+ with open(self._archive_file, "w") as f:
134
+ json.dump(dct, f, indent=2)
135
+
136
+ def _check_and_load_file_content(self, content: str) -> dict:
137
+ if not self._archive_file.exists():
138
+ raise ArchiveException(f"Archive file {self._archive_file} not found!")
139
+
140
+ with open(self._archive_file) as f:
141
+ file_dict = json.load(f)
142
+
143
+ if not content in file_dict:
144
+ raise ArchiveException(
145
+ f"{content} does not exist in archive file {self._archive_file}!"
146
+ )
147
+
148
+ ret = file_dict[content]
149
+ return ret
150
+
151
+ def _get_revision_filepath(self) -> Path:
152
+ frame = inspect.currentframe()
153
+ assert frame is not None, "got null frame"
154
+ archive_caller = frame.f_back
155
+ assert archive_caller is not None, "no archive caller frame"
156
+ opp_caller = archive_caller.f_back
157
+ assert opp_caller is not None, "no opp caller frame"
158
+ revision_caller = opp_caller.f_back
159
+ assert revision_caller is not None, "no revision caller frame"
160
+
161
+ ret = Path(revision_caller.f_code.co_filename)
162
+ return ret
@@ -0,0 +1,5 @@
1
+ class PydanticTableAlembicException(Exception):
2
+ pass
3
+
4
+ class ArchiveException(PydanticTableAlembicException):
5
+ pass
@@ -0,0 +1,241 @@
1
+ """
2
+ op adaptors
3
+ """
4
+
5
+ from alembic import op
6
+ import sqlalchemy as sa
7
+ from typing import Any, Type
8
+
9
+ from pydantic_table.alembic.archive import Archive
10
+ from pydantic_table.alembic.exceptions import (
11
+ ArchiveException,
12
+ PydanticTableAlembicException,
13
+ )
14
+ from pydantic_table.logger import logg
15
+ import pydantic_table.sqlalchemy as sap
16
+
17
+ from pydantic_table.table_model.model import TableModel
18
+ from pydantic_table.utils import dict_as_str
19
+
20
+
21
+ def create_table(
22
+ table: Type[TableModel],
23
+ foreign_keys: dict[str, str] = {},
24
+ ):
25
+ """
26
+ Invoke alembic create table based on provided TableModel schema.
27
+
28
+ primary_keys (list[str]): list of primary key column names
29
+ foreign_keys (dict[str, str]): dictionary mapping {column name: foreign key column information}
30
+ Format of foreign key is foreign_table.column
31
+
32
+ Column names must correspond to TableModel field names.
33
+ """
34
+ archive = Archive()
35
+ if archive.file_exists:
36
+ try:
37
+ # TODO: table rename
38
+ column_fields = archive.read_column_fields(table.table_name())
39
+ except ArchiveException:
40
+ raise PydanticTableAlembicException(
41
+ f"Table {table.table_info()} schema has changed relative to this revision but archive file not found! Cannot create table."
42
+ )
43
+ else:
44
+ column_fields = table.column_fields()
45
+
46
+ sa_columns = [
47
+ sap.Column(name, column_info, foreign_key=foreign_keys.get(name, None))
48
+ for name, column_info in column_fields.items()
49
+ ]
50
+
51
+ op.create_table(table.table_name(), *sa_columns)
52
+
53
+
54
+ def drop_table(table: Type[TableModel]):
55
+ """
56
+ Invoke alembic drop table based on provided TableModel schema.
57
+
58
+ If detect extra columns or missing columns i.e. any schema change,
59
+ archive table schema.
60
+ """
61
+ sa_table = sap.Table(table, autoload_with=op.get_bind())
62
+ # TODO: mutual set difference
63
+ # TODO: detect any schema change in column info (changed default, changed nullability or primary)
64
+ model_has_new_columns = any(
65
+ col_name not in table.column_fields() for col_name in sa_table.c
66
+ )
67
+ model_is_missing_columns = any(
68
+ col_name not in sa_table.c for col_name in table.column_fields()
69
+ )
70
+ if model_has_new_columns or model_is_missing_columns:
71
+ Archive().archive_table_model(sa_table)
72
+
73
+ op.drop_table(table.table_name())
74
+
75
+
76
+ def update_where(table: Type[TableModel], values: dict[str, Any], **kwargs):
77
+ """
78
+ Update table with given values {column: value}.
79
+
80
+ kwargs in format column=value for the where condition
81
+ """
82
+ tb = sap.Table(table, autoload_with=op.get_bind())
83
+ condition = _get_condition(tb, **kwargs)
84
+ logg.debug(f"Values: {dict_as_str(values)}")
85
+ op.execute(tb.update().where(condition).values(values))
86
+
87
+
88
+ def add_column(
89
+ table: Type[TableModel],
90
+ name: str,
91
+ data: list[TableModel] | TableModel = [],
92
+ foreign_key: str | None = None,
93
+ ):
94
+ """
95
+ Add column to table.
96
+
97
+ Given name must be present in table column fields or in archive.
98
+ Fill column with given data; otherwise initialize to column default (must be present).
99
+
100
+ In case column is not nullable, no default, and data is given:
101
+ - create column first as nullable to prevent crash
102
+ - add given column data (use other columns for condition)
103
+ - set column back to nullable
104
+ """
105
+ archive = Archive()
106
+ if archive.file_exists:
107
+ try:
108
+ column_info = archive.read_column_info(name)
109
+ except ArchiveException:
110
+ raise PydanticTableAlembicException(
111
+ f"Column {name} not present in table {table.table_info()} or in archive! Cannot add."
112
+ )
113
+ else:
114
+ column_info = table.column_fields()[name]
115
+
116
+ sa_column = sap.Column(name, column_info, foreign_key=foreign_key)
117
+
118
+ logg.debug(f"Adding column - column info: {column_info}")
119
+ logg.debug(
120
+ f"default={column_info.default} default factory={column_info.default_factory} required={column_info.is_required()}"
121
+ )
122
+ # is required = neither default nor default factory are defined
123
+ if not column_info.nullable and column_info.is_required():
124
+ data_list = data if isinstance(data, list) else [data]
125
+ if len(data_list) == 0:
126
+ raise PydanticTableAlembicException(
127
+ f"Column {name} in {table.table_info()} is not nullable and no default was given. You provided no data."
128
+ f"Either provide data to add, or a default; or make column nullable"
129
+ )
130
+
131
+ sa_column.nullable = True
132
+ op.add_column(table.table_name(), sa_column)
133
+ for row in data_list:
134
+ logg.debug(f"Adding data row - {row}")
135
+ logg.debug(f"Column dump - {row.column_dump()}")
136
+ logg.debug(f"Data dump - {row.data_dump()}")
137
+ update_where(
138
+ table,
139
+ values={name: row.get(name)},
140
+ **row.column_dump(exclude={name: True}),
141
+ )
142
+ # op.alter_column(table.table_name(), name, nullable=False) - error
143
+ with op.batch_alter_table(table.table_name()) as batch_op:
144
+ batch_op.alter_column(name, nullable=False)
145
+ return
146
+
147
+ op.add_column(
148
+ table.table_name(), sap.Column(name, column_info, foreign_key=foreign_key)
149
+ )
150
+
151
+
152
+ def drop_column(table: Type[TableModel], name: str):
153
+ """
154
+ Drop column.
155
+
156
+ Drop column from the table.
157
+ Archive column if not present in TableModel anymore
158
+ (schema update removes column)
159
+ """
160
+ if name not in table.column_fields():
161
+ Archive().archive_column_info(table, name)
162
+
163
+ op.drop_column(table.table_name(), name)
164
+
165
+
166
+ def insert(rows: TableModel | list[TableModel]):
167
+ """
168
+ Insert given row(s) to the table corresponding to its schema.
169
+
170
+ Read table based on table name.
171
+ Invoke alembic execute() of table insert() with model dump.
172
+
173
+ Account for backwards compatibility with added columns or dropped columns.
174
+
175
+ Check existing columns in input that were not given (missing columns):
176
+ - if not present in table, ignore (past/future schema change dropped/added this column)
177
+ - if present in table, raise error (must be given but was not)
178
+
179
+ Ignore columns that are not present in table.
180
+
181
+ Check hidden extra columns in input (extra columns):
182
+ - if present in table, include (past/future schema change added/dropped this column)
183
+ - if not present in table, raise error (non-existing columns were given)
184
+ """
185
+ row_list = rows if isinstance(rows, list) else [rows]
186
+ for row in row_list:
187
+ table = sap.Table(row.table, autoload_with=op.get_bind())
188
+
189
+ for col_name in row.missing_columns:
190
+ if col_name in table.c:
191
+ raise PydanticTableAlembicException(
192
+ f"Row given for table {row.table_name()} is missing column {col_name}!"
193
+ )
194
+
195
+ data = row.column_dump()
196
+ logg.debug(f"Inserting data - column dump: {data}")
197
+ data = {col: val for col, val in data.items() if col in table.c}
198
+ logg.debug(f"Inserting data - only columns in table: {data}")
199
+
200
+ for col_name, col_value in row.extra_data.items():
201
+ if col_name not in table.c:
202
+ raise PydanticTableAlembicException(
203
+ f"Row given for table {row.table_info()} has extra column {col_name}!"
204
+ )
205
+ data[col_name] = col_value
206
+
207
+ logg.debug(f"Inserting data + extra columns: {data}")
208
+ op.execute(table.insert().values(data))
209
+
210
+
211
+ def delete_where(table: Type[TableModel], **kwargs):
212
+ """
213
+ Delete rows from table where columns have given values.
214
+
215
+ kwargs in format column=value for the where condition.
216
+ """
217
+ tb = sap.Table(table, autoload_with=op.get_bind())
218
+ condition = _get_condition(tb, **kwargs)
219
+ op.execute(tb.delete().where(condition))
220
+
221
+
222
+ def deep_delete(rows: TableModel | list[TableModel]):
223
+ """
224
+ Delete given row(s) from the table corresponding to its schema.
225
+
226
+ Read table based on table name.
227
+ Invoke alembic execute() with table delete() matching all fields in where().
228
+
229
+ Note that this operation takes a lot of time, but is the most secure as it deletes exactly the row.
230
+ """
231
+ row_list = rows if isinstance(rows, list) else [rows]
232
+ table = row_list[0].table
233
+
234
+ for row in row_list:
235
+ delete_where(table, **row.data_dump())
236
+
237
+
238
+ def _get_condition(tb: sa.Table, **kwargs) -> sa.ColumnElement[bool]:
239
+ logg.debug(f"Condition: {dict_as_str(kwargs)}")
240
+ condition = sa.and_(*[tb.c[column] == value for column, value in kwargs.items()])
241
+ return condition
@@ -0,0 +1,136 @@
1
+ import enum
2
+ import inspect
3
+ import logging
4
+ from pathlib import Path
5
+ from typing import Any
6
+ from dotenv import dotenv_values
7
+
8
+
9
+ class AnsiStyle(str, enum.Enum):
10
+ normal = "0"
11
+ bold = "1"
12
+ start = "\033["
13
+ end = "\033[0m"
14
+
15
+ def __str__(self) -> str:
16
+ return self.value
17
+
18
+
19
+ class AnsiColor(str, enum.Enum):
20
+ green = "32"
21
+ grey = "90"
22
+ red = "31"
23
+ yellow = "33"
24
+ white = "37"
25
+
26
+ def apply(self, message: Any, bold: bool = False) -> str:
27
+ """
28
+ To be used with color based
29
+ """
30
+ style = AnsiStyle.bold if bold else AnsiStyle.normal
31
+ ret = f"{AnsiStyle.start}{style};{self.value}m{message}{AnsiStyle.end}"
32
+ return ret
33
+
34
+ def bold(self, message: Any) -> str:
35
+ """
36
+ Shortcut for bold colored text
37
+ """
38
+ ret = self.apply(message, bold=True)
39
+ return ret
40
+
41
+ def __str__(self) -> str:
42
+ return self.value
43
+
44
+
45
+ class LevelFormatter(logging.Formatter):
46
+ def __init__(self, formats, default_fmt=None, datefmt=None):
47
+ super().__init__(datefmt=datefmt)
48
+ self.formats = {
49
+ level: logging.Formatter(fmt, datefmt=datefmt)
50
+ for level, fmt in formats.items()
51
+ }
52
+ self.default_formatter = logging.Formatter(
53
+ default_fmt or "%(message)s", datefmt=datefmt
54
+ )
55
+
56
+ def format(self, record):
57
+ formatter = self.formats.get(record.levelno, self.default_formatter)
58
+ return formatter.format(record)
59
+
60
+
61
+ class Logger:
62
+ def __init__(self, log_level=logging.INFO):
63
+ self._logger = logging.getLogger(__name__)
64
+
65
+ full_format = "%(asctime)s [%(levelname)s] %(classname)s.%(funcName)s:%(lineno)d - %(message)s"
66
+ short_format = "[%(levelname)s] %(classname)s - %(message)s"
67
+ no_format = ""
68
+ level_format = full_format if log_level == logging.DEBUG else no_format
69
+ formatter = LevelFormatter(
70
+ {
71
+ logging.DEBUG: full_format,
72
+ logging.ERROR: level_format,
73
+ logging.INFO: level_format,
74
+ logging.WARNING: level_format,
75
+ }
76
+ )
77
+
78
+ handler = logging.StreamHandler()
79
+ handler.setFormatter(formatter)
80
+ handler.setLevel(log_level)
81
+ self._logger.addHandler(handler)
82
+
83
+ self._logger.setLevel(log_level)
84
+ self._logger.propagate = False
85
+
86
+ def info(self, message: Any, header: bool = False):
87
+ color = AnsiColor.green if header else AnsiColor.white
88
+ return self._log(logging.INFO, color.apply(message, header))
89
+
90
+ def error(self, message: Any):
91
+ return self._log(logging.ERROR, AnsiColor.red.apply(message))
92
+
93
+ def warning(self, message: Any, important: bool = False):
94
+ if important:
95
+ message = f"! WARNING ! {message}"
96
+ return self._log(logging.WARNING, AnsiColor.yellow.apply(message, important))
97
+
98
+ def debug(self, message: Any):
99
+ return self._log(logging.DEBUG, AnsiColor.grey.apply(message))
100
+
101
+ def _log(self, level, message: Any):
102
+ """
103
+ Common log interface for info/error/warning/debug.
104
+
105
+ Auto-detect class name of caller.
106
+ Account for stack level to display correct funcName and lineno.
107
+ Skip 2 stack levels including this method,
108
+ the info/error/warning/debug method calling it,
109
+ """
110
+ self._logger.log(
111
+ level,
112
+ message,
113
+ extra={"classname": self._get_caller_class_name()},
114
+ stacklevel=3,
115
+ )
116
+
117
+ def _get_caller_class_name(self):
118
+ """
119
+ Determine class name of where logger is called from.
120
+
121
+ Obtain frame index 3 corresponding to actual caller
122
+ (0 = this method, 1 = _log internal method, 2 = info/warning/error/debug logger call).
123
+ Return class or module name of that frame.
124
+ """
125
+ frame = inspect.stack()[3].frame
126
+ if "self" in frame.f_locals:
127
+ return type(frame.f_locals["self"]).__name__
128
+ elif "cls" in frame.f_locals:
129
+ return frame.f_locals["cls"].__name__
130
+ return frame.f_globals.get("__name__", "-")
131
+
132
+
133
+ # NOTE: dotenv_values() in some cases yielded empty .env for unexplained reason
134
+ env_config = dotenv_values(Path.cwd() / ".env")
135
+ is_debug = env_config.get("DEBUG_PYDANTIC_TABLE", "").lower() in ("true", "1")
136
+ logg = Logger(log_level=logging.DEBUG if is_debug else logging.INFO)
File without changes
@@ -0,0 +1,3 @@
1
+ from pydantic_table.sqlalchemy.base import BaseMeta as BaseMeta
2
+ from pydantic_table.sqlalchemy.column import Column as Column, ColumnFieldInfo as ColumnFieldInfo
3
+ from pydantic_table.sqlalchemy.table import Table as Table
@@ -0,0 +1,57 @@
1
+ import enum
2
+ from typing import Any, Type
3
+
4
+ from sqlalchemy import Float, Integer, String
5
+ from sqlalchemy.orm import DeclarativeBase, mapped_column
6
+ from sqlalchemy.orm.decl_api import DeclarativeAttributeIntercept
7
+
8
+ from pydantic_table.table_model.model import TableModel
9
+
10
+
11
+ class BaseFieldType(enum.Enum):
12
+ int = Integer()
13
+ float = Float(32)
14
+ str = String(255)
15
+
16
+ @classmethod
17
+ def from_type(cls, type: type):
18
+ return cls[type.__name__].value
19
+
20
+
21
+ class BaseMeta(DeclarativeAttributeIntercept):
22
+ """
23
+ Adaptor Metaclass for creating DeclarativeBase based on pydantic-table TableModel.
24
+
25
+ Translates TableModel field:
26
+ - table_name_ --> __tablename__
27
+ - annotation -> sqlalchemy type (Integer, Float, String)
28
+ - default=None -> nullable=True
29
+ """
30
+
31
+ def __new__(
32
+ cls,
33
+ name: str,
34
+ bases: tuple[type, ...],
35
+ namespace: dict[str, Any],
36
+ /,
37
+ model: Type[TableModel],
38
+ **kwds: Any,
39
+ ):
40
+ namespace["__tablename__"] = model.table_name()
41
+
42
+ for column_name, column_info in model.column_fields().items():
43
+ if column_info.annotation is None:
44
+ # TODO: model validator (this should basically be assert level)
45
+ raise ValueError(
46
+ f"pydantic model fields must be annotated for pydantic-table sqlalchemy adaptor!\n{column_info}"
47
+ )
48
+
49
+ namespace[column_name] = mapped_column(
50
+ BaseFieldType.from_type(column_info.annotation),
51
+ nullable=column_info.nullable,
52
+ primary_key=column_info.primary_key,
53
+ )
54
+
55
+ x = super().__new__(cls, name, bases, namespace, **kwds)
56
+
57
+ return x