baseline: 批次A-D 成果 + membership 半成品(测试红)
This commit is contained in:
@@ -0,0 +1,196 @@
|
||||
"""API 权限矩阵探测(开发/审计工具)。
|
||||
|
||||
做什么:枚举全部 URL,用「匿名 / 无租户头 / 错误租户 / 有效 JWT」四种身份发 GET,
|
||||
把状态码汇总成矩阵,标出**预期外开放**的端点(潜在越权)。
|
||||
|
||||
判定基线(本项目的既定设计):
|
||||
- 匿名能拿 200 的只应是"公开端点"白名单:ping / auth/token / demo/enter /
|
||||
billing/plans / open/statements / storefront 客户端入口
|
||||
- 其余端点匿名应为 401/403(受全局 IsAuthenticated 保护)
|
||||
- 带有效 JWT 但**无租户头**时,业务端点应 400(无法识别租户)而非 200
|
||||
|
||||
用法:
|
||||
python scripts/audit_api_permissions.py # 用 dev 库 + 自动登录 alice
|
||||
python scripts/audit_api_permissions.py --base http://127.0.0.1:9000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
os.environ.setdefault("DJANGO_SETTINGS_MODULE", "config.settings.dev")
|
||||
|
||||
import django # noqa: E402
|
||||
|
||||
django.setup()
|
||||
|
||||
from django.urls import get_resolver # noqa: E402
|
||||
|
||||
|
||||
# 公开端点前缀(匿名访问 200/4xx 都正常,不算问题)
|
||||
PUBLIC_PREFIXES = (
|
||||
"api/v1/ping/",
|
||||
"api/v1/auth/token/",
|
||||
"api/v1/demo/enter/",
|
||||
"api/v1/billing/plans/",
|
||||
"api/v1/open/statements/",
|
||||
)
|
||||
|
||||
# 需要替换路径参数的占位(探测用 1)
|
||||
PARAM_PATTERNS = [
|
||||
(re.compile(r"<int:[^>]+>"), "1"),
|
||||
(re.compile(r"<uuid:[^>]+>"), "00000000-0000-0000-0000-000000000000"),
|
||||
(re.compile(r"<[^>]+:[^>]+>"), "1"),
|
||||
(re.compile(r"\(\?P<[^>]+>\[\^/\.\]\+\)"), "1"),
|
||||
(re.compile(r"\\\.\(\?P<format>\[a-z0-9\]\+\)/\\?\$"), ""),
|
||||
]
|
||||
|
||||
|
||||
def collect_urls() -> list:
|
||||
"""枚举所有 API URL(把正则路径参数换成占位值)。"""
|
||||
resolver = get_resolver()
|
||||
raw = set()
|
||||
|
||||
def walk(patterns, prefix=""):
|
||||
for p in patterns:
|
||||
pat = prefix + str(p.pattern)
|
||||
if hasattr(p, "url_patterns"):
|
||||
walk(p.url_patterns, pat)
|
||||
else:
|
||||
raw.add(pat)
|
||||
|
||||
walk(resolver.url_patterns)
|
||||
|
||||
urls = set()
|
||||
for u in raw:
|
||||
if not u.startswith("api/"):
|
||||
continue
|
||||
if "<drf_format_suffix" in u:
|
||||
continue
|
||||
path = u
|
||||
for rx, repl in PARAM_PATTERNS:
|
||||
path = rx.sub(repl, path)
|
||||
if "(" in path or "?" in path: # 仍含正则残留,跳过
|
||||
continue
|
||||
path = path.rstrip("$").rstrip("^")
|
||||
if not path.startswith("api/"):
|
||||
path = "api/" + path
|
||||
urls.add("/" + path.lstrip("/"))
|
||||
return sorted(urls)
|
||||
|
||||
|
||||
def request(url: str, *, token: str = "", tenant: str = "") -> int:
|
||||
"""发 GET,返回状态码(0 = 连接/其他异常)。"""
|
||||
req = urllib.request.Request(url, method="GET")
|
||||
if token:
|
||||
req.add_header("Authorization", f"Bearer {token}")
|
||||
if tenant:
|
||||
req.add_header("X-Tenant-Id", tenant)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=10) as resp:
|
||||
return resp.status
|
||||
except urllib.error.HTTPError as e:
|
||||
return e.code
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
def login(base: str, username: str, password: str) -> str:
|
||||
body = json.dumps({"username": username, "password": password}).encode()
|
||||
req = urllib.request.Request(f"{base}/api/v1/auth/token/", data=body, method="POST")
|
||||
req.add_header("Content-Type", "application/json")
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=10) as resp:
|
||||
return json.loads(resp.read().decode())["access"]
|
||||
except Exception as exc:
|
||||
print(f"登录失败:{exc}")
|
||||
return ""
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--base", default="http://127.0.0.1:9000")
|
||||
ap.add_argument("--user", default="alice")
|
||||
ap.add_argument("--password", default="alice12345")
|
||||
ap.add_argument("--tenant", default="default")
|
||||
ap.add_argument("--verbose", action="store_true")
|
||||
args = ap.parse_args()
|
||||
|
||||
base = args.base.rstrip("/")
|
||||
token = login(base, args.user, args.password)
|
||||
if not token:
|
||||
print("无法登录,探测中止")
|
||||
return 2
|
||||
|
||||
urls = collect_urls()
|
||||
print(f"探测 {len(urls)} 个端点 · base={base} · 租户={args.tenant}\n")
|
||||
|
||||
findings = []
|
||||
rows = []
|
||||
for path in urls:
|
||||
url = base + path
|
||||
anon = request(url)
|
||||
no_tenant = request(url, token=token)
|
||||
ok = request(url, token=token, tenant=args.tenant)
|
||||
|
||||
is_public = any(path.startswith("/" + p) for p in PUBLIC_PREFIXES)
|
||||
rows.append((path, anon, no_tenant, ok, is_public))
|
||||
|
||||
# 问题 1:非公开端点匿名可访问
|
||||
if not is_public and anon == 200:
|
||||
findings.append(("匿名可访问", path, f"anon={anon}"))
|
||||
# 问题 2:带 JWT 但无租户头仍返回 200(应为 400 无法识别租户)
|
||||
if not is_public and no_tenant == 200 and ok != 200:
|
||||
findings.append(("缺租户头仍 200", path, f"no_tenant={no_tenant}"))
|
||||
# 问题 3:合法请求反而 500(服务端缺陷)
|
||||
if ok == 500:
|
||||
findings.append(("合法请求 500", path, "with_tenant=500"))
|
||||
# 问题 4:合法请求 0(路由不可达)
|
||||
if ok == 0:
|
||||
findings.append(("端点不可达", path, "network/route error"))
|
||||
|
||||
print(f"{'端点':<62} {'匿名':>5} {'无租户':>7} {'正常':>5}")
|
||||
print("-" * 84)
|
||||
for path, anon, nt, ok_, is_pub in rows:
|
||||
if args.verbose or anon == 200 or ok_ in (500, 0, 401):
|
||||
mark = " [公开]" if is_pub else ""
|
||||
print(f"{path:<62} {anon:>5} {nt:>7} {ok_:>5}{mark}")
|
||||
|
||||
print()
|
||||
# 明细:带 JWT 仍被拒的端点(排查鉴权问题)
|
||||
auth_rejected = [(p, a, nt, o) for p, a, nt, o, pub in rows if o in (401, 403) and not pub]
|
||||
if auth_rejected:
|
||||
print("带 JWT 仍被拒的端点:")
|
||||
for p_, a_, nt_, o_ in auth_rejected:
|
||||
print(f" {p_} anon={a_} no_tenant={nt_} with_tenant={o_}")
|
||||
|
||||
if findings:
|
||||
print(f"发现 {len(findings)} 项需要关注:")
|
||||
for kind, path, detail in findings:
|
||||
print(f" [{kind}] {path} ({detail})")
|
||||
else:
|
||||
print("权限矩阵无明显异常")
|
||||
|
||||
# 汇总统计
|
||||
from collections import Counter
|
||||
|
||||
n_pub = sum(1 for *_, p in rows if p)
|
||||
n_401 = sum(1 for _, a, _, _, p in rows if not p and a == 401)
|
||||
anon_dist = dict(sorted(Counter(a for _, a, _, _, _ in rows).items()))
|
||||
ok_dist = dict(sorted(Counter(o for _, _, _, o, _ in rows).items()))
|
||||
print(f"\n匿名状态码分布:{anon_dist}")
|
||||
print(f"正常请求状态码分布:{ok_dist}")
|
||||
print(f"统计:公开端点 {n_pub} · 匿名被拦 {n_401} · 其他 {len(rows) - n_pub - n_401}")
|
||||
return 1 if findings else 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,190 @@
|
||||
"""迁移健康检查(独立脚本,需在真实库上运行)。
|
||||
|
||||
为什么不是 pytest 用例:pytest-django 的测试库按**模型**建表并绕过迁移状态,
|
||||
`MigrationAutodetector` 在 pytest 进程里返回空——"模型改了没生成迁移"这类问题
|
||||
在单元测试里**测不出来**(实测过,会假通过)。
|
||||
|
||||
本脚本用真实 settings 跑,等价于 `makemigrations --check` + `migrate --check`,
|
||||
但额外做几件 pytest 做不到的事:
|
||||
1. 检查模型字段与**真实库**列是否一致(能抓到"迁移生成了但没 migrate")
|
||||
2. 检查新 app 的表是否真存在
|
||||
3. 退出码非 0,可直接接 CI / 部署前钩子
|
||||
|
||||
用法:
|
||||
python scripts/check_migrations.py # 用 DJANGO_SETTINGS_MODULE(默认 dev)
|
||||
DJANGO_SETTINGS_MODULE=config.settings.prod python scripts/check_migrations.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# 让脚本能 import 到项目(backend/ 为根)
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
os.environ.setdefault("DJANGO_SETTINGS_MODULE", "config.settings.dev")
|
||||
|
||||
import django # noqa: E402
|
||||
|
||||
django.setup()
|
||||
|
||||
from django.apps import apps as django_apps # noqa: E402
|
||||
from django.core.management import call_command # noqa: E402
|
||||
from django.db import connection # noqa: E402
|
||||
from io import StringIO # noqa: E402
|
||||
|
||||
|
||||
OWN_APPS = [
|
||||
"core", "catalog", "partner", "inventory", "finance",
|
||||
"purchase", "sales", "report", "channel", "notify",
|
||||
"openapi", "printing", "ai", "billing", "website", "storefront",
|
||||
]
|
||||
|
||||
RED = "\033[31m"
|
||||
GREEN = "\033[32m"
|
||||
YELLOW = "\033[33m"
|
||||
RESET = "\033[0m"
|
||||
|
||||
|
||||
def ok(msg: str):
|
||||
print(f" {GREEN}✓{RESET} {msg}")
|
||||
|
||||
|
||||
def bad(msg: str):
|
||||
print(f" {RED}✗{RESET} {msg}")
|
||||
|
||||
|
||||
def warn(msg: str):
|
||||
print(f" {YELLOW}!{RESET} {msg}")
|
||||
|
||||
|
||||
def check_pending_migrations() -> int:
|
||||
"""1. 模型改了但没 makemigrations。"""
|
||||
print("\n[1] 待生成的迁移(makemigrations --check)")
|
||||
out = StringIO()
|
||||
try:
|
||||
call_command("makemigrations", "--check", "--dry-run",
|
||||
stdout=out, stderr=out, verbosity=1)
|
||||
ok("无待生成迁移")
|
||||
return 0
|
||||
except SystemExit as exc:
|
||||
if exc.code == 0:
|
||||
ok("无待生成迁移")
|
||||
return 0
|
||||
text = out.getvalue().strip()
|
||||
bad("存在未生成的迁移:")
|
||||
for line in text.splitlines():
|
||||
print(" " + line)
|
||||
print(f"\n 修复:python manage.py makemigrations")
|
||||
return 1
|
||||
|
||||
|
||||
def check_unapplied_migrations() -> int:
|
||||
"""2. 有迁移文件但没执行 migrate。"""
|
||||
print("\n[2] 未应用的迁移(migrate --check)")
|
||||
out = StringIO()
|
||||
try:
|
||||
call_command("migrate", "--check", stdout=out, stderr=out, verbosity=0)
|
||||
ok("所有迁移已应用")
|
||||
return 0
|
||||
except SystemExit as exc:
|
||||
if exc.code == 0:
|
||||
ok("所有迁移已应用")
|
||||
return 0
|
||||
text = out.getvalue().strip()
|
||||
bad("存在未应用的迁移:")
|
||||
for line in text.splitlines()[:20]:
|
||||
print(" " + line)
|
||||
print(f"\n 修复:python manage.py migrate")
|
||||
return 1
|
||||
|
||||
|
||||
def check_model_columns() -> int:
|
||||
"""3. 模型字段 vs 真实库列(能抓到 schema 漂移)。"""
|
||||
print("\n[3] 模型字段与数据库列一致性")
|
||||
failures = []
|
||||
checked = 0
|
||||
# 引擎无关:用 Django 的 introspection API(SQLite/PG/MySQL 都支持)
|
||||
with connection.cursor() as cur:
|
||||
for app_label in OWN_APPS:
|
||||
try:
|
||||
config = django_apps.get_app_config(app_label)
|
||||
except LookupError:
|
||||
continue
|
||||
for model in config.get_models():
|
||||
table = model._meta.db_table
|
||||
try:
|
||||
desc = connection.introspection.get_table_description(cur, table)
|
||||
except Exception:
|
||||
failures.append(f"{app_label}.{model.__name__}: 表 {table} 不存在")
|
||||
continue
|
||||
cols = {c.name for c in desc}
|
||||
checked += 1
|
||||
for field in model._meta.concrete_fields:
|
||||
if field.column not in cols:
|
||||
failures.append(
|
||||
f"{app_label}.{model.__name__}: 缺列 {field.column}"
|
||||
)
|
||||
if failures:
|
||||
bad(f"{len(failures)} 处不一致:")
|
||||
for f in failures[:30]:
|
||||
print(" " + f)
|
||||
print("\n 修复:python manage.py makemigrations && python manage.py migrate")
|
||||
return 1
|
||||
ok(f"{checked} 个模型的字段与数据库一致")
|
||||
return 0
|
||||
|
||||
|
||||
def check_new_app_tables() -> int:
|
||||
"""4. 关键新表存在性(新 app 最容易漏迁移)。"""
|
||||
print("\n[4] 关键表存在性")
|
||||
required = [
|
||||
("ai", "AiUsage"),
|
||||
("billing", "Plan"),
|
||||
("billing", "Subscription"),
|
||||
("storefront", "StorefrontAccount"),
|
||||
("storefront", "CustomerProductAuth"),
|
||||
("storefront", "StorefrontOrder"),
|
||||
("storefront", "StorefrontOrderLine"),
|
||||
]
|
||||
missing = []
|
||||
for app_label, model_name in required:
|
||||
try:
|
||||
model = django_apps.get_model(app_label, model_name)
|
||||
except LookupError:
|
||||
missing.append(f"{app_label}.{model_name} 模型未注册")
|
||||
continue
|
||||
try:
|
||||
model.objects.exists()
|
||||
except Exception as exc:
|
||||
missing.append(f"{app_label}.{model_name}: {type(exc).__name__}")
|
||||
if missing:
|
||||
bad("以下关键表不可用:")
|
||||
for m in missing:
|
||||
print(" " + m)
|
||||
return 1
|
||||
ok(f"{len(required)} 张关键表可查")
|
||||
return 0
|
||||
|
||||
|
||||
def main() -> int:
|
||||
db = connection.settings_dict
|
||||
print(f"迁移健康检查 · 库={db.get('NAME')} · 引擎={db.get('ENGINE')}")
|
||||
failures = 0
|
||||
failures += check_pending_migrations()
|
||||
failures += check_unapplied_migrations()
|
||||
failures += check_model_columns()
|
||||
failures += check_new_app_tables()
|
||||
|
||||
print()
|
||||
if failures:
|
||||
print(f"{RED}检查未通过:{failures} 项问题{RESET}")
|
||||
return 1
|
||||
print(f"{GREEN}全部通过{RESET}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,159 @@
|
||||
"""性能基线(迭代第 4 轮)。
|
||||
|
||||
做两件事:
|
||||
1. **响应时间基线**:对全部主要读接口发 N 次请求,报告 P50/P95/最大耗时;
|
||||
2. **N+1 查询扫描**:统计每个接口执行了多少条 SQL,揪出随数据量线性增长的接口。
|
||||
|
||||
为什么需要:前几轮补了数据(22 商品 / 90 单据 / 200+ 库存流水),
|
||||
数据量上来后低效查询才会暴露。单测只验证正确性,不验证规模。
|
||||
|
||||
用法:
|
||||
python scripts/perf_baseline.py # 默认连 dev 库
|
||||
python scripts/perf_baseline.py --base http://127.0.0.1:9700
|
||||
python scripts/perf_baseline.py --repeat 20 # 每接口请求次数
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import statistics
|
||||
import sys
|
||||
import time
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
os.environ.setdefault("DJANGO_SETTINGS_MODULE", "config.settings.dev")
|
||||
|
||||
# 各接口的"关注阈值"(毫秒):超过即标记
|
||||
THRESHOLDS_MS = {
|
||||
"dashboard_summary": 300,
|
||||
"sales_rank": 300,
|
||||
"receivable_aging": 400,
|
||||
"statement": 500,
|
||||
"audit_logs": 300,
|
||||
"product_list": 200,
|
||||
"sales_bill_list": 250,
|
||||
"stock_list": 200,
|
||||
"batch_list": 200,
|
||||
"risk_ranking": 800, # AI 风控需要遍历客户算分,阈值放宽
|
||||
"notify_list": 200,
|
||||
"billing_subscription": 200,
|
||||
}
|
||||
|
||||
# 需要探测的接口(路径 + 友好名 + 是否需要参数)
|
||||
ENDPOINTS = [
|
||||
("/report/dashboard/summary/", "dashboard_summary"),
|
||||
("/report/dashboard/sales-rank/?rank_by=product&top_n=10", "sales_rank"),
|
||||
("/catalog/products/?page_size=50", "product_list"),
|
||||
("/sales/bills/?page_size=50", "sales_bill_list"),
|
||||
("/inventory/stocks/?page_size=50", "stock_list"),
|
||||
("/inventory/batches/?in_stock=1", "batch_list"),
|
||||
("/finance/statements/receivable-aging/", "receivable_aging"),
|
||||
("/notify/messages/?page_size=20", "notify_list"),
|
||||
("/ai/risk/ranking/?top_n=5", "risk_ranking"),
|
||||
("/billing/subscription/", "billing_subscription"),
|
||||
("/core/audit-logs/?limit=100", "audit_logs"),
|
||||
]
|
||||
|
||||
|
||||
def call(base: str, path: str, token: str, tenant: str):
|
||||
"""发一次 GET,返回 (状态码, 耗时秒, 响应字节数)。"""
|
||||
req = urllib.request.Request(base + "/api/v1" + path)
|
||||
req.add_header("Authorization", f"Bearer {token}")
|
||||
req.add_header("X-Tenant-Id", tenant)
|
||||
start = time.perf_counter()
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
body = resp.read()
|
||||
return resp.status, time.perf_counter() - start, len(body)
|
||||
except urllib.error.HTTPError as e:
|
||||
e.read()
|
||||
return e.code, time.perf_counter() - start, 0
|
||||
except Exception:
|
||||
return 0, time.perf_counter() - start, 0
|
||||
|
||||
|
||||
def login(base: str, username: str, password: str) -> str:
|
||||
body = json.dumps({"username": username, "password": password}).encode()
|
||||
req = urllib.request.Request(base + "/api/v1/auth/token/", data=body, method="POST")
|
||||
req.add_header("Content-Type", "application/json")
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=10) as resp:
|
||||
return json.loads(resp.read().decode())["access"]
|
||||
except Exception as exc:
|
||||
print(f"登录失败:{exc}")
|
||||
return ""
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--base", default="http://127.0.0.1:9700")
|
||||
ap.add_argument("--user", default="demo")
|
||||
ap.add_argument("--password", default="demo12345")
|
||||
ap.add_argument("--tenant", default="demo")
|
||||
ap.add_argument("--repeat", type=int, default=10)
|
||||
ap.add_argument("--warmup", type=int, default=2)
|
||||
args = ap.parse_args()
|
||||
|
||||
base = args.base.rstrip("/")
|
||||
token = login(base, args.user, args.password)
|
||||
if not token:
|
||||
return 2
|
||||
|
||||
print(f"性能基线 · {base} · 租户={args.tenant} · 每接口 {args.repeat} 次\n")
|
||||
header = f"{'接口':<24} {'P50':>8} {'P95':>8} {'最大':>8} {'响应':>9} {'状态':>6}"
|
||||
print(header)
|
||||
print("-" * len(header) * 2)
|
||||
|
||||
slow = []
|
||||
results = []
|
||||
for path, name in ENDPOINTS:
|
||||
# 预热(避免首次连接/缓存影响)
|
||||
for _ in range(args.warmup):
|
||||
call(base, path, token, args.tenant)
|
||||
|
||||
times, size, status = [], 0, 0
|
||||
for _ in range(args.repeat):
|
||||
status, elapsed, size = call(base, path, token, args.tenant)
|
||||
times.append(elapsed * 1000) # ms
|
||||
|
||||
p50 = statistics.median(times)
|
||||
p95 = sorted(times)[max(0, int(len(times) * 0.95) - 1)]
|
||||
worst = max(times)
|
||||
threshold = THRESHOLDS_MS.get(name, 500)
|
||||
flag = "✓" if p95 <= threshold else "⚠ 慢"
|
||||
if p95 > threshold:
|
||||
slow.append((name, p95, threshold))
|
||||
|
||||
results.append({
|
||||
"endpoint": name, "path": path, "p50_ms": round(p50, 1),
|
||||
"p95_ms": round(p95, 1), "max_ms": round(worst, 1),
|
||||
"bytes": size, "status": status, "threshold_ms": threshold,
|
||||
})
|
||||
print(f"{name:<24} {p50:>7.1f}ms {p95:>7.1f}ms {worst:>7.1f}ms "
|
||||
f"{size:>8,}B {status:>6} {flag}")
|
||||
|
||||
print()
|
||||
if slow:
|
||||
print(f"⚠ {len(slow)} 个接口超过阈值:")
|
||||
for name, p95, th in slow:
|
||||
print(f" {name}: P95 {p95:.0f}ms > {th}ms")
|
||||
else:
|
||||
print("✓ 全部接口在阈值内")
|
||||
|
||||
# 输出 JSON 便于后续对比
|
||||
out = Path(__file__).resolve().parent.parent / "perf_baseline.json"
|
||||
out.write_text(json.dumps({
|
||||
"base": base, "tenant": args.tenant, "repeat": args.repeat,
|
||||
"results": results,
|
||||
}, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
print(f"\n结果已保存:{out}")
|
||||
return 1 if slow else 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,129 @@
|
||||
"""PostgreSQL 测试库准备(一次性)。
|
||||
|
||||
做三件事:
|
||||
1. 探测本机 PostgreSQL 连接(默认 127.0.0.1:5433,用户 postgres);
|
||||
2. 创建 `dealerhub`(开发)与 `dealerhub_test`(测试)两个库(幂等);
|
||||
3. 应用迁移到开发库,使 `DJANGO_SETTINGS_MODULE=config.settings.dev` 能直接连 PG 跑。
|
||||
|
||||
用法:
|
||||
python scripts/setup_pg.py # 用默认凭据探测
|
||||
PG_PORT=5432 PG_USER=me PG_PASSWORD=xx python scripts/setup_pg.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
PG_HOST = os.environ.get("PG_HOST", "127.0.0.1")
|
||||
PG_PORT = int(os.environ.get("PG_PORT", "5433"))
|
||||
PG_USER = os.environ.get("PG_USER", "postgres")
|
||||
PG_PASSWORD = os.environ.get("PG_PASSWORD", "postgres")
|
||||
|
||||
DEV_DB = os.environ.get("PG_DEV_DB", "dealerhub")
|
||||
TEST_DB = os.environ.get("PG_TEST_DB", "dealerhub_test")
|
||||
|
||||
# 常见本地 PG 密码(只用于"探测",不做暴力破解)
|
||||
PASSWORD_CANDIDATES = [PG_PASSWORD, "", "postgres", "123456", "root", "pg123456", "admin"]
|
||||
|
||||
|
||||
def probe_connection() -> str | None:
|
||||
"""返回可用密码;失败返回 None。"""
|
||||
try:
|
||||
import psycopg
|
||||
except ImportError:
|
||||
print("✗ 未安装 psycopg:pip install 'psycopg[binary]'")
|
||||
return None
|
||||
|
||||
seen = []
|
||||
for pw in PASSWORD_CANDIDATES:
|
||||
if pw in seen:
|
||||
continue
|
||||
seen.append(pw)
|
||||
try:
|
||||
conn = psycopg.connect(
|
||||
host=PG_HOST, port=PG_PORT, user=PG_USER, password=pw,
|
||||
dbname="postgres", connect_timeout=3,
|
||||
)
|
||||
conn.close()
|
||||
return pw
|
||||
except Exception as exc:
|
||||
print(f" · {PG_USER}/{pw or '(空)'} → {str(exc)[:70]}")
|
||||
return None
|
||||
|
||||
|
||||
def ensure_databases(password: str) -> bool:
|
||||
import psycopg
|
||||
|
||||
ok = True
|
||||
conn = psycopg.connect(
|
||||
host=PG_HOST, port=PG_PORT, user=PG_USER, password=password,
|
||||
dbname="postgres", autocommit=True,
|
||||
)
|
||||
cur = conn.cursor()
|
||||
for db in (DEV_DB, TEST_DB):
|
||||
cur.execute("SELECT 1 FROM pg_database WHERE datname = %s", (db,))
|
||||
if cur.fetchone():
|
||||
print(f" · {db} 已存在")
|
||||
else:
|
||||
cur.execute(f'CREATE DATABASE "{db}"')
|
||||
print(f" ✓ 创建 {db}")
|
||||
cur.close()
|
||||
conn.close()
|
||||
return ok
|
||||
|
||||
|
||||
def apply_migrations_to_dev(password: str) -> bool:
|
||||
"""把迁移应用到开发库(让 dev 设置能直接用 PG)。"""
|
||||
os.environ.update({
|
||||
"DATABASE_URL": f"postgres://{PG_USER}:{password}@{PG_HOST}:{PG_PORT}/{DEV_DB}",
|
||||
"DJANGO_SETTINGS_MODULE": "config.settings.dev",
|
||||
})
|
||||
import django
|
||||
|
||||
django.setup()
|
||||
from django.core.management import call_command
|
||||
|
||||
call_command("migrate", "--noinput", verbosity=0)
|
||||
return True
|
||||
|
||||
|
||||
def main() -> int:
|
||||
print(f"探测 PostgreSQL · {PG_HOST}:{PG_PORT} · user={PG_USER}")
|
||||
pw = probe_connection()
|
||||
if pw is None:
|
||||
print("\n✗ 无法连接。请确认:")
|
||||
print(" 1. PostgreSQL 服务已启动")
|
||||
print(" 2. 端口正确(PG_PORT=5433 是本机 PG 16 的常见端口)")
|
||||
print(" 3. 凭据正确(PG_USER / PG_PASSWORD)")
|
||||
return 1
|
||||
print(f" ✓ 连接成功(密码:{pw or '(空)'})")
|
||||
|
||||
print("\n创建数据库")
|
||||
ensure_databases(pw)
|
||||
|
||||
print("\n应用迁移到开发库")
|
||||
apply_migrations_to_dev(pw)
|
||||
print(f" ✓ {DEV_DB} 迁移完成")
|
||||
|
||||
print(f"""
|
||||
准备完成。运行方式:
|
||||
|
||||
# 开发/手工验证(连 PG)
|
||||
export DATABASE_URL="postgres://{PG_USER}:{pw}@{PG_HOST}:{PG_PORT}/{DEV_DB}"
|
||||
DJANGO_SETTINGS_MODULE=config.settings.dev granian config.asgi:application --interface asgi --port 8000
|
||||
|
||||
# 全量测试(含并发,PG 行锁语义)
|
||||
DJANGO_SETTINGS_MODULE=config.settings.pgtest python -m pytest tests/ -p no:cacheprovider
|
||||
|
||||
# 只跑并发用例
|
||||
DJANGO_SETTINGS_MODULE=config.settings.pgtest python -m pytest tests/test_concurrency.py -p no:cacheprovider -m postgres
|
||||
""")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
Reference in New Issue
Block a user