156 lines
6.3 KiB
Python
156 lines
6.3 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from typing import Literal
|
|
|
|
from pydantic import BaseModel, Field, ValidationError, model_validator
|
|
|
|
from tht.mschema.models import (
|
|
Annotations,
|
|
ColumnPhysical,
|
|
ForeignKey,
|
|
PhysicalSchema,
|
|
TablePhysical,
|
|
)
|
|
|
|
|
|
class CatalogSnapshotError(ValueError):
|
|
"""The backend-produced Catalog Metadata Snapshot is absent or invalid."""
|
|
|
|
|
|
class CatalogSnapshotColumn(BaseModel):
|
|
id: str = Field(min_length=1)
|
|
name: str = Field(min_length=1)
|
|
ordinal_position: int = Field(alias="ordinalPosition", ge=1)
|
|
data_type: str = Field(alias="dataType", min_length=1)
|
|
is_nullable: bool = Field(alias="isNullable")
|
|
default_expression: str | None = Field(alias="defaultExpression")
|
|
primary_key_position: int | None = Field(alias="primaryKeyPosition", ge=1)
|
|
sensitive: bool
|
|
description: str | None
|
|
description_source: Literal["curated", "generated", "source_comment"] | None = Field(
|
|
alias="descriptionSource"
|
|
)
|
|
|
|
model_config = {"extra": "forbid", "populate_by_name": True}
|
|
|
|
|
|
class CatalogSnapshotTable(BaseModel):
|
|
id: str = Field(min_length=1)
|
|
name: str = Field(min_length=1)
|
|
description: str | None
|
|
description_source: Literal["curated", "generated", "source_comment"] | None = Field(
|
|
alias="descriptionSource"
|
|
)
|
|
columns: list[CatalogSnapshotColumn]
|
|
|
|
model_config = {"extra": "forbid", "populate_by_name": True}
|
|
|
|
@model_validator(mode="after")
|
|
def unique_columns(self):
|
|
names = [column.name for column in self.columns]
|
|
if len(names) != len(set(names)):
|
|
raise ValueError("catalog snapshot table contains duplicate columns")
|
|
return self
|
|
|
|
|
|
class CatalogSnapshotRelationship(BaseModel):
|
|
id: str = Field(min_length=1)
|
|
origin: Literal["physical", "generated", "manual"]
|
|
source_table: str = Field(alias="sourceTable", min_length=1)
|
|
source_columns: list[str] = Field(alias="sourceColumns", min_length=1)
|
|
target_table: str = Field(alias="targetTable", min_length=1)
|
|
target_columns: list[str] = Field(alias="targetColumns", min_length=1)
|
|
|
|
model_config = {"extra": "forbid", "populate_by_name": True}
|
|
|
|
@model_validator(mode="after")
|
|
def paired_columns(self):
|
|
if len(self.source_columns) != len(self.target_columns):
|
|
raise ValueError("catalog snapshot relationship columns are not paired")
|
|
return self
|
|
|
|
|
|
class CatalogMetadataSnapshot(BaseModel):
|
|
schema_version: Literal[1] = Field(alias="schemaVersion")
|
|
workspace_id: str = Field(alias="workspaceId", min_length=1)
|
|
database_id: str = Field(alias="databaseId", min_length=1)
|
|
database_name: str = Field(alias="databaseName", min_length=1)
|
|
schema_name: str = Field(alias="schemaName", min_length=1)
|
|
metadata_content_revision: int = Field(alias="metadataContentRevision", ge=0)
|
|
tables: list[CatalogSnapshotTable]
|
|
relationships: list[CatalogSnapshotRelationship]
|
|
|
|
model_config = {"extra": "forbid", "populate_by_name": True}
|
|
|
|
@model_validator(mode="after")
|
|
def valid_graph(self):
|
|
tables = {table.name: {column.name for column in table.columns} for table in self.tables}
|
|
if len(tables) != len(self.tables):
|
|
raise ValueError("catalog snapshot contains duplicate tables")
|
|
for relationship in self.relationships:
|
|
if relationship.source_table not in tables or relationship.target_table not in tables:
|
|
raise ValueError("catalog snapshot relationship references an unknown table")
|
|
if any(name not in tables[relationship.source_table] for name in relationship.source_columns):
|
|
raise ValueError("catalog snapshot relationship references an unknown source column")
|
|
if any(name not in tables[relationship.target_table] for name in relationship.target_columns):
|
|
raise ValueError("catalog snapshot relationship references an unknown target column")
|
|
return self
|
|
|
|
def to_schema_inputs(
|
|
self,
|
|
) -> tuple[PhysicalSchema, Annotations, dict[str, list[ForeignKey]]]:
|
|
relationships: dict[str, list[ForeignKey]] = {}
|
|
for relationship in self.relationships:
|
|
relationships.setdefault(relationship.source_table, []).append(
|
|
ForeignKey(
|
|
name=relationship.id,
|
|
columns=relationship.source_columns,
|
|
ref_table=relationship.target_table,
|
|
ref_columns=relationship.target_columns,
|
|
)
|
|
)
|
|
tables = {
|
|
table.name: TablePhysical(
|
|
comment=table.description or "",
|
|
columns={
|
|
column.name: ColumnPhysical(
|
|
type=column.data_type,
|
|
nullable=column.is_nullable,
|
|
pk=column.primary_key_position is not None,
|
|
default=column.default_expression,
|
|
comment=column.description or "",
|
|
eligible=not column.sensitive,
|
|
eligibility_reason="sensitive" if column.sensitive else "catalog",
|
|
)
|
|
for column in table.columns
|
|
},
|
|
foreign_keys=relationships.get(table.name, []),
|
|
)
|
|
for table in self.tables
|
|
}
|
|
return (
|
|
PhysicalSchema(
|
|
database=self.database_name,
|
|
schema=self.schema_name,
|
|
introspected_at=datetime.now(timezone.utc),
|
|
tables=tables,
|
|
),
|
|
Annotations(),
|
|
relationships,
|
|
)
|
|
|
|
|
|
def load_catalog_metadata_snapshot(path: Path, workspace_id: str | None) -> CatalogMetadataSnapshot:
|
|
if not path.is_file():
|
|
raise CatalogSnapshotError(f"catalog metadata snapshot is missing: {path}")
|
|
try:
|
|
snapshot = CatalogMetadataSnapshot.model_validate(json.loads(path.read_text()))
|
|
except (OSError, json.JSONDecodeError, ValidationError) as exc:
|
|
raise CatalogSnapshotError("catalog metadata snapshot is invalid") from exc
|
|
if workspace_id is not None and snapshot.workspace_id != workspace_id:
|
|
raise CatalogSnapshotError("catalog metadata snapshot belongs to another workspace")
|
|
return snapshot
|