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