311 lines
10 KiB
Python
311 lines
10 KiB
Python
"""
|
|
Command-line interface for single Articraft generation runs.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import asyncio
|
|
import json
|
|
import logging
|
|
import os
|
|
import sys
|
|
from pathlib import Path
|
|
from typing import Awaitable, Callable, Optional
|
|
|
|
from agent.cost import max_cost_usd_from_env, parse_max_cost_usd
|
|
from agent.payload_preview import build_provider_payload_preview
|
|
from agent.prompts import normalize_sdk_package
|
|
from agent.providers.factory import infer_provider_from_model_id, validate_provider_credentials
|
|
from agent.run_context import _default_model_id
|
|
from agent.single_run import run_from_input
|
|
from agent.tools import build_initial_user_content as _build_initial_user_content
|
|
from agent.tools import resolve_image_path as _resolve_image_path
|
|
from agent.tui.single_run import LLMWaitAwareStreamHandler
|
|
from articraft.config import (
|
|
default_model_from_env,
|
|
default_thinking_level_from_env,
|
|
load_repo_env,
|
|
)
|
|
from articraft.values import (
|
|
PROVIDER_VALUES,
|
|
THINKING_LEVEL_VALUE_SET,
|
|
THINKING_LEVEL_VALUES,
|
|
ProviderName,
|
|
)
|
|
|
|
|
|
def _resolve_data_dir(repo_root: Path, data_dir: Path | None) -> Path:
|
|
if data_dir is not None:
|
|
return data_dir.expanduser().resolve()
|
|
configured = os.getenv("ARTICRAFT_DATA_DIR")
|
|
if configured:
|
|
return Path(configured).expanduser().resolve()
|
|
return repo_root.expanduser().resolve() / "data"
|
|
|
|
|
|
def _load_qc_blurb_text(qc_blurb_path: Optional[str], *, repo_root: Path) -> Optional[str]:
|
|
if not qc_blurb_path:
|
|
return None
|
|
|
|
path = Path(qc_blurb_path)
|
|
if not path.is_absolute():
|
|
path = (repo_root / path).resolve()
|
|
if not path.exists():
|
|
raise FileNotFoundError(f"QC blurb file not found: {path}")
|
|
|
|
text = path.read_text(encoding="utf-8").replace("\r\n", "\n")
|
|
if text and not text.endswith("\n"):
|
|
text += "\n"
|
|
return text
|
|
|
|
|
|
def _build_prompt_with_qc(prompt: str, qc_blurb_text: Optional[str]) -> str:
|
|
if not qc_blurb_text:
|
|
return prompt
|
|
|
|
return (
|
|
f"{prompt.rstrip()}\n\n"
|
|
"-----\n\n"
|
|
"The following is a QC checklist to use as a final pass before declaring the model finished:\n\n"
|
|
f"{qc_blurb_text}"
|
|
)
|
|
|
|
|
|
def _resolve_model_and_provider(
|
|
args: argparse.Namespace,
|
|
parser: argparse.ArgumentParser,
|
|
) -> tuple[str | None, str]:
|
|
model_id = args.model
|
|
provider = args.provider
|
|
if model_id is None and provider is None:
|
|
model_id = default_model_from_env()
|
|
if provider is None:
|
|
provider = infer_provider_from_model_id(model_id)
|
|
if provider is None:
|
|
parser.error(
|
|
f"Unable to infer provider for model '{model_id}'. "
|
|
"Pass --provider explicitly or use a known OpenAI, Gemini, Anthropic, "
|
|
"DashScope, OpenRouter, DeepSeek, or Codex CLI model ID."
|
|
)
|
|
return model_id, provider
|
|
|
|
|
|
def _resolve_thinking_level(args: argparse.Namespace, parser: argparse.ArgumentParser) -> str:
|
|
thinking_level = args.thinking or default_thinking_level_from_env()
|
|
if thinking_level not in THINKING_LEVEL_VALUE_SET:
|
|
parser.error("ARTICRAFT_THINKING_LEVEL must be one of: " + ", ".join(THINKING_LEVEL_VALUES))
|
|
return thinking_level
|
|
|
|
|
|
def main(
|
|
argv: list[str] | None = None,
|
|
*,
|
|
run_from_input_func: Callable[..., Awaitable[int]] = run_from_input,
|
|
build_provider_payload_preview_func: Callable[..., dict] = build_provider_payload_preview,
|
|
) -> int:
|
|
logging.basicConfig(
|
|
level=logging.INFO,
|
|
format="%(levelname)s: %(message)s",
|
|
handlers=[LLMWaitAwareStreamHandler()],
|
|
)
|
|
parser = argparse.ArgumentParser(
|
|
description="Generate an articulated object and persist it to local library storage."
|
|
)
|
|
parser.add_argument("--prompt", required=True, help="Text prompt for the object.")
|
|
parser.add_argument(
|
|
"--image",
|
|
default=None,
|
|
help="Optional reference image to augment --prompt.",
|
|
)
|
|
parser.add_argument(
|
|
"--provider",
|
|
default=None,
|
|
choices=PROVIDER_VALUES,
|
|
help="LLM provider.",
|
|
)
|
|
parser.add_argument(
|
|
"--repo-root",
|
|
type=Path,
|
|
default=Path(__file__).resolve().parents[1],
|
|
help="Articraft code repository root.",
|
|
)
|
|
parser.add_argument(
|
|
"--data-dir",
|
|
type=Path,
|
|
default=None,
|
|
help="Articraft data root. Defaults to ARTICRAFT_DATA_DIR, then <repo-root>/data.",
|
|
)
|
|
parser.add_argument("--label", default=None, help="Optional label for the saved record.")
|
|
parser.add_argument("--tag", action="append", default=[], help="Optional tag. Repeatable.")
|
|
parser.add_argument(
|
|
"--category",
|
|
default=None,
|
|
help="Optional category slug to attach to the record.",
|
|
)
|
|
parser.add_argument("--model", default=None, help="Model id (provider-specific).")
|
|
parser.add_argument(
|
|
"--openai-transport",
|
|
default="http",
|
|
choices=["http", "websocket"],
|
|
help=(
|
|
"Transport for --provider openai. "
|
|
"`websocket` uses Responses WebSocket mode and enables response storage."
|
|
),
|
|
)
|
|
parser.add_argument(
|
|
"--thinking",
|
|
default=None,
|
|
choices=THINKING_LEVEL_VALUES,
|
|
help="Thinking budget level.",
|
|
)
|
|
parser.add_argument("--max-turns", type=int, default=None)
|
|
parser.add_argument(
|
|
"--max-cost-usd",
|
|
type=float,
|
|
default=None,
|
|
help="Optional per-run USD budget. Stops after the first response that pushes cumulative spend above this threshold.",
|
|
)
|
|
parser.add_argument(
|
|
"--system-prompt",
|
|
default="designer_system_prompt.txt",
|
|
help=(
|
|
"Path or generated prompt name for the system prompt file. "
|
|
"Standard designer prompt names resolve to provider-specific generated files automatically."
|
|
),
|
|
)
|
|
parser.add_argument(
|
|
"--qc-blurb",
|
|
default=None,
|
|
help="Path to a markdown QC checklist to append to the prompt.",
|
|
)
|
|
parser.add_argument(
|
|
"--dump-provider-payload",
|
|
action="store_true",
|
|
help="Print the provider request payload for turn 1 and exit (no API call).",
|
|
)
|
|
parser.add_argument(
|
|
"--dump-provider-payload-out",
|
|
default=None,
|
|
help="Write the payload JSON to this path instead of stdout.",
|
|
)
|
|
parser.add_argument(
|
|
"--dump-provider-payload-indent",
|
|
type=int,
|
|
default=2,
|
|
help="JSON indent for --dump-provider-payload (default: 2).",
|
|
)
|
|
parser.add_argument(
|
|
"--sdk-package",
|
|
default="sdk",
|
|
help=argparse.SUPPRESS,
|
|
)
|
|
args = parser.parse_args(argv)
|
|
load_repo_env(args.repo_root)
|
|
model_id_arg, provider = _resolve_model_and_provider(args, parser)
|
|
thinking_level = _resolve_thinking_level(args, parser)
|
|
try:
|
|
sdk_package = normalize_sdk_package(args.sdk_package)
|
|
max_cost_usd = (
|
|
parse_max_cost_usd(args.max_cost_usd, label="--max-cost-usd")
|
|
if args.max_cost_usd is not None
|
|
else max_cost_usd_from_env()
|
|
)
|
|
except ValueError as exc:
|
|
print(str(exc), file=sys.stderr)
|
|
return 1
|
|
|
|
openai_reasoning_summary = "auto"
|
|
|
|
try:
|
|
model_id_arg = _default_model_id(
|
|
provider=provider,
|
|
model_id=model_id_arg,
|
|
thinking_level=thinking_level,
|
|
openai_transport=args.openai_transport,
|
|
openai_reasoning_summary=openai_reasoning_summary,
|
|
)
|
|
except ValueError as exc:
|
|
print(str(exc), file=sys.stderr)
|
|
return 1
|
|
|
|
if provider != ProviderName.OPENAI.value and args.openai_transport != "http":
|
|
print("--openai-transport is only supported for --provider openai.", file=sys.stderr)
|
|
return 1
|
|
|
|
repo_root = args.repo_root.resolve()
|
|
data_root = _resolve_data_dir(repo_root, args.data_dir)
|
|
try:
|
|
qc_blurb_text = _load_qc_blurb_text(args.qc_blurb, repo_root=repo_root)
|
|
except Exception as exc:
|
|
print(f"Failed to load qc blurb: {exc}", file=sys.stderr)
|
|
return 1
|
|
|
|
try:
|
|
image_path = _resolve_image_path(args.image, provider=provider)
|
|
except Exception as exc:
|
|
print(f"Failed to load image: {exc}", file=sys.stderr)
|
|
return 1
|
|
|
|
prompt_with_qc = _build_prompt_with_qc(args.prompt, qc_blurb_text)
|
|
user_content = _build_initial_user_content(prompt_with_qc, image_path=image_path)
|
|
|
|
if args.dump_provider_payload:
|
|
model_id = _default_model_id(
|
|
provider=provider,
|
|
model_id=model_id_arg,
|
|
thinking_level=thinking_level,
|
|
openai_transport=args.openai_transport,
|
|
openai_reasoning_summary=openai_reasoning_summary,
|
|
)
|
|
|
|
payload = build_provider_payload_preview_func(
|
|
user_content,
|
|
provider=provider,
|
|
model_id=model_id,
|
|
openai_transport=args.openai_transport,
|
|
thinking_level=thinking_level,
|
|
system_prompt_path=args.system_prompt,
|
|
sdk_package=sdk_package,
|
|
openai_reasoning_summary=openai_reasoning_summary,
|
|
)
|
|
text = json.dumps(payload, indent=args.dump_provider_payload_indent, ensure_ascii=False)
|
|
if args.dump_provider_payload_out:
|
|
out_path = Path(args.dump_provider_payload_out)
|
|
if not out_path.is_absolute():
|
|
out_path = (Path.cwd() / out_path).resolve()
|
|
out_path.parent.mkdir(parents=True, exist_ok=True)
|
|
out_path.write_text(text, encoding="utf-8")
|
|
else:
|
|
print(text)
|
|
return 0
|
|
|
|
try:
|
|
validate_provider_credentials(provider)
|
|
except ValueError as exc:
|
|
print(str(exc), file=sys.stderr)
|
|
return 1
|
|
|
|
return asyncio.run(
|
|
run_from_input_func(
|
|
user_content,
|
|
prompt_text=prompt_with_qc,
|
|
display_prompt=args.prompt,
|
|
repo_root=repo_root,
|
|
data_root=data_root,
|
|
image_path=image_path,
|
|
provider=provider,
|
|
model_id=model_id_arg,
|
|
openai_transport=args.openai_transport,
|
|
thinking_level=thinking_level,
|
|
max_turns=args.max_turns,
|
|
system_prompt_path=args.system_prompt,
|
|
sdk_package=sdk_package,
|
|
openai_reasoning_summary=openai_reasoning_summary,
|
|
max_cost_usd=max_cost_usd,
|
|
label=args.label,
|
|
tags=list(args.tag or []),
|
|
category_slug=args.category,
|
|
)
|
|
)
|