refactor(manual): упрощение архитектуры учебного пайплайна и документации

- Зачем:
  - сделать примеры более доступными для начинающих, убрав лишнюю архитектурную сложность (хелперы, внешние SQL файлы).
  - сфокусировать обучение на самом Airflow, а не на структуре Python-проекта.
- Что:
  - консолидирована логика: DDL и вспомогательные функции перенесены из helpers/ и sql/ напрямую в csv_to_postgres.py и csv_to_postgres_dq.py.
  - удалена избыточная инфраструктура: папки sql/, tests/, helpers/ и файл requirements.txt больше не содержат специфичных для этого пайплайна файлов.
  - упрощена работа с БД: использование PostgresHook напрямую в задачах вместо кастомных оберток.
  - исправлен порт PostgreSQL для метаданных Airflow (5433 → 5434) в README.md и educational-setup-plan.md.
  - изменен путь генерации CSV на /opt/airflow/data/output в csv_to_postgres.py.
  - добавлен шаблон data/*.csv в .gitignore.
  - обновлен комментарий в sql_basic_dag.py о создании соединения.
  - исправлено описание практического задания в educational-tasks.md.
  - добавлены переводы строк в конце файлов csv_to_postgres.py и educational-setup-plan.md.
- Проверка:
  - запуск DAG-ов csv_to_postgres и csv_to_postgres_dq в Airflow UI.
This commit is contained in:
2026-03-08 21:38:59 +03:00
parent 24ee9f65f0
commit af6b0a9fae
15 changed files with 144 additions and 572 deletions
+1
View File
@@ -7,5 +7,6 @@ __pycache__/
*.pyc
.venv/
venv/
data/*.csv
data/output/
data/input/
+6 -10
View File
@@ -36,7 +36,7 @@ docker-compose up -d
- Пользователь: `student`
- Пароль: `student`
- **PostgreSQL для метаданных Airflow**: `localhost:5433`
- **PostgreSQL для метаданных Airflow**: `localhost:5434`
- База данных: `airflow`
- Пользователь: `airflow`
- Пароль: `airflow`
@@ -101,15 +101,9 @@ airflow-docker/
│ ├── data_processing_dag.py # Сложный ETL пайплайн
│ ├── branching_dag.py # Условная логика
│ └── error_handling_dag.py # Обработка ошибок
├── helpers/ # Вспомогательные скрипты
│ └── postgres.py # Функции для работы с БД и DQ
├── tests/ # Unit-тесты для хелперов
│ └── test_postgres_helpers.py # Тестирование DQ проверок
├── sql/ # SQL скрипты и DDL
│ └── base/ # Базовые DDL таблиц
├── data/ # Данные для упражнений
├── data/ # Данные и артефакты прогонов
│ ├── input/ # Входные данные
│ └── output/ # Результаты обработки
│ └── output/ # Сгенерированные CSV, отчеты и результаты обработки
├── logs/ # Логи Airflow
├── README.md # Эта инструкция
└── educational-tasks.md # Практические задания для студентов
@@ -199,12 +193,14 @@ airflow-docker/
docker-compose exec airflow-webserver airflow connections get postgres_training
```
`csv_to_postgres.py` по умолчанию складывает сгенерированные CSV в `/opt/airflow/data/output`, то есть в локальный каталог `airflow-docker/data/output/`.
### Порты
- `8080` - Airflow Webserver
- `5432` - PostgreSQL для тренировок
- `5433` - PostgreSQL для метаданных Airflow
- `5434` - PostgreSQL для метаданных Airflow
## 🛠️ Управление стендом
+1 -1
View File
@@ -139,7 +139,7 @@ Run automated data quality checks on the public.orders table after CSV loading.
- Minimum row count (> 0)
- Unique order_id values (no duplicates)
**Helper Functions:** Located in `dags/helpers/postgres.py`.
**Helper Functions:** All DQ check functions are defined directly in `csv_to_postgres_dq.py`.
#### 2.4 data_processing_dag.py
**Learning Objectives:**
+41 -33
View File
@@ -5,15 +5,19 @@ import os
import random
from datetime import UTC, datetime, timedelta
from pathlib import Path
from typing import List
import pandas as pd
from airflow.operators.python import PythonOperator
from helpers.postgres import get_postgres_conn
from airflow.providers.postgres.hooks.postgres import PostgresHook
from airflow import DAG
CSV_DIR = Path(os.getenv("CSV_DIR", "/opt/airflow/data"))
POSTGRES_CONN_ID = "postgres_training"
def _get_conn():
return PostgresHook(postgres_conn_id=POSTGRES_CONN_ID).get_conn()
CSV_DIR = Path(os.getenv("CSV_DIR", "/opt/airflow/data/output"))
CSV_ROWS = int(os.getenv("CSV_ROWS", "1000"))
@@ -27,9 +31,13 @@ def _create_table() -> None:
amount NUMERIC(12,2) NOT NULL
);
"""
with get_postgres_conn() as conn, conn.cursor() as cur:
cur.execute(ddl)
conn.commit()
conn = _get_conn()
try:
with conn.cursor() as cur:
cur.execute(ddl)
conn.commit()
finally:
conn.close()
def _generate_csv(rows: int, csv_dir: Path) -> str:
@@ -93,33 +101,33 @@ def _load_csv(csv_path: str) -> None:
if not csv_file.exists():
raise FileNotFoundError(f"CSV не найден: {csv_file}")
with (
get_postgres_conn() as conn,
conn.cursor() as cur,
csv_file.open("r", encoding="utf-8") as f,
):
cur.execute(
"CREATE TEMP TABLE tmp_orders (LIKE public.orders INCLUDING DEFAULTS) ON COMMIT DROP;"
)
cur.copy_expert(
"COPY tmp_orders (order_id, order_ts, customer_id, amount) FROM STDIN WITH CSV HEADER",
f,
)
conn = _get_conn()
try:
with conn.cursor() as cur, csv_file.open("r", encoding="utf-8") as f:
cur.execute(
"CREATE TEMP TABLE tmp_orders (LIKE public.orders INCLUDING DEFAULTS) ON COMMIT DROP;"
)
cur.copy_expert(
"COPY tmp_orders (order_id, order_ts, customer_id, amount) FROM STDIN WITH CSV HEADER",
f,
)
cur.execute("SELECT COUNT(*) FROM tmp_orders")
tmp_rows = cur.fetchone()[0]
cur.execute("SELECT COUNT(*) FROM tmp_orders")
tmp_rows = cur.fetchone()[0]
cur.execute(
"""
INSERT INTO public.orders(order_id, order_ts, customer_id, amount)
SELECT t.order_id, t.order_ts, t.customer_id, t.amount
FROM tmp_orders t
LEFT JOIN public.orders o ON o.order_id = t.order_id
WHERE o.order_id IS NULL
"""
)
inserted = cur.rowcount if cur.rowcount != -1 else 0
conn.commit()
cur.execute(
"""
INSERT INTO public.orders(order_id, order_ts, customer_id, amount)
SELECT t.order_id, t.order_ts, t.customer_id, t.amount
FROM tmp_orders t
LEFT JOIN public.orders o ON o.order_id = t.order_id
WHERE o.order_id IS NULL
"""
)
inserted = cur.rowcount if cur.rowcount != -1 else 0
conn.commit()
finally:
conn.close()
logging.info("Загружено строк: %s (прочитано из CSV: %s)", inserted, tmp_rows)
@@ -128,7 +136,7 @@ default_args = {"owner": "airflow", "retries": 1, "retry_delay": timedelta(secon
with DAG(
dag_id="csv_to_postgres",
start_date=datetime(2017, 1, 1),
start_date=datetime(2023, 1, 1),
schedule=None,
catchup=False,
default_args=default_args,
+88 -47
View File
@@ -4,94 +4,135 @@ import logging
from datetime import datetime, timedelta
from airflow.operators.python import PythonOperator
from helpers.postgres import (
assert_orders_have_rows,
assert_orders_no_duplicates,
assert_orders_schema,
assert_orders_table_exists,
get_postgres_conn,
)
from airflow.providers.postgres.hooks.postgres import PostgresHook
from airflow import DAG
POSTGRES_CONN_ID = "postgres_training"
def _run_check(check_callable):
"""
Оборачивает проверку качества данных в контекст подключения к Postgres.
EXPECTED_ORDERS_SCHEMA = [
("order_id", "bigint"),
("order_ts", "timestamp without time zone"),
("customer_id", "bigint"),
("amount", "numeric"),
]
Этот DAG предназначен для автоматической проверки качества данных
после CSV-пайплайна в таблице public.orders:
1. Проверяет существование таблицы
2. Проверяет соответствие схемы
3. Проверяет наличие данных
4. Проверяет отсутствие дубликатов
Args:
check_callable: Функция проверки, принимающая подключение к БД
"""
# Получаем имя функции для логов
check_name = check_callable.__name__.replace("assert_", "")
logging.info("🚀 Запуск проверки: %s", check_name)
def _get_conn():
return PostgresHook(postgres_conn_id=POSTGRES_CONN_ID).get_conn()
with get_postgres_conn() as conn:
check_callable(conn)
logging.info("✅ Проверка пройдена: %s", check_name)
def _check_table_exists():
"""Проверяет наличие таблицы public.orders."""
conn = _get_conn()
try:
with conn.cursor() as cur:
cur.execute("""
SELECT 1 FROM pg_catalog.pg_tables
WHERE schemaname = 'public' AND tablename = 'orders'
""")
if cur.fetchone() is None:
raise ValueError("Таблица public.orders не найдена")
logging.info("Таблица public.orders существует")
finally:
conn.close()
def _check_schema():
"""Проверяет соответствие схемы таблицы public.orders ожидаемой."""
conn = _get_conn()
try:
with conn.cursor() as cur:
cur.execute("""
SELECT column_name, data_type
FROM information_schema.columns
WHERE table_schema = 'public' AND table_name = 'orders'
ORDER BY ordinal_position
""")
actual = cur.fetchall()
if actual != EXPECTED_ORDERS_SCHEMA:
raise ValueError(
f"Схема не совпадает. Ожидалось: {EXPECTED_ORDERS_SCHEMA}, "
f"получено: {actual}"
)
logging.info("Схема таблицы public.orders соответствует ожидаемой")
finally:
conn.close()
def _check_has_rows():
"""Проверяет, что в таблице public.orders есть данные."""
conn = _get_conn()
try:
with conn.cursor() as cur:
cur.execute("SELECT COUNT(*) FROM public.orders")
count = cur.fetchone()[0]
if count == 0:
raise ValueError("Таблица public.orders пуста")
logging.info("Таблица public.orders содержит %s строк", count)
finally:
conn.close()
def _check_no_duplicates():
"""Проверяет отсутствие дубликатов order_id в public.orders."""
conn = _get_conn()
try:
with conn.cursor() as cur:
cur.execute("""
SELECT order_id, COUNT(*) AS cnt
FROM public.orders
GROUP BY order_id
HAVING COUNT(*) > 1
""")
duplicates = cur.fetchall()
if duplicates:
raise ValueError(f"Обнаружены дубликаты order_id: {duplicates}")
logging.info("Дубликатов order_id не обнаружено")
finally:
conn.close()
def _log_dq_summary():
"""
Логирует итоговую сводку по качеству данных.
Эта задача выполняется после всех проверок и показывает общий результат.
"""
logging.info("🎉 Все проверки качества данных пройдены успешно!")
logging.info("📊 Качество данных в таблице orders соответствует требованиям.")
"""Логирует итоговую сводку по качеству данных."""
logging.info("Все проверки качества данных пройдены успешно!")
logging.info("Качество данных в таблице orders соответствует требованиям.")
default_args = {"owner": "airflow", "retries": 1, "retry_delay": timedelta(seconds=30)}
with DAG(
dag_id="csv_to_postgres_dq",
start_date=datetime(2017, 1, 1),
start_date=datetime(2023, 1, 1),
schedule=None,
catchup=False,
default_args=default_args,
tags=["demo", "postgres", "quality", "csv", "dq"],
description="Проверки качества данных после CSV public.orders в Postgres",
description="Проверки качества данных после CSV -> public.orders в Postgres",
) as dag:
# Задача 1: Проверка существования таблицы
check_exists = PythonOperator(
task_id="check_orders_table_exists",
python_callable=_run_check,
op_args=[assert_orders_table_exists],
python_callable=_check_table_exists,
)
# Задача 2: Проверка соответствия схемы таблицы
check_schema = PythonOperator(
task_id="check_orders_schema",
python_callable=_run_check,
op_args=[assert_orders_schema],
python_callable=_check_schema,
)
# Задача 3: Проверка наличия данных
check_has_rows = PythonOperator(
task_id="check_orders_has_rows",
python_callable=_run_check,
op_args=[assert_orders_have_rows],
python_callable=_check_has_rows,
)
# Задача 4: Проверка отсутствия дубликатов
check_no_duplicates = PythonOperator(
task_id="check_order_duplicates",
python_callable=_run_check,
op_args=[assert_orders_no_duplicates],
python_callable=_check_no_duplicates,
)
# Задача 5: Итоговая сводка
dq_summary = PythonOperator(
task_id="data_quality_summary",
python_callable=_log_dq_summary,
)
# Определяем последовательность выполнения задач
check_exists >> check_schema >> check_has_rows >> check_no_duplicates >> dq_summary
-32
View File
@@ -1,32 +0,0 @@
from __future__ import annotations
"""
Учебный DAG: применяет DDL для базовой таблицы orders в Postgres.
Запускается вручную перед CSVпайплайном или после изменения схемы.
"""
from datetime import datetime, timedelta
from airflow.providers.postgres.operators.postgres import PostgresOperator
from airflow import DAG
POSTGRES_CONN_ID = "postgres_training"
default_args = {"owner": "airflow", "retries": 1, "retry_delay": timedelta(seconds=30)}
with DAG(
dag_id="orders_base_ddl",
start_date=datetime(2017, 1, 1),
schedule=None,
catchup=False,
template_searchpath="/opt/airflow/sql",
default_args=default_args,
tags=["demo", "postgres", "ddl", "orders"],
description="Создаёт/обновляет базовую таблицу orders в схеме public",
) as dag:
apply_orders_ddl = PostgresOperator(
task_id="apply_orders_ddl",
postgres_conn_id=POSTGRES_CONN_ID,
sql="base/orders_ddl.sql",
)
-223
View File
@@ -1,223 +0,0 @@
from __future__ import annotations
"""
LEGACY: Вспомогательные функции для прямого подключения к Postgres через psycopg2.
Внимание: этот модуль оставлен только для поддержки базового CSV-пайплайна.
В новых DAG (ODS/DDS/DM) используйте встроенный в Airflow PostgresOperator
и штатные механизмы XCom.
"""
import logging
import os
from typing import List, Sequence, Tuple
import psycopg2
# Настройки для подключения к Postgres. По умолчанию используем Airflow Connection,
# но при проблемах можно переключиться на ENV-подключение, установив POSTGRES_USE_AIRFLOW_CONN=false.
POSTGRES_CONN_ID = os.getenv("POSTGRES_CONN_ID", "postgres_training")
POSTGRES_USE_AIRFLOW_CONN = os.getenv("POSTGRES_USE_AIRFLOW_CONN", "true").lower() in (
"1",
"true",
"yes",
)
# Ожидаемая схема таблицы orders для проверки качества данных
EXPECTED_ORDERS_SCHEMA: List[Tuple[str, str]] = [
("order_id", "bigint"),
("order_ts", "timestamp without time zone"),
("customer_id", "bigint"),
("amount", "numeric"),
]
def get_postgres_conn():
"""
Возвращает psycopg2 connection к Postgres.
Приоритет подключения:
1. Через Airflow Connection (если настроено и доступно)
2. Прямое подключение по переменным окружения (фоллбек)
Returns:
psycopg2 connection object
"""
if POSTGRES_USE_AIRFLOW_CONN:
try:
from airflow.providers.postgres.hooks.postgres import PostgresHook
hook = PostgresHook(postgres_conn_id=POSTGRES_CONN_ID)
conn = hook.get_conn()
logging.info("✅ Подключение через Airflow Connection успешно")
return conn
except Exception as e:
logging.warning("⚠️ Не удалось подключиться через Airflow Connection: %s", e)
logging.info("🔄 Переключаемся на прямое подключение по ENV переменным")
# Фоллбек на прямое подключение по переменным окружения.
# Прямое подключение по переменным окружения
conn_params = {
"dbname": os.getenv("POSTGRES_DB", "training"),
"user": os.getenv("POSTGRES_USER", "student"),
"password": os.getenv("POSTGRES_PASSWORD", "student"),
"host": os.getenv("POSTGRES_HOST", "postgres-training"),
"port": int(os.getenv("POSTGRES_PORT", "5432")),
}
logging.info(
"🔗 Подключение к Postgres: %s:%s/%s",
conn_params["host"],
conn_params["port"],
conn_params["dbname"],
)
return psycopg2.connect(**conn_params)
def assert_orders_table_exists(conn) -> None:
"""
Проверяет наличие таблицы orders в схеме public.
Args:
conn: Подключение к Postgres
Raises:
ValueError: Если таблица не найдена
"""
logging.info("🔍 Проверяем существование таблицы public.orders...")
with conn.cursor() as cur:
cur.execute(
"""
SELECT 1
FROM pg_catalog.pg_tables
WHERE schemaname = 'public' AND tablename = 'orders'
"""
)
if cur.fetchone() is None:
raise ValueError(
"❌ Таблица public.orders не найдена; запусти DAG csv_to_postgres."
)
logging.info("✅ Таблица public.orders существует")
def fetch_orders_schema(conn) -> Sequence[Tuple[str, str]]:
"""
Получает схему таблицы orders из information_schema.
Args:
conn: Подключение к Postgres
Returns:
Список кортежей (имя_колонки, тип_данных)
"""
with conn.cursor() as cur:
cur.execute(
"""
SELECT column_name, data_type
FROM information_schema.columns
WHERE table_schema = 'public' AND table_name = 'orders'
ORDER BY ordinal_position
"""
)
return cur.fetchall()
def assert_orders_schema(conn) -> None:
"""
Проверяет, что схема таблицы orders соответствует ожидаемой.
Args:
conn: Подключение к Postgres
Raises:
ValueError: Если схема не соответствует ожидаемой
"""
logging.info("📋 Проверяем схему таблицы orders...")
schema = fetch_orders_schema(conn)
logging.info("📊 Фактическая схема: %s", list(schema))
logging.info("📊 Ожидаемая схема: %s", EXPECTED_ORDERS_SCHEMA)
if list(schema) != EXPECTED_ORDERS_SCHEMA:
raise ValueError(
f"❌ Неожиданная схема orders: {schema}. Ожидали {EXPECTED_ORDERS_SCHEMA}."
)
logging.info("✅ Схема таблицы orders соответствует ожиданиям")
def fetch_orders_count(conn) -> int:
"""
Получает количество строк в таблице orders.
Args:
conn: Подключение к Postgres
Returns:
Количество строк в таблице
"""
with conn.cursor() as cur:
cur.execute("SELECT COUNT(*) FROM public.orders")
return cur.fetchone()[0]
def assert_orders_have_rows(conn) -> None:
"""
Проверяет, что таблица orders не пустая.
Args:
conn: Подключение к Postgres
Raises:
ValueError: Если таблица пустая
"""
logging.info("📊 Проверяем наличие данных в таблице orders...")
row_count = fetch_orders_count(conn)
logging.info("📈 Количество строк в orders: %s", row_count)
if row_count <= 0:
raise ValueError(
"❌ Таблица public.orders пустая — запусти DAG csv_to_postgres перед проверкой."
)
logging.info("✅ Таблица orders содержит данные (%s строк)", row_count)
def fetch_orders_duplicates(conn) -> int:
"""
Подсчитывает количество дубликатов по order_id.
Args:
conn: Подключение к Postgres
Returns:
Количество дублирующихся order_id
"""
with conn.cursor() as cur:
cur.execute(
"""
SELECT COUNT(*) FROM (
SELECT order_id
FROM public.orders
GROUP BY order_id
HAVING COUNT(*) > 1
) d
"""
)
return cur.fetchone()[0]
def assert_orders_no_duplicates(conn) -> None:
"""
Проверяет, что в таблице нет дублей по order_id.
Args:
conn: Подключение к Postgres
Raises:
ValueError: Если обнаружены дубликаты
"""
logging.info("🔍 Проверяем отсутствие дубликатов по order_id...")
duplicates = fetch_orders_duplicates(conn)
logging.info("📊 Найдено дубликатов: %s", duplicates)
if duplicates:
raise ValueError(
f"❌ Обнаружены дубли по order_id ({duplicates} шт.) — проверь загрузку данных."
)
logging.info("✅ Дубликаты не обнаружены")
+1 -1
View File
@@ -56,7 +56,7 @@ TRUNCATE TABLE students_sample;
# Определение задач
create_table_task = PostgresOperator(
task_id='create_table',
postgres_conn_id='postgres_training', # Это соединение нужно будет создать вручную в Airflow UI
postgres_conn_id='postgres_training', # Соединение создается автоматически в airflow-init
sql=create_table_sql,
dag=dag
)
-3
View File
@@ -61,7 +61,6 @@ services:
volumes:
- ./dags:/opt/airflow/dags
- ./data:/opt/airflow/data
- ./tests:/opt/airflow/tests
depends_on:
airflow-init:
condition: service_completed_successfully
@@ -83,7 +82,6 @@ services:
volumes:
- ./dags:/opt/airflow/dags
- ./data:/opt/airflow/data
- ./tests:/opt/airflow/tests
depends_on:
airflow-init:
condition: service_completed_successfully
@@ -102,7 +100,6 @@ services:
volumes:
- ./dags:/opt/airflow/dags
- ./data:/opt/airflow/data
- ./tests:/opt/airflow/tests
command: >
bash -ceuo pipefail "
mkdir -p /opt/airflow/data &&
+1 -1
View File
@@ -26,7 +26,7 @@ services:
POSTGRES_PASSWORD: airflow
POSTGRES_DB: airflow
ports:
- "5433:5432"
- "5434:5432"
volumes:
- pgmeta:/var/lib/postgresql/data
+2 -2
View File
@@ -148,7 +148,7 @@
**Время выполнения:** 20-25 минут
**Задача:**
- Перепишите задачу `create_orders_table`. Сейчас она использует `PythonOperator` и прямое подключение через `psycopg2`.
- Перепишите задачу `create_orders_table`. Сейчас она использует `PythonOperator` и `PostgresHook` внутри Python-функции.
- Замените её на использование стандартного `PostgresOperator`, используя заранее созданный Connection.
- Убедитесь, что пайплайн продолжает работать корректно.
@@ -165,7 +165,7 @@
**Время выполнения:** 20-25 минут
**Задача:**
- Добавьте новую функцию проверки в `helpers/postgres.py`, которая будет убеждаться, что все значения в колонке `amount` строго больше нуля.
- Добавьте новую функцию проверки прямо в `csv_to_postgres_dq.py`, которая будет убеждаться, что все значения в колонке `amount` строго больше нуля.
- Добавьте вызов этой функции как новую задачу в DAG `csv_to_postgres_dq`.
- Встройте новую задачу в общую цепочку выполнения (например, перед `dq_summary`).
-2
View File
@@ -13,5 +13,3 @@ mimesis==15.1.0
# Airflow PostgreSQL provider (used in sql_basic_dag.py and data_processing_dag.py)
apache-airflow-providers-postgres==5.11.1
# Testing framework
pytest==7.4.4
-9
View File
@@ -1,9 +0,0 @@
-- DDL для базовой таблицы orders, которую использует CSV‑pipeline.
-- Выполняется идемпотентно: таблица создаётся, если ещё не существует.
CREATE TABLE IF NOT EXISTS public.orders (
order_id BIGINT PRIMARY KEY,
order_ts TIMESTAMP NOT NULL,
customer_id BIGINT NOT NULL,
amount NUMERIC(12,2) NOT NULL
);
-20
View File
@@ -1,20 +0,0 @@
from __future__ import annotations
import pytest
def patch_postgres_hook(monkeypatch, fake_hook_class) -> None:
"""
Патчит PostgresHook для тестирования.
Args:
monkeypatch: pytest monkeypatch fixture
fake_hook_class: Класс-имитация PostgresHook
"""
# Патчим PostgresHook на уровне airflow.providers.postgres.hooks.postgres
# Это нужно, так как get_postgres_conn() импортирует его оттуда
monkeypatch.setattr(
"airflow.providers.postgres.hooks.postgres.PostgresHook",
fake_hook_class,
raising=False,
)
@@ -1,185 +0,0 @@
from __future__ import annotations
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any, List, Sequence
import pytest
# Добавляем путь к dags в sys.path для импорта модулей
dags_path = str(Path(__file__).parent.parent / "dags")
if dags_path not in sys.path:
sys.path.insert(0, dags_path)
import helpers.postgres as postgres_helpers
from conftest import patch_postgres_hook
@dataclass
class FakeCursor:
fetchone_value: Any = None
fetchall_value: Sequence[Any] | None = None
rowcount: int | None = None
def __post_init__(self) -> None:
self.queries: List[Any] = []
def execute(self, query: str, params: Any | None = None) -> None:
self.queries.append((query, params))
def fetchone(self) -> Any:
return self.fetchone_value
def fetchall(self) -> Sequence[Any] | None:
return self.fetchall_value
def __enter__(self) -> FakeCursor:
return self
def __exit__(self, exc_type, exc, tb) -> None:
return None
class FakeConn:
def __init__(self, cursors: Sequence[FakeCursor]) -> None:
self._cursors = list(cursors)
self._index = 0
self.commits = 0
def cursor(self) -> FakeCursor:
cursor = self._cursors[self._index]
self._index += 1
return cursor
def commit(self) -> None:
self.commits += 1
def test_get_postgres_conn_uses_airflow_hook(monkeypatch) -> None:
class FakeHook:
def __init__(self, postgres_conn_id: str) -> None:
self.postgres_conn_id = postgres_conn_id
def get_conn(self) -> str:
return "hook_connection"
patch_postgres_hook(monkeypatch, FakeHook)
monkeypatch.setattr(postgres_helpers, "POSTGRES_CONN_ID", "demo_conn", raising=False)
monkeypatch.setattr(postgres_helpers, "POSTGRES_USE_AIRFLOW_CONN", True, raising=False)
conn = postgres_helpers.get_postgres_conn()
assert conn == "hook_connection"
def test_get_postgres_conn_fallback_to_psycopg(monkeypatch) -> None:
class BrokenHook:
def __init__(self, postgres_conn_id: str) -> None:
self.postgres_conn_id = postgres_conn_id
def get_conn(self):
raise RuntimeError("boom")
patch_postgres_hook(monkeypatch, BrokenHook)
monkeypatch.setattr(postgres_helpers, "POSTGRES_USE_AIRFLOW_CONN", True, raising=False)
monkeypatch.setattr(postgres_helpers, "POSTGRES_CONN_ID", "demo_conn", raising=False)
monkeypatch.setenv("POSTGRES_DB", "demo_db")
monkeypatch.setenv("POSTGRES_USER", "demo_user")
monkeypatch.setenv("POSTGRES_PASSWORD", "secret")
monkeypatch.setenv("POSTGRES_HOST", "postgres-host")
monkeypatch.setenv("POSTGRES_PORT", "5434")
captured_kwargs = {}
def fake_connect(**kwargs):
captured_kwargs.update(kwargs)
return "psycopg_connection"
monkeypatch.setattr(postgres_helpers.psycopg2, "connect", fake_connect)
conn = postgres_helpers.get_postgres_conn()
assert conn == "psycopg_connection"
assert captured_kwargs == {
"dbname": "demo_db",
"user": "demo_user",
"password": "secret",
"host": "postgres-host",
"port": 5434,
}
def test_get_postgres_conn_without_airflow(monkeypatch) -> None:
monkeypatch.setattr(postgres_helpers, "POSTGRES_USE_AIRFLOW_CONN", False, raising=False)
monkeypatch.setenv("POSTGRES_DB", "demo_db")
monkeypatch.setenv("POSTGRES_USER", "demo_user")
monkeypatch.setenv("POSTGRES_PASSWORD", "secret")
monkeypatch.setenv("POSTGRES_HOST", "postgres-host")
monkeypatch.setenv("POSTGRES_PORT", "5435")
captured_kwargs = {}
def fake_connect(**kwargs):
captured_kwargs.update(kwargs)
return "direct_psycopg"
monkeypatch.setattr(postgres_helpers.psycopg2, "connect", fake_connect)
conn = postgres_helpers.get_postgres_conn()
assert conn == "direct_psycopg"
assert captured_kwargs["port"] == 5435
def test_assert_orders_table_exists_ok() -> None:
conn = FakeConn([FakeCursor(fetchone_value=(1,))])
postgres_helpers.assert_orders_table_exists(conn)
def test_assert_orders_table_exists_missing() -> None:
conn = FakeConn([FakeCursor(fetchone_value=None)])
with pytest.raises(ValueError):
postgres_helpers.assert_orders_table_exists(conn)
def test_assert_orders_schema_ok() -> None:
expected = list(postgres_helpers.EXPECTED_ORDERS_SCHEMA)
conn = FakeConn([FakeCursor(fetchall_value=expected)])
postgres_helpers.assert_orders_schema(conn)
def test_assert_orders_schema_mismatch() -> None:
conn = FakeConn([FakeCursor(fetchall_value=[("order_id", "bigint")])])
with pytest.raises(ValueError):
postgres_helpers.assert_orders_schema(conn)
def test_assert_orders_have_rows_ok() -> None:
conn = FakeConn([FakeCursor(fetchone_value=(5,))])
postgres_helpers.assert_orders_have_rows(conn)
def test_assert_orders_have_rows_empty() -> None:
conn = FakeConn([FakeCursor(fetchone_value=(0,))])
with pytest.raises(ValueError):
postgres_helpers.assert_orders_have_rows(conn)
def test_assert_orders_no_duplicates_ok() -> None:
conn = FakeConn([FakeCursor(fetchone_value=(0,))])
postgres_helpers.assert_orders_no_duplicates(conn)
def test_assert_orders_no_duplicates_detected() -> None:
conn = FakeConn([FakeCursor(fetchone_value=(3,))])
with pytest.raises(ValueError):
postgres_helpers.assert_orders_no_duplicates(conn)