Sfoglia il codice sorgente

feat(io): 抽 PG connect/resolve/query 到 dw_base/io/db/postgresql,sync-gen + ddl-gen 改调

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
tianyu.chu 1 mese fa
parent
commit
99ac990021

+ 4 - 42
bin/datax-sync-template-gen.py

@@ -30,7 +30,6 @@ CLI:
 """
 import argparse
 import os
-import re
 import sys
 from configparser import ConfigParser
 from datetime import datetime
@@ -38,8 +37,7 @@ from datetime import datetime
 project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
 sys.path.append(project_root)
 
-from dw_base.datax.datasources.data_source_factory import DataSourceFactory
-from dw_base.datax.datax_constants import DS_POSTGRE_SQL_JDBC_URL
+from dw_base.io.db import postgresql as pgdb
 
 
 WORKSPACE_DEFAULT = os.path.join(
@@ -54,37 +52,12 @@ PROBE_SAMPLE_LIMIT = 1000
 PROBE_SMALL_TABLE_THRESHOLD = 100000
 
 
-def resolve_datasource(ds_ref):
-    """复用 plugin.py:34-42 的 ref → DataSource 解析逻辑。
-
-    ds_ref 形如 'postgresql/prod-hobby',首段为 db_type(同父目录名)。
-    datasource ini 落点:项目同级 ../datasource/{ds_ref}.ini。
-    """
-    ds_type = ds_ref.split('/')[0]
-    if ds_type != 'postgresql':
-        raise NotImplementedError('暂只支持 postgresql 数据源,收到: ' + ds_type)
-    ds_file_path = os.path.normpath(
-        os.path.join(project_root, '..', 'datasource', ds_ref + '.ini'))
-    if not os.path.isfile(ds_file_path):
-        raise FileNotFoundError('数据源 ini 不存在: ' + ds_file_path)
-    return DataSourceFactory.get_data_source(ds_type, ds_file_path)
-
-
-def parse_jdbc_url(jdbc_url):
-    """从 jdbc:postgresql://host:port/database 抽 (host, port, database)。"""
-    m = re.match(r'jdbc:postgresql://([^:/]+)(?::(\d+))?/(.+)', jdbc_url)
-    if not m:
-        raise ValueError('无法解析 PG jdbcUrl: ' + jdbc_url)
-    return m.group(1), int(m.group(2) or 5432), m.group(3)
-
-
 def query_columns_full(conn, schema, table):
     """带序号 / 类型 / 主键标识的全字段 metadata 查询,按 attnum 排序。
 
     返回 [(attnum, attname, comment, pg_type, pk_flag), ...]
     """
-    cur = conn.cursor()
-    cur.execute("""
+    return pgdb.query(conn, """
         SELECT
             a.attnum,
             a.attname,
@@ -102,7 +75,6 @@ def query_columns_full(conn, schema, table):
           AND a.attnum > 0 AND NOT a.attisdropped
         ORDER BY a.attnum
     """, (schema, table))
-    return cur.fetchall()
 
 
 def probe_table(conn, schema, table, full_rows):
@@ -373,18 +345,8 @@ def main():
         sys.exit(2)
     schema, table = args.t.split('.', 1)
 
-    ds = resolve_datasource(args.ds)
-    ds_dict = ds.parse()
-    jdbc_url = ds_dict[DS_POSTGRE_SQL_JDBC_URL]
-    user = ds_dict['username']
-    password = ds_dict['password']
-    host, port, database = parse_jdbc_url(jdbc_url)
-
-    import pg8000.dbapi
-    conn = pg8000.dbapi.connect(
-        host=host, port=port, database=database,
-        user=user, password=password,
-    )
+    database = pgdb.resolve(args.ds)['database']
+    conn = pgdb.connect(args.ds)
     try:
         full_rows = query_columns_full(conn, schema, table)
         if not full_rows:

+ 5 - 16
bin/hive-ddl-gen.py

@@ -4,7 +4,7 @@
 Hive DDL 生成器(raw / ods 双层)。
 
 **仅支持 PG 源**:reader.dataSource 必须是 `postgresql/{env}-{instance}`
-形式;mysql 等其他源由复用的 datax-sync-template-gen.resolve_datasource
+形式;mysql 等其他源由 dw_base.io.db.postgresql 的 _resolve_datasource
 直接 NotImplementedError。
 
 输入 sync ini,从 PG 抽字段类型 + 中文注释:
@@ -37,11 +37,11 @@ from datetime import datetime
 project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
 sys.path.append(project_root)
 
+from dw_base.io.db import postgresql as pgdb
+
 
 def _load_sync_gen():
-    """复用 datax-sync-template-gen 的 resolve_datasource / parse_jdbc_url /
-    query_columns_full / DS_POSTGRE_SQL_JDBC_URL(脚本名含连字符,importlib 加载)。
-    """
+    """复用 datax-sync-template-gen 的 query_columns_full(脚本名含连字符,importlib 加载)。"""
     spec = importlib.util.spec_from_file_location(
         'datax_sync_template_gen',
         os.path.join(project_root, 'bin', 'datax-sync-template-gen.py'),
@@ -127,18 +127,7 @@ def fetch_column_full_rows(ds_ref, schema, table):
 
 
 def _fetch_pg_column_rows(ds_ref, schema, table):
-    ds = SYNC_GEN.resolve_datasource(ds_ref)
-    ds_dict = ds.parse()
-    jdbc_url = ds_dict[SYNC_GEN.DS_POSTGRE_SQL_JDBC_URL]
-    user = ds_dict['username']
-    password = ds_dict['password']
-    host, port, database = SYNC_GEN.parse_jdbc_url(jdbc_url)
-
-    import pg8000.dbapi
-    conn = pg8000.dbapi.connect(
-        host=host, port=port, database=database,
-        user=user, password=password,
-    )
+    conn = pgdb.connect(ds_ref)
     try:
         return SYNC_GEN.query_columns_full(conn, schema, table)
     finally:

+ 75 - 0
dw_base/io/db/postgresql.py

@@ -0,0 +1,75 @@
+# -*- coding:utf-8 -*-
+"""PG 连接工厂 + 薄查询封装(pg8000)。
+
+ds_ref(形如 postgresql/prd-hobby)→ 解析项目同级 ../datasource/{ds_ref}.ini
+→ pg8000 活连接。ds_ref → 凭据 的解析复用 datax 的 DataSourceFactory。
+"""
+import os
+import re
+
+from dw_base.datax.datasources.data_source_factory import DataSourceFactory
+from dw_base.datax.datax_constants import DS_POSTGRE_SQL_JDBC_URL
+
+# dw_base/io/db/postgresql.py → 上 4 层到项目根
+_PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.dirname(
+    os.path.dirname(os.path.abspath(__file__)))))
+
+
+def parse_jdbc_url(jdbc_url):
+    """从 jdbc:postgresql://host:port/database 抽 (host, port, database)。"""
+    m = re.match(r'jdbc:postgresql://([^:/]+)(?::(\d+))?/(.+)', jdbc_url)
+    if not m:
+        raise ValueError('无法解析 PG jdbcUrl: ' + jdbc_url)
+    return m.group(1), int(m.group(2) or 5432), m.group(3)
+
+
+def _resolve_datasource(ds_ref):
+    """ds_ref → DataSource。
+
+    ds_ref 形如 'postgresql/prd-hobby',首段为 db_type(同父目录名)。
+    datasource ini 落点:项目同级 ../datasource/{ds_ref}.ini。
+    """
+    ds_type = ds_ref.split('/')[0]
+    if ds_type != 'postgresql':
+        raise NotImplementedError('暂只支持 postgresql 数据源,收到: ' + ds_type)
+    ds_file_path = os.path.normpath(
+        os.path.join(_PROJECT_ROOT, '..', 'datasource', ds_ref + '.ini'))
+    if not os.path.isfile(ds_file_path):
+        raise FileNotFoundError('数据源 ini 不存在: ' + ds_file_path)
+    return DataSourceFactory.get_data_source(ds_type, ds_file_path)
+
+
+def resolve(ds_ref):
+    """ds_ref → {host, port, database, username, password}。"""
+    d = _resolve_datasource(ds_ref).parse()
+    host, port, database = parse_jdbc_url(d[DS_POSTGRE_SQL_JDBC_URL])
+    return {
+        'host': host,
+        'port': port,
+        'database': database,
+        'username': d['username'],
+        'password': d['password'],
+    }
+
+
+def connect(ds_ref):
+    """ds_ref → pg8000 活连接。调用方负责 close。"""
+    info = resolve(ds_ref)
+    import pg8000.dbapi
+    return pg8000.dbapi.connect(
+        host=info['host'], port=info['port'], database=info['database'],
+        user=info['username'], password=info['password'],
+    )
+
+
+def query(conn, sql, params=None):
+    """跑 SQL 返回 fetchall。params 为 None 时不传参(避开空参数风格差异)。"""
+    cur = conn.cursor()
+    try:
+        if params is None:
+            cur.execute(sql)
+        else:
+            cur.execute(sql, params)
+        return cur.fetchall()
+    finally:
+        cur.close()

+ 1 - 12
tests/unit/datax/test_hive_ddl_gen.py

@@ -159,24 +159,13 @@ def _patch_main_dependencies(monkeypatch, tmp_path):
         encoding='utf-8',
     )
 
-    fake_ds = MagicMock()
-    fake_ds.parse.return_value = {
-        GEN.SYNC_GEN.DS_POSTGRE_SQL_JDBC_URL: 'jdbc:postgresql://10.0.0.1:5432/mydb',
-        'username': 'u',
-        'password': 'p',
-    }
-    monkeypatch.setattr(GEN.SYNC_GEN, 'resolve_datasource', lambda ref: fake_ds)
-
     fake_conn = MagicMock()
     fake_cur = fake_conn.cursor.return_value
     fake_cur.fetchall.return_value = [
         (1, 'id', 'id', 'bigint', 'PK'),
         (2, 'name', '姓名', 'character varying', ''),
     ]
-    fake_pg8000 = MagicMock()
-    fake_pg8000.dbapi.connect.return_value = fake_conn
-    monkeypatch.setitem(sys.modules, 'pg8000', fake_pg8000)
-    monkeypatch.setitem(sys.modules, 'pg8000.dbapi', fake_pg8000.dbapi)
+    monkeypatch.setattr(GEN.pgdb, 'connect', lambda ref: fake_conn)
 
     return str(sync_ini)
 

+ 5 - 31
tests/unit/datax/test_sync_template_gen.py

@@ -27,25 +27,6 @@ def _load_script():
 GEN = _load_script()
 
 
-def test_parse_jdbc_url_with_port():
-    host, port, db = GEN.parse_jdbc_url('jdbc:postgresql://10.0.0.1:5433/hobby_stocks')
-    assert host == '10.0.0.1'
-    assert port == 5433
-    assert db == 'hobby_stocks'
-
-
-def test_parse_jdbc_url_default_port():
-    host, port, db = GEN.parse_jdbc_url('jdbc:postgresql://pg.example.com/mydb')
-    assert host == 'pg.example.com'
-    assert port == 5432
-    assert db == 'mydb'
-
-
-def test_parse_jdbc_url_invalid():
-    with pytest.raises(ValueError, match='无法解析'):
-        GEN.parse_jdbc_url('mysql://10.0.0.1:3306/foo')
-
-
 def test_render_template_includes_required_fields():
     columns = [('id', 'id'), ('name', '姓名'), ('create_time', '创建时间')]
     out = GEN.render_template(
@@ -175,24 +156,17 @@ def test_render_template_empty_pk():
 
 def _patch_main_dependencies(monkeypatch):
     """共享 mock:让 main() 不连真 PG / 真 datasource。"""
-    fake_ds = MagicMock()
-    fake_ds.parse.return_value = {
-        GEN.DS_POSTGRE_SQL_JDBC_URL: 'jdbc:postgresql://10.0.0.1:5432/mydb',
-        'username': 'u',
-        'password': 'p',
-    }
-    monkeypatch.setattr(GEN, 'resolve_datasource', lambda ref: fake_ds)
-
     fake_conn = MagicMock()
     fake_cur = fake_conn.cursor.return_value
     fake_cur.fetchall.return_value = [
         (1, 'id', 'id', 'bigint', 'PK'),
         (2, 'name', '名称', 'character varying', ''),
     ]
-    fake_pg8000 = MagicMock()
-    fake_pg8000.dbapi.connect.return_value = fake_conn
-    monkeypatch.setitem(sys.modules, 'pg8000', fake_pg8000)
-    monkeypatch.setitem(sys.modules, 'pg8000.dbapi', fake_pg8000.dbapi)
+    monkeypatch.setattr(GEN.pgdb, 'resolve', lambda ref: {
+        'host': '10.0.0.1', 'port': 5432, 'database': 'mydb',
+        'username': 'u', 'password': 'p',
+    })
+    monkeypatch.setattr(GEN.pgdb, 'connect', lambda ref: fake_conn)
 
 
 def test_main_stdout_only_when_no_o(monkeypatch, capsys):

+ 75 - 0
tests/unit/io/db/test_postgresql.py

@@ -0,0 +1,75 @@
+# -*- coding:utf-8 -*-
+"""dw_base.io.db.postgresql 连接工厂 + 查询封装单测(不连真 PG)。"""
+import sys
+from unittest.mock import MagicMock
+
+import pytest
+
+from dw_base.io.db import postgresql as pgdb
+
+
+def test_parse_jdbc_url_with_port():
+    host, port, db = pgdb.parse_jdbc_url('jdbc:postgresql://10.0.0.1:5433/hobby_stocks')
+    assert host == '10.0.0.1'
+    assert port == 5433
+    assert db == 'hobby_stocks'
+
+
+def test_parse_jdbc_url_default_port():
+    host, port, db = pgdb.parse_jdbc_url('jdbc:postgresql://pg.example.com/mydb')
+    assert host == 'pg.example.com'
+    assert port == 5432
+    assert db == 'mydb'
+
+
+def test_parse_jdbc_url_invalid():
+    with pytest.raises(ValueError, match='无法解析'):
+        pgdb.parse_jdbc_url('mysql://10.0.0.1:3306/foo')
+
+
+def test_resolve_returns_conn_params(monkeypatch):
+    fake_ds = MagicMock()
+    fake_ds.parse.return_value = {
+        pgdb.DS_POSTGRE_SQL_JDBC_URL: 'jdbc:postgresql://10.0.0.1:5432/mydb',
+        'username': 'u',
+        'password': 'p',
+    }
+    monkeypatch.setattr(pgdb, '_resolve_datasource', lambda ref: fake_ds)
+    assert pgdb.resolve('postgresql/prd-hobby') == {
+        'host': '10.0.0.1', 'port': 5432, 'database': 'mydb',
+        'username': 'u', 'password': 'p',
+    }
+
+
+def test_connect_passes_resolved_params(monkeypatch):
+    monkeypatch.setattr(pgdb, 'resolve', lambda ref: {
+        'host': 'h', 'port': 5432, 'database': 'mydb',
+        'username': 'u', 'password': 'p',
+    })
+    fake_pg8000 = MagicMock()
+    monkeypatch.setitem(sys.modules, 'pg8000', fake_pg8000)
+    monkeypatch.setitem(sys.modules, 'pg8000.dbapi', fake_pg8000.dbapi)
+
+    pgdb.connect('postgresql/prd-hobby')
+    fake_pg8000.dbapi.connect.assert_called_once_with(
+        host='h', port=5432, database='mydb', user='u', password='p',
+    )
+
+
+def test_query_with_params():
+    conn = MagicMock()
+    cur = conn.cursor.return_value
+    cur.fetchall.return_value = [(1, 'a'), (2, 'b')]
+    rows = pgdb.query(conn, 'SELECT 1 WHERE x = %s', ('v',))
+    assert rows == [(1, 'a'), (2, 'b')]
+    cur.execute.assert_called_once_with('SELECT 1 WHERE x = %s', ('v',))
+    cur.close.assert_called_once()
+
+
+def test_query_without_params():
+    conn = MagicMock()
+    cur = conn.cursor.return_value
+    cur.fetchall.return_value = []
+    pgdb.query(conn, 'SELECT 1')
+    cur.execute.assert_called_once_with('SELECT 1')
+    cur.close.assert_called_once()