188 lines
7.1 KiB
Python
188 lines
7.1 KiB
Python
"""构建并编译 LangGraph,提供 run_country() 入口。
|
||
|
||
图结构(线性流水线,节点全部带兜底):
|
||
START -> seed -> fetch -> filter -> score -> screen -> prompt_build
|
||
-> compose(生成纯印花设计稿 + 导出简报)-> product(底图/模特/三图合成/模板)
|
||
-> oss_upload(压缩 3:4 / ≥1340×1785 / <2MB + 上传阿里云 OSS)-> END
|
||
"""
|
||
import time
|
||
from pathlib import Path
|
||
from typing import Any, Dict, Optional
|
||
|
||
from langgraph.graph import END, StateGraph
|
||
|
||
from graph.loader import build_country_config
|
||
from graph.nodes import (
|
||
compose_node,
|
||
fetch_node,
|
||
filter_node,
|
||
oss_upload_node,
|
||
product_node,
|
||
prompt_node,
|
||
score_node,
|
||
screen_node,
|
||
seed_node,
|
||
seed_shot_node,
|
||
template_export_node,
|
||
)
|
||
from graph.state import AgentState
|
||
|
||
|
||
def build_graph():
|
||
"""构建 StateGraph 并编译。"""
|
||
builder = StateGraph(AgentState)
|
||
builder.add_node("seed", seed_node)
|
||
builder.add_node("fetch", fetch_node)
|
||
builder.add_node("filter", filter_node)
|
||
builder.add_node("score", score_node)
|
||
builder.add_node("screen", screen_node)
|
||
builder.add_node("prompt_build", prompt_node)
|
||
builder.add_node("product", product_node)
|
||
builder.add_node("compose", compose_node)
|
||
builder.add_node("oss_upload", oss_upload_node)
|
||
builder.add_node("seed_shot", seed_shot_node)
|
||
builder.add_node("template_export", template_export_node)
|
||
|
||
builder.add_edge("__start__", "seed")
|
||
builder.add_edge("seed", "fetch")
|
||
builder.add_edge("fetch", "filter")
|
||
builder.add_edge("filter", "score")
|
||
builder.add_edge("score", "screen")
|
||
builder.add_edge("screen", "prompt_build")
|
||
builder.add_edge("prompt_build", "compose") # compose:生成纯印花设计稿(放前面)
|
||
builder.add_edge("compose", "product") # product:底图/模特/三图合成/模板
|
||
builder.add_edge("product", "oss_upload") # oss_upload:压缩 + 上传图床
|
||
builder.add_edge("oss_upload", "seed_shot") # seed_shot:种草图生成(模板+模特特征 yaml)→ 上传
|
||
builder.add_edge("seed_shot", "template_export") # template_export:最终结果导入商品上传模板
|
||
builder.add_edge("template_export", END)
|
||
return builder.compile()
|
||
|
||
|
||
def build_pinterest_graph():
|
||
"""Pinterest 参考模式图(独立于 Google Trends 采集链路):
|
||
pinterest_search → pinterest_scrape → pinterest_analyze → compose → product
|
||
→ oss_upload → seed_shot → template_export
|
||
"""
|
||
from graph.nodes import (
|
||
pinterest_analyze_node,
|
||
pinterest_scrape_node,
|
||
pinterest_search_node,
|
||
)
|
||
|
||
builder = StateGraph(AgentState)
|
||
builder.add_node("pinterest_search", pinterest_search_node)
|
||
builder.add_node("pinterest_scrape", pinterest_scrape_node)
|
||
builder.add_node("pinterest_analyze", pinterest_analyze_node)
|
||
builder.add_node("compose", compose_node)
|
||
builder.add_node("product", product_node)
|
||
builder.add_node("oss_upload", oss_upload_node)
|
||
builder.add_node("seed_shot", seed_shot_node)
|
||
builder.add_node("template_export", template_export_node)
|
||
|
||
builder.add_edge("__start__", "pinterest_search")
|
||
builder.add_edge("pinterest_search", "pinterest_scrape")
|
||
builder.add_edge("pinterest_scrape", "pinterest_analyze")
|
||
builder.add_edge("pinterest_analyze", "compose")
|
||
builder.add_edge("compose", "product")
|
||
builder.add_edge("product", "oss_upload")
|
||
builder.add_edge("oss_upload", "seed_shot")
|
||
builder.add_edge("seed_shot", "template_export")
|
||
builder.add_edge("template_export", END)
|
||
return builder.compile()
|
||
|
||
|
||
def run_country(
|
||
country: str,
|
||
global_config: Dict[str, Any],
|
||
project_root: Path,
|
||
output_root: Optional[Path] = None,
|
||
base_image: Optional[str] = None,
|
||
task_timestamp: Optional[str] = None,
|
||
) -> Dict[str, Any]:
|
||
"""运行单个国家的完整流水线,返回最终 state(含 errors / stats / briefs)。
|
||
|
||
project_root:数据文件根(configs / prompts,打包后为 _MEIPASS 只读目录)。
|
||
output_root :产物输出根(默认=project_root;打包后传 exe 旁运行目录,
|
||
避免把 output/ 写进临时解压目录导致重启丢失)。
|
||
task_timestamp:任务时间戳(每次点击运行 = 一个任务);None 时自动生成。
|
||
"""
|
||
compiled = build_graph()
|
||
cc = build_country_config(global_config, country, project_root)
|
||
prompts_dir = project_root / "prompts" / country
|
||
cache_dir = (output_root or project_root) / "output" / country # 缓存/去重(根目录)
|
||
ts = task_timestamp or time.strftime("%Y%m%d_%H%M%S")
|
||
_base = ts
|
||
_i = 1
|
||
while (cache_dir / ts).exists(): # 时间戳文件夹唯一(同秒多任务防冲突/覆盖)
|
||
ts = f"{_base}_{_i}"
|
||
_i += 1
|
||
output_dir = cache_dir / ts # 本次任务产物(时间戳文件夹)
|
||
|
||
state: Dict[str, Any] = {
|
||
"country": country,
|
||
"config": global_config,
|
||
"country_config": cc,
|
||
"prompts_dir": str(prompts_dir),
|
||
"cache_dir": str(cache_dir),
|
||
"output_dir": str(output_dir),
|
||
"raw_rows": [],
|
||
"filtered_rows": [],
|
||
"scored_rows": [],
|
||
"screened": [],
|
||
"briefs": [],
|
||
"composite": [],
|
||
"designs": [],
|
||
"errors": [],
|
||
"stats": {},
|
||
"task_timestamp": ts, # 任务开始时间戳(OSS 路径段 / 产物文件夹名)
|
||
"oss_seq": 0, # 货号计数(000 起,最多 999)
|
||
}
|
||
if base_image:
|
||
state["base_image"] = base_image
|
||
|
||
result = compiled.invoke(state)
|
||
return result
|
||
|
||
|
||
def run_pinterest_ref(
|
||
country: str,
|
||
global_config: Dict[str, Any],
|
||
project_root: Path,
|
||
output_root: Optional[Path] = None,
|
||
task_timestamp: Optional[str] = None,
|
||
) -> Dict[str, Any]:
|
||
"""Pinterest 参考模式入口:独立于 Google Trends 的完整流程。
|
||
|
||
种子词 → LLM 搜索词(json_schema + 动态注入防重复)→ 爬图 → LLM 分析图片
|
||
→ 设计简报 → 设计稿 → 产品图 → 上传 → 种草图 → 模板导出。
|
||
参数语义与 run_country 一致(project_root=数据根,output_root=产物根)。
|
||
"""
|
||
compiled = build_pinterest_graph()
|
||
cc = build_country_config(global_config, country, project_root)
|
||
prompts_dir = project_root / "prompts" / country
|
||
cache_dir = (output_root or project_root) / "output" / country
|
||
ts = task_timestamp or time.strftime("%Y%m%d_%H%M%S")
|
||
_base = ts
|
||
_i = 1
|
||
while (cache_dir / ts).exists(): # 时间戳文件夹唯一(同秒多任务防冲突/覆盖)
|
||
ts = f"{_base}_{_i}"
|
||
_i += 1
|
||
output_dir = cache_dir / ts
|
||
|
||
state: Dict[str, Any] = {
|
||
"country": country,
|
||
"config": global_config,
|
||
"country_config": cc,
|
||
"prompts_dir": str(prompts_dir),
|
||
"cache_dir": str(cache_dir),
|
||
"output_dir": str(output_dir),
|
||
"briefs": [],
|
||
"composite": [],
|
||
"designs": [],
|
||
"errors": [],
|
||
"stats": {},
|
||
"task_timestamp": ts,
|
||
"oss_seq": 0,
|
||
}
|
||
return compiled.invoke(state)
|