kit初版,模型引入,agent优化

This commit is contained in:
2026-04-13 00:41:23 +08:00
parent 9b053e302b
commit c435ab15cf
14097 changed files with 5032 additions and 2676248 deletions

View File

@@ -1,76 +1,220 @@
"""Replicate 图像生成服务封装"""
"""图像生成服务 — Provider 抽象层"""
import base64
import mimetypes
import os
import uuid
from abc import ABC, abstractmethod
from pathlib import Path
from typing import Any
import httpx
import replicate
from replicate import Client as ReplicateClient
GENERATED_DIR = Path(__file__).parent.parent.parent / "generated"
from app.config import get_image_model_config, get_image_aspect_ratio, get_image_output_format
BACKEND_ROOT = Path(__file__).parent.parent.parent
GENERATED_DIR = BACKEND_ROOT / "generated"
def _make_replicate_client() -> ReplicateClient:
"""
创建 Replicate 客户端。
Replicate SDK 内部创建 httpx transport 时不读取代理环境变量,
需要我们手动把代理配置注入到 transport 中。
"""
timeout = httpx.Timeout(5.0, read=300.0, write=30.0, connect=30.0, pool=10.0)
proxy_url = os.environ.get("HTTPS_PROXY") or os.environ.get("HTTP_PROXY")
transport_kwargs: dict[str, Any] = {}
if proxy_url:
transport_kwargs["proxy"] = proxy_url
our_transport = httpx.AsyncHTTPTransport(**transport_kwargs)
client = ReplicateClient(
timeout=timeout,
transport=our_transport,
)
return client
_replicate_client = _make_replicate_client()
# ─── 公共工具 ──────────────────────────────────────────
def to_data_uri(image_path: str) -> str:
"""将本地路径或已有 URL 转为 Replicate 可接受的格式data URI 或原始 URL"""
if image_path.startswith("data:") or image_path.startswith("http"):
return image_path
local = BACKEND_ROOT / image_path.lstrip("/")
if local.exists():
mime = mimetypes.guess_type(str(local))[0] or "image/png"
b64 = base64.b64encode(local.read_bytes()).decode()
return f"data:{mime};base64,{b64}"
return image_path
async def _download_image(url: str) -> str:
"""下载远程图片到本地 generated/ 目录,返回本地 URL 路径。"""
filename = f"{uuid.uuid4().hex}.png"
filepath = GENERATED_DIR / filename
async with httpx.AsyncClient() as client:
proxy = os.environ.get("HTTPS_PROXY") or os.environ.get("HTTP_PROXY")
async with httpx.AsyncClient(proxy=proxy, timeout=httpx.Timeout(60.0)) as client:
resp = await client.get(url, follow_redirects=True)
resp.raise_for_status()
filepath.write_bytes(resp.content)
return f"/generated/{filename}"
# ─── Provider 抽象基类 ─────────────────────────────────
class ImageProvider(ABC):
"""所有图像生成 provider 的基类。"""
@abstractmethod
async def generate(
self,
model_config: dict[str, Any],
prompt: str,
num_images: int = 1,
ref_image_url: str | None = None,
) -> list[str]:
"""生成图片并返回本地 URL 列表。"""
...
class ReplicateProvider(ImageProvider):
"""通过 Replicate API 调用模型。"""
async def generate(
self,
model_config: dict[str, Any],
prompt: str,
num_images: int = 1,
ref_image_url: str | None = None,
) -> list[str]:
replicate_model_id = model_config["model_id"]
default_params: dict[str, Any] = model_config.get("default_params", {})
local_urls: list[str] = []
try:
input_params = self._build_input(
model_config, default_params, prompt, num_images, ref_image_url
)
output = await _replicate_client.async_run(
replicate_model_id, input=input_params, wait=False
)
items = output if isinstance(output, list) else [output]
for item in items:
url = str(item)
if url.startswith("https://") or url.startswith("http://") or url.startswith("data:"):
local_urls.append(await _download_image(url))
else:
local_urls.append(f"[生成失败: 模型返回非图片内容: {url[:200]}]")
except Exception as e:
detail = str(e) or f"{type(e).__name__}: {repr(e)}"
local_urls.append(f"[生成失败: {detail}]")
return local_urls
@staticmethod
def _build_input(
model_config: dict[str, Any],
default_params: dict[str, Any],
prompt: str,
num_images: int = 1,
ref_image_url: str | None = None,
) -> dict[str, Any]:
"""
根据模型配置构建 Replicate input 参数。
通过 model_config 中的标志字段自动适配不同模型:
- supports_ref_image / ref_image_param: 参考图注入
- num_images_param: 各模型的批量生成参数名(如 number_of_images / num_outputs
"""
model_id = model_config["model_id"]
supports_ref = model_config.get("supports_ref_image", False)
ref_param_name = model_config.get("ref_image_param", "image")
num_images_param = model_config.get("num_images_param", "num_outputs")
# ── 支持参考图的模型IP-Adapter 系列)──
if supports_ref:
params: dict[str, Any] = {"prompt": prompt}
for k, v in default_params.items():
params[k] = v
params[num_images_param] = num_images
if ref_image_url:
params[ref_param_name] = to_data_uri(ref_image_url)
return params
# ── Flux 系列(纯文生图)──
is_flux = "flux" in model_id.lower()
params = {"prompt": prompt, num_images_param: num_images}
if is_flux:
params["aspect_ratio"] = default_params.get(
"aspect_ratio", get_image_aspect_ratio()
)
params["output_format"] = default_params.get(
"output_format", get_image_output_format()
)
else:
# SDXL 类模型
params["width"] = default_params.get("width", 1024)
params["height"] = default_params.get("height", 1024)
if "num_inference_steps" in default_params:
params["num_inference_steps"] = default_params["num_inference_steps"]
if "guidance_scale" in default_params:
params["guidance_scale"] = default_params["guidance_scale"]
return params
# ─── Provider 注册 ─────────────────────────────────────
_PROVIDERS: dict[str, ImageProvider] = {
"replicate": ReplicateProvider(),
}
# ─── 公开 API ──────────────────────────────────────────
class GenerateResult:
"""图片生成结果,包含生成的 URL 列表和实际使用的模型信息。"""
def __init__(self, urls: list[str], model_name: str, model_id: str):
self.urls = urls
self.model_name = model_name
self.model_id = model_id
async def generate_images(
prompt: str,
num_images: int = 1,
ref_image_url: str | None = None,
) -> list[str]:
model_id: str | None = None,
) -> GenerateResult:
"""
调用 Replicate 生成图片
统一入口:根据 model_id 查注册表,分发到对应 provider
返回本地可访问的图片 URL 列表
model_id 为空时使用 .env 中配置的默认模型
"""
local_urls = []
config = get_image_model_config(model_id)
provider_name = config.get("provider", "replicate")
provider = _PROVIDERS.get(provider_name)
for _ in range(num_images):
try:
if ref_image_url and ref_image_url.startswith("/"):
# 本地路径转为 file URI 不适用于 Replicate
# 需要用户上传的图先通过后端 URL 访问
# MVP 阶段:参考图作为 prompt 的文字补充,不直接传给模型
# 仅使用 flux-schnell 文生图
output = await replicate.async_run(
"black-forest-labs/flux-schnell",
input={
"prompt": prompt,
"num_outputs": 1,
"aspect_ratio": "1:1",
"output_format": "png",
},
)
else:
output = await replicate.async_run(
"black-forest-labs/flux-schnell",
input={
"prompt": prompt,
"num_outputs": 1,
"aspect_ratio": "1:1",
"output_format": "png",
},
)
model_name = config.get("name", "未知模型")
resolved_id = config.get("id", model_id or "unknown")
# output 是 FileOutput 列表或单个 URL
if isinstance(output, list):
for item in output:
url = str(item)
local_url = await _download_image(url)
local_urls.append(local_url)
else:
url = str(output)
local_url = await _download_image(url)
local_urls.append(local_url)
if not provider:
return GenerateResult([f"[未知 provider: {provider_name}]"], model_name, resolved_id)
except Exception as e:
local_urls.append(f"[生成失败: {e}]")
return local_urls
urls = await provider.generate(config, prompt, num_images, ref_image_url)
return GenerateResult(urls, model_name, resolved_id)