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.
@@ -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
+ ]