diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 153677e9a..a5034016f 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -37,11 +37,11 @@ jobs: PYTHONPATH: . run: mypy app - - name: Tests (pytest) + - name: Tests and coverage (pytest) working-directory: backend env: PYTHONPATH: . - run: pytest -q + run: pytest -q --cov --cov-report=term-missing --cov-fail-under=100 frontend: runs-on: ubuntu-latest diff --git a/CHANGELOG.md b/CHANGELOG.md index 35613431a..c2a0bdb6f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,8 +1,13 @@ # Changelog ## Unreleased + +- [BE] ๐Ÿ” **์‹ ๋ขฐ๋œ ๋กœ์ปฌ PostgreSQL ์Šค๋ƒ…์ƒท CLI**: ์›น API์˜ SSRF ํ—ˆ์šฉ ์ •์ฑ…์„ ์™„ํ™”ํ•˜๊ฑฐ๋‚˜ ๋น„๋ฐ€๋ฒˆํ˜ธ๊ฐ€ ํฌํ•จ๋œ DSN์„ ๋…ธ์ถœํ•˜์ง€ ์•Š๊ณ , ๋ช…์‹œ์ ์œผ๋กœ ๊ฒ€์ฆ๋œ Unix-domain socket์„ ํ†ตํ•ด ๋กœ์ปฌ ๋ฐ์ดํ„ฐ๋ฒ ์ด์Šค ์Šคํ‚ค๋งˆ๋ฅผ JSON์œผ๋กœ ์บก์ฒ˜ํ•  ์ˆ˜ ์žˆ๋Š” `pg-erd-snapshot` ๋ช…๋ น์„ ์ถ”๊ฐ€ํ–ˆ์Šต๋‹ˆ๋‹ค. ๊ณต์šฉ ์ˆ˜์ง‘๊ธฐ๋ฅผ ์›นยทCLI ๊ฒฝ๋กœ์—์„œ ์žฌ์‚ฌ์šฉํ•˜๊ณ  Citus ์œ ๋ฌดยท์นดํƒˆ๋กœ๊ทธ ํ˜ธํ™˜์„ฑยท์ถœ๋ ฅ ๋ฐ ์˜ค๋ฅ˜ redaction์„ ํšŒ๊ท€ ํ…Œ์ŠคํŠธ๋กœ ๊ฒ€์ฆํ•˜๋ฉฐ, ๋ฐฑ์—”๋“œ CI์—์„œ ์„ ํƒ๋œ ํ•ต์‹ฌ ๋ชจ๋“ˆ์˜ 100% coverage evidence๋ฅผ ๊ฐ•์ œํ•ฉ๋‹ˆ๋‹ค. +- [BE] ๐Ÿงญ **๋กœ์ปฌ ์Šค๋ƒ…์ƒท ์—ฐ๊ฒฐ ์‹œ๊ฐ„ ์ดˆ๊ณผ ์ฒ˜๋ฆฌ**: Unix socket ์—ฐ๊ฒฐ์ด ์‹œ๊ฐ„ ์ดˆ๊ณผ๋˜๋ฉด ๋‚ด๋ถ€ ๋“œ๋ผ์ด๋ฒ„ ๋ฉ”์‹œ์ง€๋ฅผ ๋…ธ์ถœํ•˜์ง€ ์•Š๊ณ  ์•ˆ์ •์ ์ธ ์‹คํŒจ ์ฝ”๋“œ์™€ ์˜ค๋ฅ˜ ์œ ํ˜•๋งŒ ๋ฐ˜ํ™˜ํ•ฉ๋‹ˆ๋‹ค. ์—ฐ๊ฒฐ์ด ์ค€๋น„๋˜๋ฉด ๊ฐ™์€ ๋ช…๋ น์„ ๋‹ค์‹œ ์‹คํ–‰ํ•˜์‹ญ์‹œ์˜ค. + - [BE] ๐Ÿ”’ **Cryptography 50+ ๋ณด์•ˆ ๊ฒฝ๊ณ„ ๊ฐฑ์‹ **: `pyproject.toml`๊ณผ ๋‘ hash-locked ์š”๊ตฌ์‚ฌํ•ญ ํŒŒ์ผ์„ ๋™์ผํ•œ Cryptography 50+ ํ•ด์„์œผ๋กœ ์ •ํ•ฉํ™”ํ•˜์—ฌ PKCS#7 ์˜ค๋ฅ˜ยทํƒ€์ด๋ฐ ๊ตฌ๋ถ„์œผ๋กœ ์ธํ•œ CVE-2026-69247 ์™„ํ™”๋ฅผ ์‹ค์ œ ์„ค์น˜ยท๊ฒ€์ฆ ๊ฒฝ๋กœ์— ๋ฐ˜์˜ํ–ˆ์Šต๋‹ˆ๋‹ค. - [FE] โšก **๊ฒ€์ƒ‰ ๋…ธ๋“œ ์ฐธ์กฐ ์•ˆ์ •ํ™” ๋ฐ ์ˆœ์ฐจ ์Šค๋ƒ…์ƒท ํด๋ง**: ๊ฐ™์€ ์ •๊ทœํ™” ๊ฒ€์ƒ‰์–ด์™€ ์›๋ณธ ํ…Œ์ด๋ธ” ๋ฐ์ดํ„ฐ์—๋Š” ์žฅ์‹๋œ `node.data` ์ฐธ์กฐ๋ฅผ ์žฌ์‚ฌ์šฉํ•˜์—ฌ ๋“œ๋ž˜๊ทธ ์ค‘ ๋ถˆํ•„์š”ํ•œ ํ•˜์œ„ ๋ Œ๋”๋ง๊ณผ ํ• ๋‹น์„ ์ค„์ž…๋‹ˆ๋‹ค. ์Šค๋ƒ…์ƒท ํด๋ง์€ ์ด์ „ ์š”์ฒญ์ด ๋๋‚œ ๋’ค์—๋งŒ ๋‹ค์Œ ์š”์ฒญ์„ ์˜ˆ์•ฝํ•˜๋ฉฐ, ์„ ํƒ ๋ณ€๊ฒฝยท์–ธ๋งˆ์šดํŠธ ํ›„ ๋„์ฐฉํ•œ ์˜ค๋ž˜๋œ ์„ฑ๊ณต ๋˜๋Š” ์‹คํŒจ ์‘๋‹ต์„ ๋ฌด์‹œํ•ฉ๋‹ˆ๋‹ค. + - [BE] ๐Ÿ”’ **๊ณต์œ  export ์ „ ๊ฒฝ๋กœ redaction**: ๊ณต๊ฐœ share์˜ SQL / index-design / reversing-spec export์—์„œ ์ฝ”๋ฉ˜ํŠธยท`example_value`๋ฅผ ์ œ๊ฑฐํ•ฉ๋‹ˆ๋‹ค. ๋‹จ์œ„ ํ…Œ์ŠคํŠธ๋กœ ๋ˆ„์ถœ์„ ์ฐจ๋‹จํ•ฉ๋‹ˆ๋‹ค. - [BE] ๐Ÿ› ๏ธ **ํ•จ์ˆ˜ ์ธ๋ฑ์Šค ์ค‘๋ณต ์˜คํƒ ์ˆ˜์ •**: `lower(email)` ๋“ฑ expression index๋ฅผ ํ‰๋ฌธ ์ปฌ๋Ÿผ ์ธ๋ฑ์Šค์˜ ์ค‘๋ณต์œผ๋กœ ์ž˜๋ชป ํŒ๋‹จํ•˜์ง€ ์•Š๋„๋ก ๊ด„ํ˜ธ ํŒŒ์„œ๋ฅผ ๊ฐ•ํ™”ํ–ˆ์Šต๋‹ˆ๋‹ค. - [Docs] README๋ฅผ ์ƒ์šฉ ๊ธฐ์ค€ ๊ธฐ๋Šฅ ์„ค๋ช…์œผ๋กœ ๊ฐฑ์‹  (MVP skeleton ํ‘œํ˜„ ์ œ๊ฑฐ, share redactionยทdiff/export ๋ฐ˜์˜). diff --git a/README.md b/README.md index 29924af51..f66646599 100644 --- a/README.md +++ b/README.md @@ -135,6 +135,22 @@ hypercorn --config python:app.hypercorn_config app.main:app \ --access-logfile - --error-logfile - ``` +๋กœ์ปฌ ๋งˆ์ด๊ทธ๋ ˆ์ด์…˜ DB๋ฅผ ์›น API์˜ SSRF ํ—ˆ์šฉ ๋ชฉ๋ก์— ๋„ฃ์ง€ ์•Š๊ณ  ์Šค๋ƒ…์ƒทํ•˜๋ ค๋ฉด +Unix-domain socket ์ „์šฉ CLI๋ฅผ ์‚ฌ์šฉํ•ฉ๋‹ˆ๋‹ค. ๋น„๋ฐ€๋ฒˆํ˜ธ๊ฐ€ ํฌํ•จ๋œ DSN์„ ์ธ์ž๋‚˜ ๋กœ๊ทธ๋กœ +์ „๋‹ฌํ•˜์ง€ ์•Š์œผ๋ฉฐ ๊ฒฐ๊ณผ JSON์€ stdout์œผ๋กœ๋งŒ ์ถœ๋ ฅํ•ฉ๋‹ˆ๋‹ค. + +```bash +cd backend +pg-erd-snapshot \ + --host /tmp \ + --database my_local_database \ + --schema public \ + --pretty > snapshot.json +``` + +์›๊ฒฉ PostgreSQL์€ ๊ธฐ์กด ์—ฐ๊ฒฐ API์™€ host allowlist๋ฅผ ๊ณ„์† ์‚ฌ์šฉํ•ฉ๋‹ˆ๋‹ค. ์ด CLI๋Š” TCP +ํ˜ธ์ŠคํŠธ๋ฅผ ๋ฐ›์ง€ ์•Š์œผ๋ฏ€๋กœ ์›น API์˜ loopback/private-network ์ฐจ๋‹จ์„ ์šฐํšŒํ•˜์ง€ ์•Š์Šต๋‹ˆ๋‹ค. + Snowflake ๋ฆฌ๋ฒ„์Šค ์—”์ง€๋‹ˆ์–ด๋ง์„ ์‚ฌ์šฉํ•  ๊ฐœ๋ฐœ ํ™˜๊ฒฝ์—์„œ๋Š” ๋ฐฑ์—”๋“œ ๊ฐ€์ƒํ™˜๊ฒฝ์—์„œ ์„ ํƒ ์˜์กด์„ฑ์„ ํ•จ๊ป˜ ์„ค์น˜ํ•ฉ๋‹ˆ๋‹ค. diff --git a/backend/app/local_snapshot_cli.py b/backend/app/local_snapshot_cli.py new file mode 100644 index 000000000..5ce9da554 --- /dev/null +++ b/backend/app/local_snapshot_cli.py @@ -0,0 +1,147 @@ +"""Trusted-local PostgreSQL snapshot CLI. + +The web API intentionally rejects loopback and private database targets. This +separate operator CLI accepts only an absolute Unix-domain socket directory, so +developers can reverse a local migration database without weakening the remote +API's SSRF boundary or putting a password-bearing DSN in the process list. +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import os +import re +import sys +from pathlib import Path +from typing import Sequence + +import asyncpg + +from app.pg_introspect.snapshot_collect import collect_postgres_snapshot + +_SCHEMA_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_$]{0,62}$") + + +def _socket_directory(value: str) -> str: + """Accept an existing absolute directory suitable for a Unix socket.""" + + path = Path(value) + if not path.is_absolute(): + raise argparse.ArgumentTypeError( + "--host must be an absolute PostgreSQL Unix socket directory" + ) + if not path.is_dir(): + raise argparse.ArgumentTypeError("--host Unix socket directory does not exist") + return str(path) + + +def _schema_name(value: str) -> str: + """Accept one unquoted PostgreSQL schema identifier.""" + + if not _SCHEMA_RE.fullmatch(value): + raise argparse.ArgumentTypeError("--schema is not a valid PostgreSQL identifier") + return value + + +def _port(value: str) -> int: + """Parse a PostgreSQL port in the valid TCP and Unix-socket range.""" + + try: + port = int(value) + except ValueError as exc: + raise argparse.ArgumentTypeError("--port must be an integer") from exc + if not 1 <= port <= 65535: + raise argparse.ArgumentTypeError("--port must be between 1 and 65535") + return port + + +def build_parser() -> argparse.ArgumentParser: + """Build the argument parser for the trusted-local snapshot CLI.""" + + parser = argparse.ArgumentParser( + prog="pg-erd-snapshot", + description=( + "Write a pg-erd-cloud snapshot for a trusted local PostgreSQL Unix socket " + "to stdout." + ), + ) + parser.add_argument( + "--database", + default=os.environ.get("PGDATABASE"), + required=os.environ.get("PGDATABASE") is None, + help="database name (defaults to PGDATABASE)", + ) + host_default = os.environ.get("PGHOST") + parser.add_argument( + "--host", + type=_socket_directory, + default=host_default, + required=host_default is None, + help="absolute Unix socket directory (defaults to PGHOST; no implicit fallback)", + ) + parser.add_argument( + "--port", + type=_port, + default=os.environ.get("PGPORT", "5432"), + metavar="1..65535", + ) + parser.add_argument( + "--user", + default=os.environ.get("PGUSER"), + help="database user (defaults to PGUSER or the operating-system user)", + ) + parser.add_argument("--schema", type=_schema_name, default=None) + parser.add_argument("--pretty", action="store_true") + return parser + + +async def capture_local_snapshot(args: argparse.Namespace) -> dict: + """Connect over a validated Unix socket and collect one catalog snapshot.""" + + conn = await asyncpg.connect( + database=args.database, + host=args.host, + password="", + port=args.port, + user=args.user, + timeout=10, + ) + try: + return await collect_postgres_snapshot(conn, args.schema) + finally: + await conn.close() + + +def main(argv: Sequence[str] | None = None) -> int: + """Capture a trusted-local snapshot and write JSON to standard output.""" + + args = build_parser().parse_args(argv) + try: + snapshot = asyncio.run(capture_local_snapshot(args)) + except ( + OSError, + TimeoutError, + asyncio.TimeoutError, + asyncpg.PostgresError, + ) as exc: + print( + f"pg-erd-snapshot failed: {type(exc).__name__}", + file=sys.stderr, + ) + return 1 + json.dump( + snapshot, + sys.stdout, + ensure_ascii=False, + indent=2 if args.pretty else None, + sort_keys=True, + separators=None if args.pretty else (",", ":"), + ) + sys.stdout.write("\n") + return 0 + + +if __name__ == "__main__": # pragma: no cover + raise SystemExit(main()) diff --git a/backend/app/pg_introspect/introspect.py b/backend/app/pg_introspect/introspect.py index 598fb29c6..bc405aeb8 100644 --- a/backend/app/pg_introspect/introspect.py +++ b/backend/app/pg_introspect/introspect.py @@ -1,16 +1,15 @@ +"""Guarded PostgreSQL connectivity and snapshot introspection helpers.""" + from __future__ import annotations -import datetime as dt import ssl from urllib.parse import parse_qsl, urlparse import asyncpg -from app.pg_introspect import queries -from app.pg_introspect.column_examples import add_column_examples from app.pg_introspect.dsn_guard import validate_postgres_dsn_target from app.pg_introspect.forward_ddl import ForwardDdlBatch -from app.sanitize import sanitize_for_storage +from app.pg_introspect.snapshot_collect import collect_postgres_snapshot class _ServerHostnameSSLContext(ssl.SSLContext): @@ -19,11 +18,13 @@ class _ServerHostnameSSLContext(ssl.SSLContext): _server_hostname: str def __new__(cls, server_hostname: str) -> "_ServerHostnameSSLContext": + """Create a client TLS context that retains the verified DSN hostname.""" context = super().__new__(cls, ssl.PROTOCOL_TLS_CLIENT) context._server_hostname = server_hostname return context def __init__(self, server_hostname: str) -> None: + """Keep the SSL context initialized by ``__new__`` without resetting it.""" return None def wrap_bio( @@ -34,6 +35,7 @@ def wrap_bio( server_hostname: str | bytes | None = None, session: ssl.SSLSession | None = None, ) -> ssl.SSLObject: + """Wrap a TLS BIO while forcing certificate verification to the DSN host.""" return super().wrap_bio( incoming, outgoing, @@ -44,11 +46,13 @@ def wrap_bio( def _requires_verified_tls_hostname(dsn: str) -> bool: + """Return whether the DSN requests PostgreSQL ``verify-full`` TLS mode.""" query = dict(parse_qsl(urlparse(dsn).query, keep_blank_values=True)) return query.get("sslmode", "").lower() == "verify-full" def _verified_tls_context(dsn: str, server_hostname: str) -> ssl.SSLContext: + """Build a TLS context from DSN certificate options with hostname binding.""" query = dict(parse_qsl(urlparse(dsn).query, keep_blank_values=True)) context = _ServerHostnameSSLContext(server_hostname) if query.get("sslrootcert"): @@ -63,6 +67,7 @@ def _verified_tls_context(dsn: str, server_hostname: str) -> ssl.SSLContext: async def _connect_guarded_postgres( dsn: str, *, timeout: float ) -> asyncpg.Connection: + """Validate a DSN and connect only to its resolved, permitted host targets.""" target = await validate_postgres_dsn_target(dsn) connect_host: str | list[str] = ( target.hosts[0] if len(target.hosts) == 1 else list(target.hosts) @@ -136,49 +141,6 @@ async def introspect_postgres(dsn: str, schema_filter: str | None) -> dict: # Note: avoid logging DSN. conn = await _connect_guarded_postgres(dsn, timeout=10) try: - version = await conn.fetchval("SHOW server_version") - schema_name = schema_filter - include_system = False - - schemas = await conn.fetch(queries.SCHEMAS_SQL, schema_name, include_system) - relations = await conn.fetch(queries.RELATIONS_SQL, schema_name, include_system) - columns = await conn.fetch(queries.COLUMNS_SQL, schema_name, include_system) - constraints = await conn.fetch( - queries.CONSTRAINTS_SQL, schema_name, include_system - ) - indexes = await conn.fetch(queries.INDEXES_SQL, schema_name, include_system) - pk_columns = await conn.fetch( - queries.PK_COLUMNS_SQL, schema_name, include_system - ) - fk_edges = await conn.fetch(queries.FK_EDGES_SQL, schema_name, include_system) - citus_distributed_tables = [] - has_citus = await conn.fetchval( - "SELECT EXISTS (SELECT 1 FROM pg_catalog.pg_extension WHERE extname = 'citus')" - ) - if has_citus: - try: - citus_distributed_tables = await conn.fetch( - queries.CITUS_DISTRIBUTED_TABLES_SQL, - schema_name, - include_system, - ) - except asyncpg.UndefinedTableError: - citus_distributed_tables = [] - - snapshot = { - "captured_at": dt.datetime.now(dt.timezone.utc).isoformat(), - "server_version": str(version), - "schema_filter": schema_filter, - "schemas": [dict(r) for r in schemas], - "relations": [dict(r) for r in relations], - "columns": add_column_examples([dict(r) for r in columns]), - "constraints": [dict(r) for r in constraints], - "indexes": [dict(r) for r in indexes], - "pk_columns": [dict(r) for r in pk_columns], - "fk_edges": [dict(r) for r in fk_edges], - "citus_distributed_tables": [dict(r) for r in citus_distributed_tables], - } - - return sanitize_for_storage(snapshot) # type: ignore[return-value] + return await collect_postgres_snapshot(conn, schema_filter) finally: await conn.close() diff --git a/backend/app/pg_introspect/snapshot_collect.py b/backend/app/pg_introspect/snapshot_collect.py new file mode 100644 index 000000000..be7dc2b82 --- /dev/null +++ b/backend/app/pg_introspect/snapshot_collect.py @@ -0,0 +1,71 @@ +"""Shared PostgreSQL catalog snapshot collector. + +This module deliberately has no application-settings import. Network trust and +credential policy are established before callers hand it an open connection. +""" + +from __future__ import annotations + +import datetime as dt + +import asyncpg + +from app.pg_introspect import queries +from app.pg_introspect.column_examples import add_column_examples +from app.sanitize import sanitize_for_storage + + +async def collect_postgres_snapshot( + conn: asyncpg.Connection, schema_filter: str | None +) -> dict: + """Collect the canonical snapshot from an authorized PostgreSQL connection.""" + + version = await conn.fetchval("SHOW server_version") + schema_name = schema_filter + include_system = False + + # asyncpg rejects overlapping operations on one Connection. Keep these + # catalog reads sequential unless a future caller gives the collector a + # pool and an explicit multi-connection consistency policy. + schemas = await conn.fetch(queries.SCHEMAS_SQL, schema_name, include_system) + relations = await conn.fetch(queries.RELATIONS_SQL, schema_name, include_system) + columns = await conn.fetch(queries.COLUMNS_SQL, schema_name, include_system) + constraints = await conn.fetch( + queries.CONSTRAINTS_SQL, schema_name, include_system + ) + indexes = await conn.fetch(queries.INDEXES_SQL, schema_name, include_system) + pk_columns = await conn.fetch( + queries.PK_COLUMNS_SQL, schema_name, include_system + ) + fk_edges = await conn.fetch(queries.FK_EDGES_SQL, schema_name, include_system) + citus_distributed_tables = [] + has_citus = await conn.fetchval( + "SELECT EXISTS (SELECT 1 FROM pg_catalog.pg_extension WHERE extname = 'citus')" + ) + if has_citus: + try: + citus_distributed_tables = await conn.fetch( + queries.CITUS_DISTRIBUTED_TABLES_SQL, + schema_name, + include_system, + ) + except asyncpg.UndefinedTableError: + citus_distributed_tables = [] + + snapshot = { + "captured_at": dt.datetime.now(dt.timezone.utc).isoformat(), + "server_version": str(version), + "schema_filter": schema_filter, + "schemas": [dict(row) for row in schemas], + "relations": [dict(row) for row in relations], + "columns": add_column_examples([dict(row) for row in columns]), + "constraints": [dict(row) for row in constraints], + "indexes": [dict(row) for row in indexes], + "pk_columns": [dict(row) for row in pk_columns], + "fk_edges": [dict(row) for row in fk_edges], + "citus_distributed_tables": [ + dict(row) for row in citus_distributed_tables + ], + } + + return sanitize_for_storage(snapshot) # type: ignore[return-value] diff --git a/backend/pyproject.toml b/backend/pyproject.toml index b2d47dd2a..ae599403e 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -32,6 +32,9 @@ dependencies = [ "python-multipart>=0.0.32", ] +[project.scripts] +pg-erd-snapshot = "app.local_snapshot_cli:main" + [project.optional-dependencies] snowflake = [ "snowflake-connector-python>=4,<5", @@ -65,9 +68,11 @@ asyncio_mode = "auto" include = [ "app/csrf.py", "app/db_introspect.py", + "app/local_snapshot_cli.py", "app/metrics.py", "app/permissions.py", "app/pg_introspect/queries.py", + "app/pg_introspect/snapshot_collect.py", "app/pooler.py", "app/schemas.py", "app/security.py", diff --git a/backend/tests/test_docstrings.py b/backend/tests/test_docstrings.py index 4856ca134..9464e08e3 100644 --- a/backend/tests/test_docstrings.py +++ b/backend/tests/test_docstrings.py @@ -7,7 +7,11 @@ BACKEND_ROOT = Path(__file__).resolve().parents[1] -CHECKED_MODULES = (BACKEND_ROOT / "app" / "snowflake_introspect" / "introspect.py",) +CHECKED_MODULES = ( + BACKEND_ROOT / "app" / "local_snapshot_cli.py", + BACKEND_ROOT / "app" / "pg_introspect" / "snapshot_collect.py", + BACKEND_ROOT / "app" / "snowflake_introspect" / "introspect.py", +) def _public_defs( diff --git a/backend/tests/test_local_snapshot_cli.py b/backend/tests/test_local_snapshot_cli.py new file mode 100644 index 000000000..bd61e6c2b --- /dev/null +++ b/backend/tests/test_local_snapshot_cli.py @@ -0,0 +1,203 @@ +from __future__ import annotations + +import argparse +import asyncio +import json +from pathlib import Path +from typing import Any + +import pytest + +from app import local_snapshot_cli + + +class FakeConnection: + """Track closure for a deterministic local snapshot connection.""" + + def __init__(self) -> None: + self.closed = False + + async def close(self) -> None: + """Record that the snapshot collector closed the connection.""" + + self.closed = True + + +def test_socket_directory_requires_existing_absolute_directory(tmp_path: Path) -> None: + """Reject relative or missing socket directories.""" + + assert local_snapshot_cli._socket_directory(str(tmp_path)) == str(tmp_path) + with pytest.raises(argparse.ArgumentTypeError): + local_snapshot_cli._socket_directory("relative/socket") + with pytest.raises(argparse.ArgumentTypeError): + local_snapshot_cli._socket_directory(str(tmp_path / "missing")) + + +def test_environment_host_default_cannot_bypass_socket_boundary( + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + """Validate an explicit PGHOST default with the Unix-socket boundary.""" + + monkeypatch.setenv("PGDATABASE", "catalog") + monkeypatch.setenv("PGHOST", "db.example.com") + + with pytest.raises(SystemExit) as exc_info: + local_snapshot_cli.build_parser().parse_args([]) + + assert exc_info.value.code == 2 + assert "absolute PostgreSQL Unix socket directory" in capsys.readouterr().err + + +def test_parser_requires_explicit_host_without_pghost( + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + """Require --host when PGHOST does not provide a socket directory.""" + + monkeypatch.setenv("PGDATABASE", "catalog") + monkeypatch.delenv("PGHOST", raising=False) + + with pytest.raises(SystemExit) as exc_info: + local_snapshot_cli.build_parser().parse_args([]) + + assert exc_info.value.code == 2 + assert "--host" in capsys.readouterr().err + + +@pytest.mark.parametrize("value", ["not-a-port", "0", "65536"]) +def test_port_rejects_invalid_values(value: str) -> None: + """Reject nonnumeric and out-of-range PostgreSQL ports.""" + + with pytest.raises(argparse.ArgumentTypeError): + local_snapshot_cli._port(value) + + +def test_schema_name_rejects_sql_fragments() -> None: + """Accept one identifier and reject SQL-fragment schema values.""" + + assert local_snapshot_cli._schema_name("catalog_v2") == "catalog_v2" + with pytest.raises(argparse.ArgumentTypeError): + local_snapshot_cli._schema_name("public; DROP SCHEMA public") + + +@pytest.mark.asyncio +async def test_capture_uses_local_connection_without_environment_password_or_dsn( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + """Block PGPASSWORD inheritance, bound connection inputs, and close afterward.""" + + captured: dict[str, Any] = {} + connection = FakeConnection() + + async def fake_connect(**kwargs: object) -> FakeConnection: + captured.update(kwargs) + return connection + + async def fake_collect(conn: FakeConnection, schema: str | None) -> dict: + assert conn is connection + return {"schema_filter": schema, "relations": []} + + monkeypatch.setenv("PGPASSWORD", "must-not-be-inherited") + monkeypatch.setattr(local_snapshot_cli.asyncpg, "connect", fake_connect) + monkeypatch.setattr( + local_snapshot_cli, + "collect_postgres_snapshot", + fake_collect, + ) + args = argparse.Namespace( + database="catalog", + host=str(tmp_path), + port=5432, + user="operator", + schema="public", + pretty=False, + ) + + snapshot = await local_snapshot_cli.capture_local_snapshot(args) + + assert snapshot == {"schema_filter": "public", "relations": []} + assert captured == { + "database": "catalog", + "host": str(tmp_path), + "password": "", + "port": 5432, + "timeout": 10, + "user": "operator", + } + assert "dsn" not in captured + assert captured["password"] != "must-not-be-inherited" + assert connection.closed is True + + +@pytest.mark.parametrize("pretty", [False, True]) +def test_main_writes_snapshot_json( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + capsys: pytest.CaptureFixture[str], + pretty: bool, +) -> None: + """Write equivalent compact and pretty snapshot JSON representations.""" + + expected = {"schema_filter": "public", "relations": []} + + async def fake_capture(args: argparse.Namespace) -> dict: + assert args.database == "catalog" + assert args.host == str(tmp_path) + assert args.schema == "public" + return expected + + monkeypatch.setattr(local_snapshot_cli, "capture_local_snapshot", fake_capture) + argv = [ + "--database", + "catalog", + "--host", + str(tmp_path), + "--schema", + "public", + ] + if pretty: + argv.append("--pretty") + + status = local_snapshot_cli.main(argv) + + output = capsys.readouterr() + assert status == 0 + assert output.err == "" + assert json.loads(output.out) == expected + if pretty: + assert output.out.startswith("{\n ") + else: + assert output.out == '{"relations":[],"schema_filter":"public"}\n' + + +@pytest.mark.parametrize( + ("error", "error_name"), + [ + (OSError("socket denied"), "OSError"), + (TimeoutError("connection timed out"), "TimeoutError"), + (asyncio.TimeoutError("async connection timed out"), "TimeoutError"), + (local_snapshot_cli.asyncpg.PostgresError("database denied"), "PostgresError"), + ], +) +def test_main_redacts_connection_errors( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + capsys: pytest.CaptureFixture[str], + error: Exception, + error_name: str, +) -> None: + """Return a stable error type without disclosing connection details.""" + + async def fail_capture(_args: argparse.Namespace) -> dict: + raise error + + monkeypatch.setattr(local_snapshot_cli, "capture_local_snapshot", fail_capture) + + status = local_snapshot_cli.main( + ["--database", "catalog", "--host", str(tmp_path)] + ) + + assert status == 1 + assert capsys.readouterr().err == f"pg-erd-snapshot failed: {error_name}\n" diff --git a/backend/tests/test_pg_snapshot_collect.py b/backend/tests/test_pg_snapshot_collect.py new file mode 100644 index 000000000..041be1b93 --- /dev/null +++ b/backend/tests/test_pg_snapshot_collect.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +from typing import Literal + +import pytest + +from app.pg_introspect import snapshot_collect + + +CitusMode = Literal["absent", "present", "missing_catalog"] + + +class MissingCitusCatalogError(Exception): + """Test-only stand-in for asyncpg.UndefinedTableError.""" + + +class FakeConnection: + """Return deterministic catalog results for the shared snapshot collector.""" + + def __init__(self, citus_mode: CitusMode) -> None: + self.citus_mode = citus_mode + self.fetch_count = 0 + + async def fetchval(self, query: str, *_args: object) -> object: + """Return a stable server version or Citus-extension flag.""" + + if query == "SHOW server_version": + return "17.10" + return self.citus_mode != "absent" + + async def fetch(self, *_args: object) -> list[dict[str, object]]: + """Return empty catalogs and model optional Citus catalog behavior.""" + + self.fetch_count += 1 + if self.fetch_count == 8: + if self.citus_mode == "missing_catalog": + raise MissingCitusCatalogError + return [{"logicalrelid": "public.orders"}] + return [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("citus_mode", "expected_citus", "expected_fetch_count"), + [ + ("absent", [], 7), + ("present", [{"logicalrelid": "public.orders"}], 8), + ("missing_catalog", [], 8), + ], +) +async def test_collect_postgres_snapshot_handles_each_citus_state( + monkeypatch: pytest.MonkeyPatch, + citus_mode: CitusMode, + expected_citus: list[dict[str, object]], + expected_fetch_count: int, +) -> None: + """Validate collection when Citus is absent, present, or lacks its catalog.""" + + monkeypatch.setattr( + snapshot_collect.asyncpg, + "UndefinedTableError", + MissingCitusCatalogError, + ) + connection = FakeConnection(citus_mode) + + result = await snapshot_collect.collect_postgres_snapshot(connection, "public") + + assert result["server_version"] == "17.10" + assert result["schema_filter"] == "public" + assert result["schemas"] == [] + assert result["relations"] == [] + assert result["columns"] == [] + assert result["constraints"] == [] + assert result["indexes"] == [] + assert result["pk_columns"] == [] + assert result["fk_edges"] == [] + assert result["citus_distributed_tables"] == expected_citus + assert isinstance(result["captured_at"], str) + assert connection.fetch_count == expected_fetch_count