462 lines
16 KiB
Python
462 lines
16 KiB
Python
from __future__ import annotations
|
||
|
||
import json
|
||
import os
|
||
import re
|
||
import time
|
||
from datetime import datetime
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
import requests
|
||
|
||
from gpt_image2_generate import (
|
||
DEFAULT_CONNECT_TIMEOUT,
|
||
DEFAULT_INITIAL_POLL_DELAY,
|
||
DEFAULT_MAX_POLL_WAIT,
|
||
DEFAULT_MODEL,
|
||
DEFAULT_PROMPT,
|
||
DEFAULT_PIXEL_SIZE,
|
||
DEFAULT_QUALITY,
|
||
DEFAULT_READ_TIMEOUT,
|
||
DEFAULT_POLL_CONNECT_TIMEOUT,
|
||
DEFAULT_POLL_INTERVAL,
|
||
DEFAULT_POLL_READ_TIMEOUT,
|
||
ImageInput,
|
||
build_request_session,
|
||
dump_json,
|
||
env_int,
|
||
extract_image_bytes_from_item,
|
||
image_source_to_input,
|
||
is_pixel_size,
|
||
load_env_file,
|
||
make_run_dir,
|
||
normalize_base_url,
|
||
normalize_size,
|
||
payload_to_form_fields,
|
||
resolve_api_size,
|
||
resolve_min_output_edge,
|
||
resolve_request_style,
|
||
save_error_artifact,
|
||
save_http_response,
|
||
save_image_bytes,
|
||
)
|
||
|
||
|
||
DEFAULT_OUTPUT_STEM = "gpt-image-2-multi"
|
||
SUPPORTED_ASPECT_RATIOS = {
|
||
"1:1",
|
||
"16:9",
|
||
"9:16",
|
||
"4:3",
|
||
"3:4",
|
||
"3:2",
|
||
"2:3",
|
||
"5:4",
|
||
"4:5",
|
||
"2:1",
|
||
"1:2",
|
||
"21:9",
|
||
"9:21",
|
||
}
|
||
|
||
|
||
def _json_dumps(data: Any) -> str:
|
||
return json.dumps(data, ensure_ascii=False, indent=2)
|
||
|
||
|
||
def split_reference_values(raw_value: str) -> list[str]:
|
||
value = raw_value.strip()
|
||
if not value:
|
||
return []
|
||
|
||
if value.startswith("["):
|
||
parsed = json.loads(value)
|
||
if not isinstance(parsed, list):
|
||
raise RuntimeError("REFERENCE_IMAGES 必须是 JSON 数组")
|
||
return [str(item).strip() for item in parsed if str(item).strip()]
|
||
|
||
return [part.strip() for part in value.split("|") if part.strip()]
|
||
|
||
|
||
def collect_reference_sources() -> list[str]:
|
||
raw_list = os.getenv("REFERENCE_IMAGES", "").strip()
|
||
if raw_list:
|
||
refs = split_reference_values(raw_list)
|
||
else:
|
||
refs = []
|
||
for index in range(1, 17):
|
||
value = os.getenv(f"REFERENCE_IMAGE_{index}", "").strip()
|
||
if value:
|
||
refs.append(value)
|
||
if not refs:
|
||
single_candidates = [
|
||
os.getenv("REFERENCE_IMAGE", "").strip(),
|
||
os.getenv("REFERENCE_IMAGE_FILE", "").strip(),
|
||
os.getenv("REFERENCE_IMAGE_URL", "").strip(),
|
||
os.getenv("REFERENCE_IMAGE_DATA_URL", "").strip(),
|
||
]
|
||
refs.extend([value for value in single_candidates if value])
|
||
|
||
deduped: list[str] = []
|
||
seen: set[str] = set()
|
||
for ref in refs:
|
||
if ref not in seen:
|
||
deduped.append(ref)
|
||
seen.add(ref)
|
||
return deduped
|
||
|
||
|
||
def collect_mask_source() -> str | None:
|
||
candidates = {
|
||
"MASK_IMAGE": os.getenv("MASK_IMAGE", "").strip(),
|
||
"MASK_IMAGE_FILE": os.getenv("MASK_IMAGE_FILE", "").strip(),
|
||
"MASK_IMAGE_URL": os.getenv("MASK_IMAGE_URL", "").strip(),
|
||
"MASK_IMAGE_DATA_URL": os.getenv("MASK_IMAGE_DATA_URL", "").strip(),
|
||
}
|
||
provided = [value for value in candidates.values() if value]
|
||
if len(provided) > 1:
|
||
raise RuntimeError("MASK_IMAGE 相关变量只能设置一个")
|
||
return provided[0] if provided else None
|
||
|
||
|
||
def resolve_output_stem() -> str:
|
||
raw_value = os.getenv("IMAGE_OUTPUT_STEM", DEFAULT_OUTPUT_STEM).strip()
|
||
candidate = re.sub(r"[^A-Za-z0-9._-]+", "_", Path(raw_value).name).strip("._")
|
||
return candidate or DEFAULT_OUTPUT_STEM
|
||
|
||
|
||
def validate_requested_size(requested_size: str) -> None:
|
||
if not (is_pixel_size(requested_size) or requested_size in SUPPORTED_ASPECT_RATIOS):
|
||
raise RuntimeError(f"不支持的 IMAGE_SIZE: {requested_size}")
|
||
|
||
|
||
def build_payload(base_url: str, reference_sources: list[str]) -> tuple[dict[str, Any], str, str, str, str, str | None, str | None]:
|
||
model = os.getenv("IMAGE_MODEL", DEFAULT_MODEL).strip()
|
||
prompt = os.getenv("IMAGE_PROMPT", DEFAULT_PROMPT).strip() or DEFAULT_PROMPT
|
||
|
||
requested_size = normalize_size(os.getenv("IMAGE_SIZE", DEFAULT_PIXEL_SIZE))
|
||
validate_requested_size(requested_size)
|
||
request_style = resolve_request_style(base_url)
|
||
api_size = resolve_api_size(requested_size, request_style)
|
||
quality = os.getenv("IMAGE_QUALITY", DEFAULT_QUALITY).strip()
|
||
resolution = os.getenv("IMAGE_RESOLUTION", "").strip()
|
||
background = os.getenv("IMAGE_BACKGROUND", "").strip() or None
|
||
input_fidelity = os.getenv("IMAGE_INPUT_FIDELITY", "").strip() or None
|
||
output_format = os.getenv("IMAGE_OUTPUT_FORMAT", "").strip() or None
|
||
|
||
n = env_int("IMAGE_N", max(1, len(reference_sources)))
|
||
if n < 1:
|
||
raise RuntimeError("IMAGE_N 必须大于等于 1")
|
||
|
||
payload: dict[str, Any] = {
|
||
"model": model,
|
||
"prompt": prompt,
|
||
"n": n,
|
||
"size": api_size,
|
||
"quality": quality,
|
||
}
|
||
if not (request_style == "openai" and model.startswith("gpt-image")):
|
||
payload["response_format"] = "b64_json"
|
||
if not is_pixel_size(requested_size) and request_style != "openai":
|
||
payload["resolution"] = resolution or "4K"
|
||
if background:
|
||
payload["background"] = background
|
||
if input_fidelity:
|
||
payload["input_fidelity"] = input_fidelity
|
||
if output_format:
|
||
payload["output_format"] = output_format
|
||
|
||
mask_source = collect_mask_source()
|
||
if mask_source and not reference_sources:
|
||
raise RuntimeError("MASK_IMAGE 只能和参考图一起使用")
|
||
|
||
request_mode = "edit" if reference_sources else "generate"
|
||
return payload, model, requested_size, request_mode, mask_source, background, input_fidelity
|
||
|
||
|
||
def build_edit_files(
|
||
request_style: str,
|
||
reference_inputs: list[ImageInput],
|
||
mask_input: ImageInput | None,
|
||
) -> list[tuple[str, tuple[str, bytes, str]]]:
|
||
field_name = "image" if request_style == "openai" and len(reference_inputs) == 1 else "image[]"
|
||
files: list[tuple[str, tuple[str, bytes, str]]] = [
|
||
(field_name, (ref.filename, ref.data, ref.mime_type))
|
||
for ref in reference_inputs
|
||
]
|
||
if mask_input is not None:
|
||
files.append(("mask", (mask_input.filename, mask_input.data, mask_input.mime_type)))
|
||
return files
|
||
|
||
|
||
def extract_task_id(data: Any) -> str | None:
|
||
if not isinstance(data, dict):
|
||
return None
|
||
for key in ("task_id", "id"):
|
||
value = data.get(key)
|
||
if isinstance(value, str) and value:
|
||
return value
|
||
nested = data.get("data")
|
||
if isinstance(nested, dict):
|
||
for key in ("task_id", "id"):
|
||
value = nested.get(key)
|
||
if isinstance(value, str) and value:
|
||
return value
|
||
return None
|
||
|
||
|
||
def get_task_error(task: dict[str, Any]) -> str:
|
||
nested = task.get("data")
|
||
error = nested.get("error") if isinstance(nested, dict) else task.get("error")
|
||
if isinstance(error, dict):
|
||
return str(error.get("message") or error)
|
||
if error:
|
||
return str(error)
|
||
return "task failed"
|
||
|
||
|
||
def extract_image_items(data: Any) -> list[dict[str, Any]]:
|
||
if isinstance(data, list):
|
||
return [item for item in data if isinstance(item, dict)]
|
||
|
||
if isinstance(data, dict):
|
||
for key in ("images", "data"):
|
||
nested = data.get(key)
|
||
if isinstance(nested, list):
|
||
return [item for item in nested if isinstance(item, dict)]
|
||
|
||
result = data.get("result")
|
||
if isinstance(result, dict):
|
||
images = result.get("images")
|
||
if isinstance(images, list):
|
||
return [item for item in images if isinstance(item, dict)]
|
||
|
||
return []
|
||
|
||
|
||
def poll_task_images(
|
||
task_id: str,
|
||
base_url: str,
|
||
headers: dict[str, str],
|
||
session: requests.Session,
|
||
run_dir: Path,
|
||
max_wait_seconds: int,
|
||
initial_poll_delay: int,
|
||
poll_interval: int,
|
||
) -> list[dict[str, Any]]:
|
||
task_url = f"{base_url}/v1/tasks/{task_id}"
|
||
deadline = time.monotonic() + max_wait_seconds
|
||
poll_index = 0
|
||
|
||
time.sleep(initial_poll_delay)
|
||
while True:
|
||
if time.monotonic() >= deadline:
|
||
raise TimeoutError(f"任务轮询超时: task_id={task_id}")
|
||
|
||
poll_index += 1
|
||
task_resp = session.get(
|
||
task_url,
|
||
headers=headers,
|
||
timeout=(DEFAULT_POLL_CONNECT_TIMEOUT, DEFAULT_POLL_READ_TIMEOUT),
|
||
)
|
||
task_record = save_http_response(run_dir / f"poll-{poll_index:03d}.json", task_resp)
|
||
task_resp.raise_for_status()
|
||
task = task_record.get("json")
|
||
if not isinstance(task, dict):
|
||
raise RuntimeError(f"任务响应不是 JSON 对象: {task_record}")
|
||
|
||
data = task.get("data")
|
||
if not isinstance(data, dict):
|
||
raise RuntimeError(f"任务响应格式异常: {task}")
|
||
|
||
status = data.get("status")
|
||
print(f"任务轮询 #{poll_index}: status={status}")
|
||
if status == "completed":
|
||
result = data.get("result", {})
|
||
images = result.get("images", [])
|
||
if not images:
|
||
raise RuntimeError(f"任务已完成,但没有返回图片: {task}")
|
||
return [item for item in images if isinstance(item, dict)]
|
||
|
||
if status == "failed":
|
||
raise RuntimeError(get_task_error(task))
|
||
|
||
time.sleep(poll_interval)
|
||
|
||
|
||
def save_output_images(
|
||
items: list[dict[str, Any]],
|
||
session: requests.Session,
|
||
run_dir: Path,
|
||
stem: str,
|
||
min_output_edge: int | None,
|
||
) -> list[str]:
|
||
paths: list[str] = []
|
||
for index, item in enumerate(items, start=1):
|
||
image_bytes = extract_image_bytes_from_item(item, session)
|
||
output_path = run_dir / f"{stem}-{index}.png"
|
||
save_image_bytes(image_bytes, output_path, min_output_edge=min_output_edge)
|
||
paths.append(str(output_path))
|
||
return paths
|
||
|
||
|
||
def main() -> None:
|
||
load_env_file()
|
||
|
||
api_key = os.getenv("OPENAI_API_KEY") or os.getenv("API_KEY")
|
||
base_url = normalize_base_url(
|
||
os.getenv("OPENAI_BASE_URL") or os.getenv("BASE_URL") or "https://api.toskaxy.xyz"
|
||
)
|
||
if not api_key:
|
||
raise RuntimeError("未找到 API Key,请先在 .env 中设置 OPENAI_API_KEY")
|
||
|
||
reference_sources = collect_reference_sources()
|
||
payload, model, requested_size, request_mode, mask_source, background, input_fidelity = build_payload(
|
||
base_url,
|
||
reference_sources,
|
||
)
|
||
run_dir = make_run_dir()
|
||
session = build_request_session(base_url)
|
||
request_style = resolve_request_style(base_url)
|
||
request_timeout = env_int("IMAGE_REQUEST_TIMEOUT", DEFAULT_READ_TIMEOUT)
|
||
max_poll_wait = env_int("IMAGE_MAX_POLL_WAIT", DEFAULT_MAX_POLL_WAIT)
|
||
initial_poll_delay = env_int("IMAGE_INITIAL_POLL_DELAY", DEFAULT_INITIAL_POLL_DELAY)
|
||
poll_interval = env_int("IMAGE_POLL_INTERVAL", DEFAULT_POLL_INTERVAL)
|
||
endpoint_path = "/v1/images/edits" if request_mode == "edit" else "/v1/images/generations"
|
||
min_output_edge = resolve_min_output_edge(requested_size, payload.get("resolution"), request_style)
|
||
prompt = payload["prompt"]
|
||
stem = f"{resolve_output_stem()}-{datetime.now().strftime('%Y%m%d-%H%M%S-%f')}"
|
||
|
||
print(
|
||
"开始生成图片:",
|
||
f"mode={request_mode}",
|
||
f"request_style={request_style}",
|
||
f"model={model}",
|
||
f"refs={len(reference_sources)}",
|
||
f"n={payload['n']}",
|
||
f"size={requested_size}",
|
||
f"resolution={payload.get('resolution') or 'none'}",
|
||
f"request_timeout={request_timeout}s",
|
||
f"max_poll_wait={max_poll_wait}s",
|
||
f"initial_poll_delay={initial_poll_delay}s",
|
||
f"poll_interval={poll_interval}s",
|
||
f"endpoint={base_url}{endpoint_path}",
|
||
f"responses={run_dir}",
|
||
)
|
||
|
||
dump_json(
|
||
run_dir / "request.json",
|
||
{
|
||
"base_url": base_url,
|
||
"endpoint_path": endpoint_path,
|
||
"request_mode": request_mode,
|
||
"request_style": request_style,
|
||
"model": model,
|
||
"prompt": prompt,
|
||
"size": requested_size,
|
||
"resolution": payload.get("resolution"),
|
||
"n": payload["n"],
|
||
"reference_sources": reference_sources,
|
||
"mask_source": mask_source,
|
||
"background": background,
|
||
"input_fidelity": input_fidelity,
|
||
"min_output_edge": min_output_edge,
|
||
"output_dir": str(run_dir),
|
||
"request_timeout_seconds": request_timeout,
|
||
"max_poll_wait_seconds": max_poll_wait,
|
||
"initial_poll_delay_seconds": initial_poll_delay,
|
||
"poll_interval_seconds": poll_interval,
|
||
"payload": payload,
|
||
},
|
||
)
|
||
|
||
try:
|
||
if request_mode == "edit":
|
||
reference_inputs = [
|
||
image_source_to_input(source, f"reference-{index}", session)
|
||
for index, source in enumerate(reference_sources, start=1)
|
||
]
|
||
mask_input = image_source_to_input(mask_source, "mask-image", session) if mask_source else None
|
||
resp = session.post(
|
||
f"{base_url}{endpoint_path}",
|
||
data=payload_to_form_fields(payload),
|
||
files=build_edit_files(request_style, reference_inputs, mask_input),
|
||
headers={"Authorization": f"Bearer {api_key}"},
|
||
timeout=(DEFAULT_CONNECT_TIMEOUT, request_timeout),
|
||
)
|
||
else:
|
||
resp = session.post(
|
||
f"{base_url}{endpoint_path}",
|
||
json=payload,
|
||
headers={
|
||
"Authorization": f"Bearer {api_key}",
|
||
"Content-Type": "application/json",
|
||
},
|
||
timeout=(DEFAULT_CONNECT_TIMEOUT, request_timeout),
|
||
)
|
||
|
||
post_record = save_http_response(run_dir / "post-response.json", resp)
|
||
resp.raise_for_status()
|
||
data = post_record.get("json")
|
||
if not isinstance(data, dict):
|
||
raise RuntimeError(f"post response is not a JSON object: {post_record}")
|
||
|
||
response_data = data.get("data")
|
||
items = extract_image_items(response_data)
|
||
if items:
|
||
image_paths = save_output_images(items, session, run_dir, stem, min_output_edge)
|
||
dump_json(
|
||
run_dir / "result.json",
|
||
{
|
||
"images": image_paths,
|
||
"response_path": str(run_dir / "post-response.json"),
|
||
"request_mode": request_mode,
|
||
"request_style": request_style,
|
||
"model": model,
|
||
"size": requested_size,
|
||
"n": payload["n"],
|
||
"reference_count": len(reference_sources),
|
||
},
|
||
)
|
||
print(_json_dumps({"images": image_paths, "output_dir": str(run_dir)}))
|
||
return
|
||
|
||
task_id = extract_task_id(response_data)
|
||
if task_id:
|
||
items = poll_task_images(
|
||
task_id,
|
||
base_url,
|
||
{"Authorization": f"Bearer {api_key}"},
|
||
session,
|
||
run_dir,
|
||
max_poll_wait,
|
||
initial_poll_delay,
|
||
poll_interval,
|
||
)
|
||
image_paths = save_output_images(items, session, run_dir, stem, min_output_edge)
|
||
dump_json(
|
||
run_dir / "result.json",
|
||
{
|
||
"images": image_paths,
|
||
"response_path": str(run_dir / "post-response.json"),
|
||
"request_mode": request_mode,
|
||
"request_style": request_style,
|
||
"model": model,
|
||
"size": requested_size,
|
||
"n": payload["n"],
|
||
"reference_count": len(reference_sources),
|
||
},
|
||
)
|
||
print(_json_dumps({"images": image_paths, "output_dir": str(run_dir)}))
|
||
return
|
||
|
||
raise RuntimeError(f"unexpected response: {data}")
|
||
except BaseException as exc:
|
||
save_error_artifact(run_dir / "error.json", exc, {"output_dir": str(run_dir)})
|
||
raise
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|