Files

188 lines
7.2 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""网络配图 e2eimage_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()