autostorage 0.0.7__tar.gz
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.
- autostorage-0.0.7/PKG-INFO +10 -0
- autostorage-0.0.7/pyproject.toml +101 -0
- autostorage-0.0.7/src/autostorage/__init__.py +36 -0
- autostorage-0.0.7/src/autostorage/calcn/__init__.py +18 -0
- autostorage-0.0.7/src/autostorage/calcn/core.py +107 -0
- autostorage-0.0.7/src/autostorage/calcn/registry.py +113 -0
- autostorage-0.0.7/src/autostorage/calcn/util.py +67 -0
- autostorage-0.0.7/src/autostorage/database.py +200 -0
- autostorage-0.0.7/src/autostorage/models/__init__.py +45 -0
- autostorage-0.0.7/src/autostorage/models/calculation.py +358 -0
- autostorage-0.0.7/src/autostorage/models/data.py +49 -0
- autostorage-0.0.7/src/autostorage/models/geometry.py +141 -0
- autostorage-0.0.7/src/autostorage/models/links.py +122 -0
- autostorage-0.0.7/src/autostorage/models/listeners.py +97 -0
- autostorage-0.0.7/src/autostorage/models/optional.py +57 -0
- autostorage-0.0.7/src/autostorage/models/reaction.py +116 -0
- autostorage-0.0.7/src/autostorage/models/stationary.py +143 -0
- autostorage-0.0.7/src/autostorage/types/__init__.py +21 -0
- autostorage-0.0.7/src/autostorage/types/fields.py +10 -0
- autostorage-0.0.7/src/autostorage/types/sqlalchemy.py +57 -0
- autostorage-0.0.7/src/autostorage/utils/__init__.py +5 -0
- autostorage-0.0.7/src/autostorage/utils/sql_model.py +62 -0
|
@@ -0,0 +1,10 @@
|
|
|
1
|
+
Metadata-Version: 2.3
|
|
2
|
+
Name: autostorage
|
|
3
|
+
Version: 0.0.7
|
|
4
|
+
Author: Andreas V. Copan
|
|
5
|
+
Author-email: Andreas V. Copan <avcopan@uga.edu>
|
|
6
|
+
Requires-Dist: automol==0.0.10
|
|
7
|
+
Requires-Dist: pint>=0.25.2
|
|
8
|
+
Requires-Dist: qcdata>=0.17.0
|
|
9
|
+
Requires-Dist: sqlmodel>=0.0.31
|
|
10
|
+
Requires-Python: >=3.12
|
|
@@ -0,0 +1,101 @@
|
|
|
1
|
+
[project]
|
|
2
|
+
name = "autostorage"
|
|
3
|
+
version = "0.0.7"
|
|
4
|
+
authors = [{name = "Andreas V. Copan", email = "avcopan@uga.edu"}]
|
|
5
|
+
requires-python = ">= 3.12"
|
|
6
|
+
dependencies = [
|
|
7
|
+
"automol==0.0.10",
|
|
8
|
+
"pint>=0.25.2",
|
|
9
|
+
"qcdata>=0.17.0",
|
|
10
|
+
"sqlmodel>=0.0.31",
|
|
11
|
+
]
|
|
12
|
+
|
|
13
|
+
[build-system]
|
|
14
|
+
build-backend = "uv_build"
|
|
15
|
+
requires = ["uv_build"]
|
|
16
|
+
|
|
17
|
+
# UV build configurations
|
|
18
|
+
[tool.uv.build-backend]
|
|
19
|
+
module-name = "autostorage"
|
|
20
|
+
|
|
21
|
+
# Ruff configurations
|
|
22
|
+
[tool.ruff]
|
|
23
|
+
exclude = [
|
|
24
|
+
"docs",
|
|
25
|
+
"**/*.ipynb",
|
|
26
|
+
]
|
|
27
|
+
|
|
28
|
+
[tool.ruff.lint]
|
|
29
|
+
select = ["ALL"]
|
|
30
|
+
ignore = [
|
|
31
|
+
"COM812", # missing trailing comma (handled by format)
|
|
32
|
+
"RUF022", # `__all__` is not sorted
|
|
33
|
+
"TID252", # relative import
|
|
34
|
+
"D203", # confict: incorrect blank line before class (uses D211 instead)
|
|
35
|
+
"D213", # conflict: multi-line summary second line (uses D212 instead)
|
|
36
|
+
]
|
|
37
|
+
|
|
38
|
+
[tool.ruff.lint.per-file-ignores]
|
|
39
|
+
"tests/**.py" = ["S101"]
|
|
40
|
+
|
|
41
|
+
[tool.ruff.lint.pydocstyle]
|
|
42
|
+
convention = "numpy"
|
|
43
|
+
|
|
44
|
+
# Ty configurations
|
|
45
|
+
[tool.ty.environment]
|
|
46
|
+
python = ".pixi/envs/dev/bin/python"
|
|
47
|
+
|
|
48
|
+
[tool.ty.src]
|
|
49
|
+
include = [
|
|
50
|
+
"src",
|
|
51
|
+
"tests",
|
|
52
|
+
]
|
|
53
|
+
|
|
54
|
+
# Import-linter configurations
|
|
55
|
+
[tool.importlinter]
|
|
56
|
+
root_package = "autostorage"
|
|
57
|
+
exclude_type_checking_imports = true
|
|
58
|
+
|
|
59
|
+
[[tool.importlinter.contracts]]
|
|
60
|
+
name = "AutoStore Layering"
|
|
61
|
+
type = "layers"
|
|
62
|
+
layers = [
|
|
63
|
+
"autostorage.database",
|
|
64
|
+
"autostorage.models",
|
|
65
|
+
"autostorage.calcn",
|
|
66
|
+
]
|
|
67
|
+
|
|
68
|
+
[[tool.importlinter.contracts]]
|
|
69
|
+
name = "AutoStore Models Layering"
|
|
70
|
+
type = "layers"
|
|
71
|
+
layers = [
|
|
72
|
+
"autostorage.models.listeners",
|
|
73
|
+
"autostorage.models.reaction",
|
|
74
|
+
"autostorage.models.stationary",
|
|
75
|
+
"autostorage.models.data",
|
|
76
|
+
"autostorage.models.calculation",
|
|
77
|
+
"autostorage.models.geometry",
|
|
78
|
+
"autostorage.models.links",
|
|
79
|
+
]
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
# Pytest configurations
|
|
83
|
+
[tool.pytest.ini_options]
|
|
84
|
+
testpaths = ["tests", "src"]
|
|
85
|
+
addopts = [
|
|
86
|
+
"--doctest-modules",
|
|
87
|
+
"--cov=autostorage",
|
|
88
|
+
"--cov-report=term-missing",
|
|
89
|
+
"--cov-report=html",
|
|
90
|
+
]
|
|
91
|
+
doctest_optionflags = "NORMALIZE_WHITESPACE"
|
|
92
|
+
|
|
93
|
+
# Pytest-cov configurations
|
|
94
|
+
[tool.coverage.run]
|
|
95
|
+
branch = true
|
|
96
|
+
source = ["autostorage"]
|
|
97
|
+
|
|
98
|
+
[tool.coverage.report]
|
|
99
|
+
show_missing = true
|
|
100
|
+
skip_covered = true
|
|
101
|
+
fail_under = 80
|
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
"""Interface for database storage."""
|
|
2
|
+
|
|
3
|
+
__version__ = "0.0.7"
|
|
4
|
+
|
|
5
|
+
from . import models
|
|
6
|
+
from .calcn import Calculation
|
|
7
|
+
from .database import Database
|
|
8
|
+
from .models import (
|
|
9
|
+
CalculationGeometryLink,
|
|
10
|
+
CalculationRow,
|
|
11
|
+
EnergyRow,
|
|
12
|
+
GeometryRow,
|
|
13
|
+
StageRow,
|
|
14
|
+
StationaryPointRow,
|
|
15
|
+
StationaryStageLink,
|
|
16
|
+
StepRow,
|
|
17
|
+
) # import core Row objects
|
|
18
|
+
from .types import Role
|
|
19
|
+
from .utils import verify_single_iteration
|
|
20
|
+
|
|
21
|
+
__all__ = [
|
|
22
|
+
"models",
|
|
23
|
+
"qc",
|
|
24
|
+
"CalculationGeometryLink",
|
|
25
|
+
"Calculation",
|
|
26
|
+
"Database",
|
|
27
|
+
"CalculationRow",
|
|
28
|
+
"EnergyRow",
|
|
29
|
+
"GeometryRow",
|
|
30
|
+
"StageRow",
|
|
31
|
+
"StationaryPointRow",
|
|
32
|
+
"StationaryStageLink",
|
|
33
|
+
"StepRow",
|
|
34
|
+
"Role",
|
|
35
|
+
"verify_single_iteration",
|
|
36
|
+
]
|
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
"""Calculation data."""
|
|
2
|
+
|
|
3
|
+
from .core import Calculation, project, projected_hash
|
|
4
|
+
from .registry import HashRegistry, calculation_hash, hash_registry
|
|
5
|
+
from .util import CalculationDict, KeywordDict, hash_from_dict, project_keywords
|
|
6
|
+
|
|
7
|
+
__all__ = [
|
|
8
|
+
"Calculation",
|
|
9
|
+
"project",
|
|
10
|
+
"projected_hash",
|
|
11
|
+
"HashRegistry",
|
|
12
|
+
"calculation_hash",
|
|
13
|
+
"hash_registry",
|
|
14
|
+
"CalculationDict",
|
|
15
|
+
"KeywordDict",
|
|
16
|
+
"hash_from_dict",
|
|
17
|
+
"project_keywords",
|
|
18
|
+
]
|
|
@@ -0,0 +1,107 @@
|
|
|
1
|
+
"""Calculation metadata."""
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
from pydantic import BaseModel, Field
|
|
6
|
+
|
|
7
|
+
from .util import CalculationDict, hash_from_dict, project_keywords
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class Calculation(BaseModel):
|
|
11
|
+
"""Calculation input parameters and metadata.
|
|
12
|
+
|
|
13
|
+
Attributes
|
|
14
|
+
----------
|
|
15
|
+
# - Program Input -------
|
|
16
|
+
program
|
|
17
|
+
Quantum chemistry program used (psi4, ORCA, ...)
|
|
18
|
+
program_keywords
|
|
19
|
+
(Optional) Quantum chemistry program keywords.
|
|
20
|
+
super_program
|
|
21
|
+
(Optional) Geometry optimizer program (geomeTRIC, ...).
|
|
22
|
+
super_keywords
|
|
23
|
+
(Optional) Geometry optimizer keywords.
|
|
24
|
+
cmdline_args
|
|
25
|
+
(Optional) Command line arguments.
|
|
26
|
+
input
|
|
27
|
+
(Optional) Input file. [ PLACEHOLDER ]
|
|
28
|
+
files
|
|
29
|
+
(Optional) Additional input files. [ PLACEHOLDER ]
|
|
30
|
+
# - Methods -------------
|
|
31
|
+
calc_type
|
|
32
|
+
Calculation type (energy, optimization, ...)
|
|
33
|
+
method
|
|
34
|
+
Computational method (B3LYP, MP2, ...)
|
|
35
|
+
basis
|
|
36
|
+
(Optional) Basis set.
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
# - Program Input -------
|
|
40
|
+
program: str
|
|
41
|
+
program_keywords: dict[str, Any] = Field(default_factory=dict)
|
|
42
|
+
super_program: str | None = Field(default=None)
|
|
43
|
+
super_keywords: dict[str, Any] = Field(default_factory=dict)
|
|
44
|
+
cmdline_args: list[str] = Field(default_factory=list)
|
|
45
|
+
|
|
46
|
+
# - Methods -------------
|
|
47
|
+
calc_type: str | None = Field(default=None)
|
|
48
|
+
method: str
|
|
49
|
+
basis: str | None = Field(default=None)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def projected_hash(calc: Calculation, template: Calculation | CalculationDict) -> str:
|
|
53
|
+
"""
|
|
54
|
+
Project calculation onto template and generate hash.
|
|
55
|
+
|
|
56
|
+
Parameters
|
|
57
|
+
----------
|
|
58
|
+
calc
|
|
59
|
+
Calculation metadata.
|
|
60
|
+
template
|
|
61
|
+
Calculation metadata template.
|
|
62
|
+
|
|
63
|
+
Returns
|
|
64
|
+
-------
|
|
65
|
+
Hash string.
|
|
66
|
+
"""
|
|
67
|
+
calc_dct = project(calc, template)
|
|
68
|
+
return hash_from_dict(calc_dct)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def project(
|
|
72
|
+
calc: Calculation, template: Calculation | CalculationDict
|
|
73
|
+
) -> CalculationDict:
|
|
74
|
+
"""
|
|
75
|
+
Project calculation onto template.
|
|
76
|
+
|
|
77
|
+
Parameters
|
|
78
|
+
----------
|
|
79
|
+
calc
|
|
80
|
+
Calculation metadata.
|
|
81
|
+
template
|
|
82
|
+
Calculation metadata template.
|
|
83
|
+
|
|
84
|
+
Returns
|
|
85
|
+
-------
|
|
86
|
+
Projected calculation dictionary.
|
|
87
|
+
"""
|
|
88
|
+
# Dump template to dictionary
|
|
89
|
+
template_dct = (
|
|
90
|
+
template.model_dump(exclude_unset=True)
|
|
91
|
+
if isinstance(template, Calculation)
|
|
92
|
+
else template
|
|
93
|
+
)
|
|
94
|
+
# Work on a deep copy to avoid accidental modifications
|
|
95
|
+
calc_copy = calc.model_copy(deep=True)
|
|
96
|
+
# Project program_keywords if 'keywords' is in the template
|
|
97
|
+
if "program_keywords" in template_dct:
|
|
98
|
+
calc_copy.program_keywords = project_keywords(
|
|
99
|
+
calc_copy.program_keywords,
|
|
100
|
+
template=template_dct.get("program_keywords", {}),
|
|
101
|
+
)
|
|
102
|
+
if "super_keywords" in template_dct:
|
|
103
|
+
calc_copy.super_keywords = project_keywords(
|
|
104
|
+
calc_copy.super_keywords, template=template_dct.get("super_keywords", {})
|
|
105
|
+
)
|
|
106
|
+
# Include fields from template
|
|
107
|
+
return calc_copy.model_dump(exclude_none=True, include=set(template_dct.keys()))
|
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
"""Calculation hash registry."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Callable
|
|
4
|
+
|
|
5
|
+
from .core import Calculation, projected_hash
|
|
6
|
+
from .util import hash_from_dict
|
|
7
|
+
|
|
8
|
+
HashFunc = Callable[[Calculation], str]
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class HashRegistry:
|
|
12
|
+
"""Hash registry."""
|
|
13
|
+
|
|
14
|
+
def __init__(self) -> None:
|
|
15
|
+
"""Initialize hash registry."""
|
|
16
|
+
self._registry: dict[str, HashFunc] = {}
|
|
17
|
+
|
|
18
|
+
def register(self, name: str) -> Callable[[HashFunc], HashFunc]:
|
|
19
|
+
"""Register function in hash registry (decorator)."""
|
|
20
|
+
|
|
21
|
+
def decorator(func: HashFunc) -> HashFunc:
|
|
22
|
+
if name in self._registry:
|
|
23
|
+
msg = f"Hash '{name}' is already registered"
|
|
24
|
+
raise ValueError(msg)
|
|
25
|
+
self._registry[name] = func
|
|
26
|
+
return func
|
|
27
|
+
|
|
28
|
+
return decorator
|
|
29
|
+
|
|
30
|
+
def get(self, name: str) -> HashFunc:
|
|
31
|
+
"""Get registered hash by name."""
|
|
32
|
+
try:
|
|
33
|
+
return self._registry[name]
|
|
34
|
+
except KeyError as err:
|
|
35
|
+
msg = f"Unknown hash type '{name}'. Available: {sorted(self._registry)}"
|
|
36
|
+
raise KeyError(msg) from err
|
|
37
|
+
|
|
38
|
+
def available(self) -> tuple[str, ...]:
|
|
39
|
+
"""Get registered hash names."""
|
|
40
|
+
return tuple(self._registry)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
hash_registry = HashRegistry()
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
@hash_registry.register("full")
|
|
47
|
+
def hash_full(calc: Calculation) -> str:
|
|
48
|
+
"""
|
|
49
|
+
Generate maximal hash for calculation.
|
|
50
|
+
|
|
51
|
+
Parameters
|
|
52
|
+
----------
|
|
53
|
+
calc
|
|
54
|
+
Calculation metadata.
|
|
55
|
+
|
|
56
|
+
Returns
|
|
57
|
+
-------
|
|
58
|
+
Hash string.
|
|
59
|
+
"""
|
|
60
|
+
input_fields = {
|
|
61
|
+
"program",
|
|
62
|
+
"program_keywords",
|
|
63
|
+
"super_program",
|
|
64
|
+
"super_keywords",
|
|
65
|
+
"cmdline_args",
|
|
66
|
+
"calc_type",
|
|
67
|
+
"method",
|
|
68
|
+
"basis",
|
|
69
|
+
}
|
|
70
|
+
calc_dct = calc.model_dump(include=input_fields)
|
|
71
|
+
return hash_from_dict(calc_dct)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
@hash_registry.register("minimal")
|
|
75
|
+
def hash_minimal(calc: Calculation) -> str:
|
|
76
|
+
"""
|
|
77
|
+
Generate minimal hash for calculation.
|
|
78
|
+
|
|
79
|
+
Parameters
|
|
80
|
+
----------
|
|
81
|
+
calc
|
|
82
|
+
Calculation metadata.
|
|
83
|
+
|
|
84
|
+
Returns
|
|
85
|
+
-------
|
|
86
|
+
Hash string.
|
|
87
|
+
"""
|
|
88
|
+
template = {
|
|
89
|
+
"program": "PROGRAM",
|
|
90
|
+
"method": "METHOD",
|
|
91
|
+
"basis": "BASIS",
|
|
92
|
+
}
|
|
93
|
+
return projected_hash(calc, template)
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def calculation_hash(calc: Calculation, name: str = "minimal") -> str:
|
|
97
|
+
"""
|
|
98
|
+
Hash calculation metadata using named hash from registry.
|
|
99
|
+
|
|
100
|
+
Parameters
|
|
101
|
+
----------
|
|
102
|
+
calc
|
|
103
|
+
Calculation metadata.
|
|
104
|
+
name
|
|
105
|
+
Hash registry name, e.g. "full" or "minimal". Use
|
|
106
|
+
`hash_registry.available()` to get list of available names.
|
|
107
|
+
|
|
108
|
+
Returns
|
|
109
|
+
-------
|
|
110
|
+
Hash string.
|
|
111
|
+
"""
|
|
112
|
+
func = hash_registry.get(name)
|
|
113
|
+
return func(calc)
|
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
"""Calculation data utilities."""
|
|
2
|
+
|
|
3
|
+
import hashlib
|
|
4
|
+
import json
|
|
5
|
+
from typing import Any, Union
|
|
6
|
+
|
|
7
|
+
KeywordDict = dict[str, Union[str, "KeywordDict", None]]
|
|
8
|
+
CalculationDict = dict[str, Any]
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def hash_from_dict(calc_dct: CalculationDict) -> str:
|
|
12
|
+
"""
|
|
13
|
+
Generate hash from calculation dictionary.
|
|
14
|
+
|
|
15
|
+
Parameters
|
|
16
|
+
----------
|
|
17
|
+
calc_dct
|
|
18
|
+
Calculation dictionary.
|
|
19
|
+
|
|
20
|
+
Returns
|
|
21
|
+
-------
|
|
22
|
+
Hash string.
|
|
23
|
+
"""
|
|
24
|
+
calc_json = json.dumps(
|
|
25
|
+
calc_dct, sort_keys=True, ensure_ascii=True, separators=(",", ":")
|
|
26
|
+
).encode("utf-8")
|
|
27
|
+
return hashlib.sha256(calc_json).hexdigest()
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def project_keywords(keywords: KeywordDict, template: object) -> KeywordDict:
|
|
31
|
+
"""
|
|
32
|
+
Project keywords dictionary onto template.
|
|
33
|
+
|
|
34
|
+
Parameters
|
|
35
|
+
----------
|
|
36
|
+
keywords
|
|
37
|
+
Keywords dictionary.
|
|
38
|
+
template
|
|
39
|
+
Keywords dictionary template.
|
|
40
|
+
|
|
41
|
+
Returns
|
|
42
|
+
-------
|
|
43
|
+
Projected keywords dictionary.
|
|
44
|
+
|
|
45
|
+
Raises
|
|
46
|
+
------
|
|
47
|
+
TypeError
|
|
48
|
+
If keywords template is not a dictionary.
|
|
49
|
+
TypeError
|
|
50
|
+
If keywords template keys are not strings.
|
|
51
|
+
"""
|
|
52
|
+
if not isinstance(template, dict):
|
|
53
|
+
msg = "Keywords template must be a dictionary."
|
|
54
|
+
raise TypeError(msg)
|
|
55
|
+
|
|
56
|
+
projected_dict: dict[str, object] = {}
|
|
57
|
+
for key, val in template.items():
|
|
58
|
+
if not isinstance(key, str):
|
|
59
|
+
msg = "Keywords template keys must be strings."
|
|
60
|
+
raise TypeError(msg)
|
|
61
|
+
|
|
62
|
+
if key in keywords:
|
|
63
|
+
if isinstance(val, dict) and isinstance(keywords[key], dict):
|
|
64
|
+
projected_dict[key] = project_keywords(keywords[key], val)
|
|
65
|
+
else:
|
|
66
|
+
projected_dict[key] = keywords[key]
|
|
67
|
+
return projected_dict
|
|
@@ -0,0 +1,200 @@
|
|
|
1
|
+
"""Database connection."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Iterator
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
|
|
6
|
+
from sqlalchemy.orm import selectinload
|
|
7
|
+
from sqlmodel import Session, SQLModel, create_engine, select
|
|
8
|
+
|
|
9
|
+
from .models import * # noqa: F403
|
|
10
|
+
from .types import SQLModelT
|
|
11
|
+
from .utils import row_to_dict
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class Database:
|
|
15
|
+
"""
|
|
16
|
+
Database connection manager.
|
|
17
|
+
|
|
18
|
+
Attributes
|
|
19
|
+
----------
|
|
20
|
+
path
|
|
21
|
+
Path to SQLite database file.
|
|
22
|
+
engine
|
|
23
|
+
SQLAlchemy engine instance.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
def __init__(self, path: str | Path, *, echo: bool = False) -> None:
|
|
27
|
+
"""
|
|
28
|
+
Initialize database connection manager.
|
|
29
|
+
|
|
30
|
+
Parameters
|
|
31
|
+
----------
|
|
32
|
+
path
|
|
33
|
+
Path to the SQLite database file.
|
|
34
|
+
echo, optional
|
|
35
|
+
If True, SQL statements will be logged to the standard output.
|
|
36
|
+
If False, no logging is performed.
|
|
37
|
+
"""
|
|
38
|
+
self.path = Path(path)
|
|
39
|
+
self.engine = create_engine(f"sqlite:///{self.path}", echo=echo)
|
|
40
|
+
SQLModel.metadata.create_all(self.engine)
|
|
41
|
+
|
|
42
|
+
def session(self) -> Session:
|
|
43
|
+
"""Create a new database session."""
|
|
44
|
+
return Session(self.engine)
|
|
45
|
+
|
|
46
|
+
def add(
|
|
47
|
+
self,
|
|
48
|
+
row: SQLModelT,
|
|
49
|
+
*,
|
|
50
|
+
eager_load: bool = False,
|
|
51
|
+
) -> SQLModelT:
|
|
52
|
+
"""
|
|
53
|
+
Add row to database.
|
|
54
|
+
|
|
55
|
+
Parameters
|
|
56
|
+
----------
|
|
57
|
+
row
|
|
58
|
+
Instance of a database model class.
|
|
59
|
+
eager_load
|
|
60
|
+
If True, fully eager loads sqlmodel relationships with model return.
|
|
61
|
+
|
|
62
|
+
Returns
|
|
63
|
+
-------
|
|
64
|
+
Updated row instance.
|
|
65
|
+
|
|
66
|
+
Raises
|
|
67
|
+
------
|
|
68
|
+
SQLAlchemyError
|
|
69
|
+
Database row failed to write.
|
|
70
|
+
"""
|
|
71
|
+
with self.session() as session:
|
|
72
|
+
session.add(row)
|
|
73
|
+
session.commit()
|
|
74
|
+
session.refresh(row)
|
|
75
|
+
|
|
76
|
+
if eager_load: # Add eager loading for all relationships
|
|
77
|
+
model = type(row)
|
|
78
|
+
statement = select(model)
|
|
79
|
+
for rel_name in model.__sqlmodel_relationships__:
|
|
80
|
+
statement = statement.options(
|
|
81
|
+
selectinload(getattr(model, rel_name))
|
|
82
|
+
)
|
|
83
|
+
matches = session.exec(statement).first()
|
|
84
|
+
if not matches:
|
|
85
|
+
msg = f"{row = } did not add to database."
|
|
86
|
+
raise RuntimeError(msg)
|
|
87
|
+
return matches
|
|
88
|
+
|
|
89
|
+
return row
|
|
90
|
+
|
|
91
|
+
def delete(self, row: SQLModelT) -> None:
|
|
92
|
+
"""
|
|
93
|
+
Delete a row from the database.
|
|
94
|
+
|
|
95
|
+
Parameters
|
|
96
|
+
----------
|
|
97
|
+
row
|
|
98
|
+
Instance of a database model class.
|
|
99
|
+
"""
|
|
100
|
+
with self.session() as session:
|
|
101
|
+
session.delete(row)
|
|
102
|
+
session.commit()
|
|
103
|
+
|
|
104
|
+
def find(
|
|
105
|
+
self,
|
|
106
|
+
row: SQLModelT,
|
|
107
|
+
*,
|
|
108
|
+
eager_load: bool = False,
|
|
109
|
+
exclude_defaults: bool = True,
|
|
110
|
+
exclude_id: bool = False,
|
|
111
|
+
) -> Iterator[SQLModelT]:
|
|
112
|
+
"""
|
|
113
|
+
Find matching rows in database.
|
|
114
|
+
|
|
115
|
+
If no matches, adds and yields row instance.
|
|
116
|
+
|
|
117
|
+
Parameters
|
|
118
|
+
----------
|
|
119
|
+
row
|
|
120
|
+
Instance of a database model class.
|
|
121
|
+
session
|
|
122
|
+
(Optional) Instance of an active session.
|
|
123
|
+
eager_load
|
|
124
|
+
If True, fully eager loads sqlmodel relationships with model return.
|
|
125
|
+
exclude_defaults
|
|
126
|
+
If True, exclude default values from model dump.
|
|
127
|
+
exclude_id
|
|
128
|
+
If True, exclude id field from model dump (if applicable).
|
|
129
|
+
|
|
130
|
+
Yields
|
|
131
|
+
------
|
|
132
|
+
Instance of a database "model".
|
|
133
|
+
"""
|
|
134
|
+
data = row_to_dict(
|
|
135
|
+
row, exclude_defaults=exclude_defaults, exclude_id=exclude_id
|
|
136
|
+
)
|
|
137
|
+
with self.session() as session:
|
|
138
|
+
model = type(row)
|
|
139
|
+
statement = select(model)
|
|
140
|
+
|
|
141
|
+
for k, v in data.items():
|
|
142
|
+
statement = statement.where(getattr(model, k) == v)
|
|
143
|
+
|
|
144
|
+
if eager_load: # Add eager loading for all relationships
|
|
145
|
+
for rel_name in model.__sqlmodel_relationships__:
|
|
146
|
+
statement = statement.options(
|
|
147
|
+
selectinload(getattr(model, rel_name))
|
|
148
|
+
)
|
|
149
|
+
|
|
150
|
+
yield from session.exec(statement)
|
|
151
|
+
|
|
152
|
+
def find_or_add(
|
|
153
|
+
self,
|
|
154
|
+
row: SQLModelT,
|
|
155
|
+
*,
|
|
156
|
+
eager_load: bool = False,
|
|
157
|
+
exclude_defaults: bool = True,
|
|
158
|
+
exclude_id: bool = False,
|
|
159
|
+
) -> Iterator[SQLModelT]:
|
|
160
|
+
"""
|
|
161
|
+
Find matching rows in database.
|
|
162
|
+
|
|
163
|
+
If no matches, adds and yields row instance.
|
|
164
|
+
|
|
165
|
+
Parameters
|
|
166
|
+
----------
|
|
167
|
+
row
|
|
168
|
+
Instance of a database model class.
|
|
169
|
+
eager_load
|
|
170
|
+
If True, fully eager loads sqlmodel relationships with model return.
|
|
171
|
+
exclude_defaults
|
|
172
|
+
If True, exclude default values from model dump.
|
|
173
|
+
exclude_id
|
|
174
|
+
If True, exclude id field from model dump (if applicable).
|
|
175
|
+
|
|
176
|
+
Yields
|
|
177
|
+
------
|
|
178
|
+
Instance of a database "model".
|
|
179
|
+
"""
|
|
180
|
+
# Flag to avoid loading the whole iterator into memory.
|
|
181
|
+
found_any = False
|
|
182
|
+
|
|
183
|
+
for matching_row in self.find(
|
|
184
|
+
row,
|
|
185
|
+
eager_load=eager_load,
|
|
186
|
+
exclude_defaults=exclude_defaults,
|
|
187
|
+
exclude_id=exclude_id,
|
|
188
|
+
):
|
|
189
|
+
found_any = True
|
|
190
|
+
yield matching_row
|
|
191
|
+
|
|
192
|
+
if not found_any:
|
|
193
|
+
yield self.add(row=row)
|
|
194
|
+
|
|
195
|
+
def close(self) -> None:
|
|
196
|
+
"""Close the database connection.
|
|
197
|
+
|
|
198
|
+
Seems to be needed only for testing with in-memory databases.
|
|
199
|
+
"""
|
|
200
|
+
self.engine.dispose()
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
"""SQLModel row definitions.
|
|
2
|
+
|
|
3
|
+
Layers to avoid circular imports at load time:
|
|
4
|
+
links
|
|
5
|
+
→ geometry
|
|
6
|
+
→ calculation
|
|
7
|
+
→ data
|
|
8
|
+
→ stationary
|
|
9
|
+
→ reaction
|
|
10
|
+
→ listeners
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from . import listeners # 1st registers @event.listens_for # noqa: F401
|
|
14
|
+
from .calculation import CalculationHashRow, CalculationRow, ProvenanceRow
|
|
15
|
+
from .data import EnergyRow
|
|
16
|
+
from .geometry import GeometryRow
|
|
17
|
+
from .links import (
|
|
18
|
+
CalculationGeometryLink,
|
|
19
|
+
StationaryIdentityLink,
|
|
20
|
+
StationaryStageLink,
|
|
21
|
+
)
|
|
22
|
+
from .reaction import StageRow, StepRow
|
|
23
|
+
from .stationary import IdentityRow, MetricRow, StationaryPointRow
|
|
24
|
+
|
|
25
|
+
__all__ = [
|
|
26
|
+
# links
|
|
27
|
+
"CalculationGeometryLink",
|
|
28
|
+
"StationaryIdentityLink",
|
|
29
|
+
"StationaryStageLink",
|
|
30
|
+
# geometry
|
|
31
|
+
"GeometryRow",
|
|
32
|
+
# calculation
|
|
33
|
+
"CalculationRow",
|
|
34
|
+
"ProvenanceRow",
|
|
35
|
+
"CalculationHashRow",
|
|
36
|
+
# data
|
|
37
|
+
"EnergyRow",
|
|
38
|
+
# stationary
|
|
39
|
+
"StationaryPointRow",
|
|
40
|
+
"IdentityRow",
|
|
41
|
+
"MetricRow",
|
|
42
|
+
# reaction
|
|
43
|
+
"StageRow",
|
|
44
|
+
"StepRow",
|
|
45
|
+
]
|