ch-migrate-cli 0.5.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.
ch_migrate/config.py ADDED
@@ -0,0 +1,116 @@
1
+ """Configuration loading for ch-migrate-cli."""
2
+
3
+ from pathlib import Path
4
+ from typing import Any
5
+
6
+ import yaml
7
+
8
+ from ch_migrate.secrets import get_secret
9
+
10
+
11
+ def load_config(config_path: Path) -> dict[str, Any]:
12
+ """
13
+ Load YAML configuration file.
14
+
15
+ Args:
16
+ config_path: Path to config.yaml
17
+
18
+ Returns:
19
+ Parsed configuration dictionary
20
+
21
+ Raises:
22
+ FileNotFoundError: If config file doesn't exist
23
+ """
24
+ if not config_path.exists():
25
+ raise FileNotFoundError(f"Configuration file not found: {config_path}")
26
+
27
+ with open(config_path) as f:
28
+ return yaml.safe_load(f) # type: ignore[no-any-return]
29
+
30
+
31
+ def get_env_config(env_name: str, config_path: Path) -> dict[str, Any]:
32
+ """
33
+ Get configuration for a specific environment, merging defaults.
34
+
35
+ Loads non-secret config from YAML, secrets from environment variables.
36
+ Supports both new field names (migration_user, migration_password) and
37
+ legacy names (user, password) for backward compatibility.
38
+
39
+ Args:
40
+ env_name: Environment name (dev, staging, production, etc.)
41
+ config_path: Path to config.yaml
42
+
43
+ Returns:
44
+ Complete environment configuration with secrets
45
+
46
+ Raises:
47
+ ValueError: If environment not found or required secrets missing
48
+ """
49
+ config = load_config(config_path)
50
+
51
+ environments = config.get("environments", {})
52
+ if env_name not in environments:
53
+ available = ", ".join(environments.keys()) or "(none)"
54
+ raise ValueError(f"Unknown environment: {env_name}. Available: {available}")
55
+
56
+ # Merge defaults with environment-specific config
57
+ defaults = config.get("defaults", {})
58
+ env_config = {**defaults, **environments[env_name]}
59
+
60
+ # Add project name from config if available
61
+ project_config = config.get("project", {})
62
+ if isinstance(project_config, dict) and "name" in project_config:
63
+ env_config["project"] = project_config["name"]
64
+
65
+ # Get SSM config if present (for secret lookups)
66
+ ssm_config = environments[env_name].get("ssm", {})
67
+ if ssm_config:
68
+ env_config["ssm"] = ssm_config
69
+
70
+ # Get AWS region for SSM lookups (optional, uses AWS default if not set)
71
+ aws_region = env_config.get("aws_region")
72
+
73
+ # Load secrets using unified get_secret() - uses SSM if configured, otherwise env vars
74
+ # Migration password is required
75
+ password = get_secret(
76
+ env_name,
77
+ "migration_password",
78
+ ssm_path=ssm_config.get("migration_password"),
79
+ aws_region=aws_region,
80
+ required=True,
81
+ )
82
+ env_config["password"] = password # Legacy field for backward compat
83
+
84
+ # Optional: admin password (for bootstrap)
85
+ env_config["admin_password"] = get_secret(
86
+ env_name,
87
+ "admin_password",
88
+ ssm_path=ssm_config.get("admin_password"),
89
+ aws_region=aws_region,
90
+ required=False,
91
+ )
92
+
93
+ # Optional: dict_reader password (for dictionaries)
94
+ env_config["dict_reader_password"] = get_secret(
95
+ env_name,
96
+ "dict_reader_password",
97
+ ssm_path=ssm_config.get("dict_reader_password"),
98
+ aws_region=aws_region,
99
+ required=False,
100
+ )
101
+
102
+ # Optional: MCP user password (for read-only access)
103
+ env_config["mcp_password"] = get_secret(
104
+ env_name,
105
+ "mcp_password",
106
+ ssm_path=ssm_config.get("mcp_password"),
107
+ aws_region=aws_region,
108
+ required=False,
109
+ )
110
+
111
+ # Pass through top-level hooks section (used by env.py for pre/post migrate)
112
+ hooks_config = config.get("hooks")
113
+ if hooks_config:
114
+ env_config["hooks"] = hooks_config
115
+
116
+ return env_config
@@ -0,0 +1,83 @@
1
+ """Shared ClickHouse connection helpers for CLI commands."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import io
6
+ import sys
7
+ from contextlib import contextmanager
8
+ from typing import Any
9
+
10
+
11
+ @contextmanager
12
+ def _suppress_stderr():
13
+ """Suppress stderr during clickhouse_connect operations.
14
+
15
+ clickhouse_connect prints "Unexpected Http Driver Exception" directly
16
+ to stderr on connection failures, bypassing the logging framework.
17
+ """
18
+ old_stderr = sys.stderr
19
+ sys.stderr = io.StringIO()
20
+ try:
21
+ yield
22
+ finally:
23
+ sys.stderr = old_stderr
24
+
25
+
26
+ def get_client(env_config: dict[str, Any]) -> Any:
27
+ """Create a clickhouse_connect client using migration user credentials.
28
+
29
+ Args:
30
+ env_config: Environment config dict from get_env_config().
31
+
32
+ Returns:
33
+ A clickhouse_connect Client instance.
34
+ """
35
+ import clickhouse_connect
36
+
37
+ secure = env_config.get("secure", True)
38
+ return clickhouse_connect.get_client(
39
+ host=env_config["host"],
40
+ port=env_config.get("port", 8443 if secure else 8123),
41
+ username=env_config.get("migration_user") or env_config.get("user", ""),
42
+ password=env_config.get("password", ""),
43
+ secure=secure,
44
+ interface="https" if secure else "http",
45
+ connect_timeout=10,
46
+ send_receive_timeout=15,
47
+ )
48
+
49
+
50
+ def get_current_heads(env_config: dict[str, Any]) -> set[str]:
51
+ """Query alembic_version table for current head revision(s).
52
+
53
+ Alembic stores only the current head(s) in alembic_version, not
54
+ every historically applied revision. Use resolve_applied_revisions()
55
+ to expand these into the full set of applied revisions.
56
+
57
+ Args:
58
+ env_config: Environment config dict from get_env_config().
59
+
60
+ Returns:
61
+ Set of current head revision ID strings (usually just one).
62
+ """
63
+ from clickhouse_connect.driver.binding import quote_identifier
64
+
65
+ with _suppress_stderr():
66
+ client = get_client(env_config)
67
+ try:
68
+ db = env_config["database"]
69
+ engines = client.query(
70
+ "SELECT engine FROM system.tables "
71
+ "WHERE database = {db:String} AND name = 'alembic_version'",
72
+ parameters={"db": db},
73
+ ).result_rows
74
+ if not engines:
75
+ return set()
76
+ # Legacy tables require FINAL; the official plain MergeTree rejects it.
77
+ final = " FINAL" if engines[0][0].endswith("ReplacingMergeTree") else ""
78
+ result = client.query(
79
+ f"SELECT version_num FROM {quote_identifier(db)}.alembic_version{final}"
80
+ )
81
+ return {row[0] for row in result.result_rows}
82
+ finally:
83
+ client.close()
ch_migrate/deps.py ADDED
@@ -0,0 +1,123 @@
1
+ """Dependency graph analysis and migration validation."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import re
6
+ from dataclasses import dataclass
7
+ from typing import Any
8
+
9
+ from ch_migrate.introspect import (
10
+ DependencyGraph,
11
+ DepType,
12
+ ObjectNode,
13
+ get_dependencies,
14
+ )
15
+
16
+
17
+ @dataclass
18
+ class MigrationWarning:
19
+ severity: str # "error" or "warning"
20
+ message: str
21
+ affected_objects: list[str]
22
+
23
+
24
+ # Patterns that detect destructive operations
25
+ _RE_DROP_TABLE = re.compile(
26
+ r"DROP\s+TABLE\s+(?:IF\s+EXISTS\s+)?(?:`?(\w+)`?\.)?`?(\w+)`?",
27
+ re.IGNORECASE,
28
+ )
29
+
30
+ _RE_DROP_VIEW = re.compile(
31
+ r"DROP\s+(?:MATERIALIZED\s+)?VIEW\s+(?:IF\s+EXISTS\s+)?(?:`?(\w+)`?\.)?`?(\w+)`?",
32
+ re.IGNORECASE,
33
+ )
34
+
35
+ _RE_DROP_DICT = re.compile(
36
+ r"DROP\s+DICTIONARY\s+(?:IF\s+EXISTS\s+)?(?:`?(\w+)`?\.)?`?(\w+)`?",
37
+ re.IGNORECASE,
38
+ )
39
+
40
+
41
+ def validate_migration(sql: str, graph: DependencyGraph) -> list[MigrationWarning]:
42
+ """Check if migration SQL would break dependencies in the graph.
43
+
44
+ Args:
45
+ sql: The migration SQL to validate.
46
+ graph: A live DependencyGraph from introspect.get_dependencies().
47
+
48
+ Returns:
49
+ List of warnings about potential dependency breakage.
50
+ """
51
+ warnings: list[MigrationWarning] = []
52
+
53
+ # Check DROP TABLE
54
+ for m in _RE_DROP_TABLE.finditer(sql):
55
+ table_name = m.group(2)
56
+ if table_name in graph.nodes:
57
+ affected = graph.affected_by_drop(table_name)
58
+ if affected:
59
+ # Distinguish schema vs data_flow impact
60
+ schema_deps = []
61
+ data_flow_deps = []
62
+ for edge in graph.edges:
63
+ if edge.source == table_name:
64
+ if edge.dep_type == DepType.SCHEMA:
65
+ schema_deps.append(edge.target)
66
+ elif edge.dep_type == DepType.DATA_FLOW:
67
+ data_flow_deps.append(edge.target)
68
+
69
+ if schema_deps:
70
+ warnings.append(MigrationWarning(
71
+ severity="error",
72
+ message=f"DROP TABLE {table_name} breaks schema dependencies: {', '.join(schema_deps)}",
73
+ affected_objects=schema_deps,
74
+ ))
75
+ if data_flow_deps:
76
+ warnings.append(MigrationWarning(
77
+ severity="warning",
78
+ message=f"DROP TABLE {table_name} breaks data flow to: {', '.join(data_flow_deps)}",
79
+ affected_objects=data_flow_deps,
80
+ ))
81
+
82
+ # Check DROP VIEW / DROP MATERIALIZED VIEW
83
+ for m in _RE_DROP_VIEW.finditer(sql):
84
+ view_name = m.group(2)
85
+ if view_name in graph.nodes:
86
+ affected = graph.affected_by_drop(view_name)
87
+ if affected:
88
+ affected_names = [n.name for n in affected]
89
+ warnings.append(MigrationWarning(
90
+ severity="warning",
91
+ message=f"DROP VIEW {view_name} affects: {', '.join(affected_names)}",
92
+ affected_objects=affected_names,
93
+ ))
94
+
95
+ # Check DROP DICTIONARY
96
+ for m in _RE_DROP_DICT.finditer(sql):
97
+ dict_name = m.group(2)
98
+ if dict_name in graph.nodes:
99
+ affected = graph.affected_by_drop(dict_name)
100
+ if affected:
101
+ affected_names = [n.name for n in affected]
102
+ warnings.append(MigrationWarning(
103
+ severity="warning",
104
+ message=f"DROP DICTIONARY {dict_name} affects: {', '.join(affected_names)}",
105
+ affected_objects=affected_names,
106
+ ))
107
+
108
+ return warnings
109
+
110
+
111
+ def build_dependency_graph(client: Any, database: str) -> DependencyGraph:
112
+ """Build a dependency graph from the live database.
113
+
114
+ Convenience wrapper around introspect.get_dependencies().
115
+
116
+ Args:
117
+ client: clickhouse-connect client.
118
+ database: Database name.
119
+
120
+ Returns:
121
+ A DependencyGraph with nodes and typed edges.
122
+ """
123
+ return get_dependencies(client, database)
ch_migrate/diff.py ADDED
@@ -0,0 +1,232 @@
1
+ """Schema comparison: field-by-field structured diff between two Schema objects."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass, field
6
+ from enum import Enum
7
+ from typing import Literal
8
+
9
+ from ch_migrate.introspect import (
10
+ ColumnDefinition,
11
+ Schema,
12
+ TableDefinition,
13
+ )
14
+
15
+
16
+ class DiffStatus(str, Enum):
17
+ IN_SYNC = "in_sync"
18
+ MODIFIED = "modified"
19
+ LOCAL_ONLY = "local_only"
20
+ REMOTE_ONLY = "remote_only"
21
+
22
+
23
+ @dataclass
24
+ class FieldDiff:
25
+ field_name: str
26
+ local_value: str | None
27
+ remote_value: str | None
28
+ message: str
29
+
30
+
31
+ @dataclass
32
+ class SchemaDiff:
33
+ name: str
34
+ obj_type: str # "table", "view", "materialized_view", "dictionary"
35
+ status: DiffStatus
36
+ field_diffs: list[FieldDiff] = field(default_factory=list)
37
+
38
+
39
+ def _compare_columns(
40
+ local_cols: list[ColumnDefinition],
41
+ remote_cols: list[ColumnDefinition],
42
+ ) -> list[FieldDiff]:
43
+ """Compare column lists field-by-field."""
44
+ diffs: list[FieldDiff] = []
45
+ local_map = {c.name: c for c in local_cols}
46
+ remote_map = {c.name: c for c in remote_cols}
47
+
48
+ all_names = dict.fromkeys([c.name for c in local_cols] + [c.name for c in remote_cols])
49
+
50
+ for name in all_names:
51
+ local_col = local_map.get(name)
52
+ remote_col = remote_map.get(name)
53
+
54
+ if local_col and not remote_col:
55
+ diffs.append(FieldDiff(
56
+ field_name=f"column '{name}'",
57
+ local_value=local_col.type,
58
+ remote_value=None,
59
+ message=f"column '{name}' {local_col.type} is in the snapshot but not in the database",
60
+ ))
61
+ elif remote_col and not local_col:
62
+ diffs.append(FieldDiff(
63
+ field_name=f"column '{name}'",
64
+ local_value=None,
65
+ remote_value=remote_col.type,
66
+ message=f"column '{name}' {remote_col.type} is in the database but not in the snapshot",
67
+ ))
68
+ elif local_col and remote_col:
69
+ if local_col.type != remote_col.type:
70
+ diffs.append(FieldDiff(
71
+ field_name=f"column '{name}' type",
72
+ local_value=local_col.type,
73
+ remote_value=remote_col.type,
74
+ message=f"column '{name}' type differs: {local_col.type} vs {remote_col.type}",
75
+ ))
76
+ if local_col.default_kind != remote_col.default_kind:
77
+ diffs.append(FieldDiff(
78
+ field_name=f"column '{name}' default_kind",
79
+ local_value=local_col.default_kind,
80
+ remote_value=remote_col.default_kind,
81
+ message=f"column '{name}' default kind differs",
82
+ ))
83
+ if local_col.default_expr != remote_col.default_expr:
84
+ diffs.append(FieldDiff(
85
+ field_name=f"column '{name}' default_expr",
86
+ local_value=local_col.default_expr,
87
+ remote_value=remote_col.default_expr,
88
+ message=f"column '{name}' default expression differs",
89
+ ))
90
+ if local_col.codec != remote_col.codec:
91
+ diffs.append(FieldDiff(
92
+ field_name=f"column '{name}' codec",
93
+ local_value=local_col.codec,
94
+ remote_value=remote_col.codec,
95
+ message=f"column '{name}' codec differs: {local_col.codec} vs {remote_col.codec}",
96
+ ))
97
+
98
+ return diffs
99
+
100
+
101
+ def _compare_tables(local: TableDefinition, remote: TableDefinition) -> list[FieldDiff]:
102
+ """Field-by-field comparison of two TableDefinitions."""
103
+ diffs: list[FieldDiff] = []
104
+
105
+ # Engine
106
+ if local.engine != remote.engine:
107
+ diffs.append(FieldDiff(
108
+ field_name="engine",
109
+ local_value=local.engine,
110
+ remote_value=remote.engine,
111
+ message=f"engine differs: {local.engine} vs {remote.engine}",
112
+ ))
113
+
114
+ # Columns
115
+ diffs.extend(_compare_columns(local.columns, remote.columns))
116
+
117
+ # ORDER BY
118
+ if local.order_by != remote.order_by:
119
+ diffs.append(FieldDiff(
120
+ field_name="order_by",
121
+ local_value=", ".join(local.order_by),
122
+ remote_value=", ".join(remote.order_by),
123
+ message=f"ORDER BY differs",
124
+ ))
125
+
126
+ # PARTITION BY
127
+ if local.partition_by != remote.partition_by:
128
+ diffs.append(FieldDiff(
129
+ field_name="partition_by",
130
+ local_value=local.partition_by,
131
+ remote_value=remote.partition_by,
132
+ message=f"PARTITION BY differs",
133
+ ))
134
+
135
+ # TTL
136
+ if local.ttl != remote.ttl:
137
+ diffs.append(FieldDiff(
138
+ field_name="ttl",
139
+ local_value=local.ttl,
140
+ remote_value=remote.ttl,
141
+ message=f"TTL differs",
142
+ ))
143
+
144
+ # Settings
145
+ if local.settings != remote.settings:
146
+ diffs.append(FieldDiff(
147
+ field_name="settings",
148
+ local_value=str(local.settings),
149
+ remote_value=str(remote.settings),
150
+ message=f"SETTINGS differ",
151
+ ))
152
+
153
+ return diffs
154
+
155
+
156
+ def _normalize_ddl(raw: str) -> str:
157
+ """Normalize raw DDL for fallback string comparison."""
158
+ import re
159
+ s = raw.strip()
160
+ s = re.sub(r"\s+", " ", s)
161
+ return s
162
+
163
+
164
+ def _compare_raw_ddl(local_ddl: str, remote_ddl: str) -> list[FieldDiff]:
165
+ """Fallback: normalized string comparison for unparseable objects."""
166
+ if _normalize_ddl(local_ddl) != _normalize_ddl(remote_ddl):
167
+ return [FieldDiff(
168
+ field_name="raw_ddl",
169
+ local_value=local_ddl[:200] if local_ddl else None,
170
+ remote_value=remote_ddl[:200] if remote_ddl else None,
171
+ message="DDL definition differs (raw comparison)",
172
+ )]
173
+ return []
174
+
175
+
176
+ def compare_schemas(local: Schema, live: Schema) -> list[SchemaDiff]:
177
+ """Compare two Schema objects and return a list of differences.
178
+
179
+ Args:
180
+ local: Schema from local snapshot files.
181
+ live: Schema from the live database.
182
+
183
+ Returns:
184
+ List of SchemaDiff objects. Empty list means schemas are in sync.
185
+ """
186
+ results: list[SchemaDiff] = []
187
+
188
+ type_map: list[tuple[str, dict, dict]] = [
189
+ ("table", local.tables, live.tables),
190
+ ("view", local.views, live.views),
191
+ ("materialized_view", local.materialized_views, live.materialized_views),
192
+ ("dictionary", local.dictionaries, live.dictionaries),
193
+ ]
194
+
195
+ for obj_type, local_objs, live_objs in type_map:
196
+ all_names = dict.fromkeys(list(local_objs.keys()) + list(live_objs.keys()))
197
+
198
+ for name in all_names:
199
+ local_obj = local_objs.get(name)
200
+ live_obj = live_objs.get(name)
201
+
202
+ if local_obj and not live_obj:
203
+ results.append(SchemaDiff(
204
+ name=name, obj_type=obj_type, status=DiffStatus.LOCAL_ONLY,
205
+ ))
206
+ elif live_obj and not local_obj:
207
+ results.append(SchemaDiff(
208
+ name=name, obj_type=obj_type, status=DiffStatus.REMOTE_ONLY,
209
+ ))
210
+ else:
211
+ # Both exist — compare
212
+ field_diffs: list[FieldDiff] = []
213
+
214
+ if obj_type == "table" and isinstance(local_obj, TableDefinition) and isinstance(live_obj, TableDefinition):
215
+ field_diffs = _compare_tables(local_obj, live_obj)
216
+ else:
217
+ # Fallback to raw DDL comparison for views, MVs, dicts
218
+ local_ddl = getattr(local_obj, "raw_ddl", "")
219
+ live_ddl = getattr(live_obj, "raw_ddl", "")
220
+ field_diffs = _compare_raw_ddl(local_ddl, live_ddl)
221
+
222
+ if field_diffs:
223
+ results.append(SchemaDiff(
224
+ name=name, obj_type=obj_type,
225
+ status=DiffStatus.MODIFIED, field_diffs=field_diffs,
226
+ ))
227
+ else:
228
+ results.append(SchemaDiff(
229
+ name=name, obj_type=obj_type, status=DiffStatus.IN_SYNC,
230
+ ))
231
+
232
+ return results