188 lines
7.2 KiB
Python
188 lines
7.2 KiB
Python
"""网络配图 e2e:image_strategy=web 全链路(test 通道)+ 空图源优雅降级。
|
||
|
||
用法:python tests/e2e_image.py
|
||
环境:IMAGE_SEARCH_MODE=test + IMAGE_SEARCH_TEST_DIR=<样例图目录>
|
||
"""
|
||
|
||
import json
|
||
import os
|
||
import subprocess
|
||
import sys
|
||
import time
|
||
import urllib.request
|
||
import urllib.error
|
||
from pathlib import Path
|
||
|
||
ROOT = Path(__file__).resolve().parents[1]
|
||
sys.path.insert(0, str(ROOT))
|
||
|
||
BASE = "http://127.0.0.1:5102"
|
||
USER = f"img_e2e_{int(time.time())}"
|
||
PASS = "img_e2e_123"
|
||
SAMPLE = ("# 配图测试主题\n本季度完成检索服务上线\n准确率达到百分之九十八\n"
|
||
"# 运营情况\n用户增长明显\n文档数量翻倍\n")
|
||
|
||
TMP = ROOT / "tests" / "_up"
|
||
|
||
|
||
def make_test_images():
|
||
from PIL import Image
|
||
TMP.mkdir(exist_ok=True)
|
||
for i, color in enumerate([(21, 101, 192), (46, 125, 50), (123, 31, 162)]):
|
||
Image.new("RGB", (800, 500), color).save(str(TMP / f"sample_{i}.jpg"))
|
||
|
||
|
||
COOKIE = {"value": ""}
|
||
|
||
|
||
def call(method, path, payload=None):
|
||
data = None if payload is None else json.dumps(payload).encode("utf-8")
|
||
req = urllib.request.Request(BASE + path, data=data, method=method)
|
||
if payload is not None:
|
||
req.add_header("Content-Type", "application/json")
|
||
if COOKIE["value"]:
|
||
req.add_header("Cookie", COOKIE["value"])
|
||
try:
|
||
with urllib.request.urlopen(req, timeout=180) as resp:
|
||
sc = resp.headers.get("Set-Cookie")
|
||
if sc and not COOKIE["value"]:
|
||
COOKIE["value"] = sc.split(";")[0]
|
||
return resp.status, json.loads(resp.read() or b"{}")
|
||
except urllib.error.HTTPError as e:
|
||
try:
|
||
return e.code, json.loads(e.read() or b"{}")
|
||
except Exception:
|
||
return e.code, {}
|
||
|
||
|
||
def generate_and_wait(title):
|
||
s, b = call("POST", "/api/generate", {
|
||
"title": title, "content": SAMPLE,
|
||
"scene": "report", "language": "zh", "canvas": "ppt169",
|
||
"image_strategy": "web", "color_scheme": "blue",
|
||
"page_min": 5, "page_max": 7,
|
||
})
|
||
if b.get("code") != 0:
|
||
return None, str(b)[:100]
|
||
task_id = b["data"]["task_id"]
|
||
for _ in range(120):
|
||
time.sleep(1)
|
||
s, b = call("GET", f"/api/generate/{task_id}/status")
|
||
st = b.get("data", {}).get("status")
|
||
if st in ("done", "failed"):
|
||
return (task_id if st == "done" else None), json.dumps(b.get("data"), ensure_ascii=False)[:120]
|
||
return None, "timeout"
|
||
|
||
|
||
def main():
|
||
failures = []
|
||
|
||
def check(name, cond, extra=""):
|
||
print(f" [{'OK' if cond else 'NG'}] {name} {extra}")
|
||
if not cond:
|
||
failures.append(name)
|
||
|
||
make_test_images()
|
||
env = {**os.environ, "PYTHONIOENCODING": "utf-8",
|
||
"LLM_BACKEND": "off",
|
||
"IMAGE_SEARCH_MODE": "test",
|
||
"IMAGE_SEARCH_TEST_DIR": str(TMP),
|
||
"MAX_IMAGE_SLIDES": "3"}
|
||
subprocess.run([sys.executable, "-m", "src.db.init"], cwd=str(ROOT),
|
||
capture_output=True, env=env)
|
||
proc = subprocess.Popen(
|
||
[sys.executable, "-c",
|
||
f"import sys; sys.path.insert(0, r'{ROOT}');"
|
||
"from src.web.app import app; app.run(host='127.0.0.1', port=5102)"],
|
||
cwd=str(ROOT), env=env, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||
)
|
||
try:
|
||
up = False
|
||
for _ in range(30):
|
||
time.sleep(0.5)
|
||
try:
|
||
call("GET", "/api/me")
|
||
up = True
|
||
break
|
||
except Exception:
|
||
continue
|
||
check("server up", up)
|
||
|
||
s, b = call("POST", "/api/register",
|
||
{"username": USER, "password": PASS, "confirm_password": PASS})
|
||
check("register", s == 200 and b.get("code") == 0)
|
||
|
||
# ---- 场景1:配图成功 ----
|
||
task_id, info = generate_and_wait("配图全链路测试")
|
||
check("generation done (with images)", task_id is not None, info)
|
||
if task_id:
|
||
s, b = call("GET", f"/api/generate/{task_id}/result")
|
||
rec = b["data"]["record"]
|
||
project_dir = ROOT / "src" / "web" / "data" / "projects" / rec["project_dir"] \
|
||
if rec.get("project_dir") else None
|
||
# 兜底:从 versions 表拿不到目录名时按标题匹配最新目录
|
||
if not project_dir or not project_dir.exists():
|
||
dirs = sorted((ROOT / "src" / "web" / "data" / "projects").glob("配图全链路测试_*"))
|
||
project_dir = dirs[-1]
|
||
plan = json.loads((project_dir / "plan.json").read_text(encoding="utf-8"))
|
||
with_img = [sl for sl in plan["slides"] if sl.get("image")]
|
||
check("plan has image slides", len(with_img) >= 2, f"n={len(with_img)}")
|
||
check("image files exist", all(Path(sl["image"]).exists() for sl in with_img))
|
||
check("images inside project dir",
|
||
all(str(sl["image"]).startswith(str(project_dir)) for sl in with_img))
|
||
|
||
s, raw = call("GET", f"/api/files/{rec['id']}/download", raw=True) \
|
||
if False else (None, None)
|
||
# 直接读导出文件验证 PPTX 内嵌图片形状
|
||
pptx_path = project_dir / "exports" / rec["file_name"] \
|
||
if rec.get("file_name") else None
|
||
if not pptx_path or not pptx_path.exists():
|
||
files = sorted((project_dir / "exports").glob("*.pptx"))
|
||
pptx_path = files[-1]
|
||
from pptx import Presentation
|
||
prs = Presentation(str(pptx_path))
|
||
pic_slides = [i for i, sl in enumerate(prs.slides)
|
||
if any(sh.shape_type == 13 for sh in sl.shapes)]
|
||
check("pptx contains pictures", len(pic_slides) >= 2, f"slides={pic_slides}")
|
||
# 保真验证已在生成流程内通过(done 状态),此处确认无 image 键误报的副作用
|
||
check("no cover/end images", all(
|
||
sl.get("type") == "content" for sl in with_img))
|
||
|
||
# ---- 场景2:空图源 → 优雅降级 ----
|
||
empty_dir = TMP / "empty_src"
|
||
empty_dir.mkdir(exist_ok=True)
|
||
proc.terminate()
|
||
proc.wait()
|
||
env2 = {**env, "IMAGE_SEARCH_TEST_DIR": str(empty_dir)}
|
||
proc = subprocess.Popen(
|
||
[sys.executable, "-c",
|
||
f"import sys; sys.path.insert(0, r'{ROOT}');"
|
||
"from src.web.app import app; app.run(host='127.0.0.1', port=5102)"],
|
||
cwd=str(ROOT), env=env2, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||
)
|
||
for _ in range(30):
|
||
time.sleep(0.5)
|
||
try:
|
||
call("GET", "/api/me")
|
||
break
|
||
except Exception:
|
||
continue
|
||
task_id, info = generate_and_wait("空图源降级测试")
|
||
check("generation done (fallback no-image)", task_id is not None, info)
|
||
if task_id:
|
||
s, b = call("GET", f"/api/generate/{task_id}/result")
|
||
rec = b["data"]["record"]
|
||
check("record still created", bool(rec and rec["id"]), str(rec)[:60])
|
||
finally:
|
||
proc.terminate()
|
||
|
||
import shutil
|
||
shutil.rmtree(TMP, ignore_errors=True)
|
||
print("\n" + ("IMAGE E2E ALL PASSED" if not failures
|
||
else "FAILURES: " + "; ".join(failures)))
|
||
sys.exit(1 if failures else 0)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|