Add per-operation AI provider routing

Each operation (inpaint, txt2img, img2img, outpaint) can now use a different
provider. Resolution order: per-op override → global AI_PROVIDER default.

Example: txt2img→openai, inpaint→invokeai, everything else→invokeai default.

Backend:
- config.py: add AI_PROVIDER_INPAINT / TXT2IMG / IMG2IMG / OUTPAINT settings
- remote_provider.py: get_remote_provider(operation) resolves override then default;
  _build_provider() extracted as shared factory; _OP_FIELD maps op→setting name
- ai_tools.py: each endpoint passes its operation to _require_remote();
  GET /api/config runs per-op health checks concurrently, returns operations map
  and overrides; POST /api/config accepts and applies per-op override fields

Frontend:
- ai_provider_settings.js: four new selects (inpaint/txt2img/img2img/outpaint);
  persists to localStorage and sends per-op fields to POST /api/config
- provider-badge.js: shows override summary (e.g. "invokeai · txt2img→openai")
  and per-op health in tooltip
- .env.example: document per-op override env vars with examples

https://claude.ai/code/session_01B58MaJCU1R6KwBDJCp8AfN
This commit is contained in:
Claude
2026-06-09 18:14:12 +00:00
parent b2a12c356f
commit d01c11f948
6 changed files with 191 additions and 62 deletions
+10 -1
View File
@@ -13,9 +13,18 @@ class Settings(BaseSettings):
# AI Provider
# Local: blank or "mock" — always available, no config needed
# Remote (set ONE): openai | invokeai | comfyui | replicate | stability
# Remote default (used for any operation without a specific override):
# openai | invokeai | comfyui | replicate | stability
ai_provider: str = "mock"
# Per-operation provider overrides — blank means use ai_provider default.
# Operations: inpaint, txt2img, img2img, outpaint
# Example: AI_PROVIDER_TXT2IMG=openai (use OpenAI for text-to-image only)
ai_provider_inpaint: str = "" # remote inpaint / replace selection
ai_provider_txt2img: str = "" # text-to-image
ai_provider_img2img: str = "" # image-to-image
ai_provider_outpaint: str = "" # expand canvas
# Provider API Keys
openai_api_key: str = ""
openai_model: str = "dall-e-3"
+68 -42
View File
@@ -75,13 +75,15 @@ def _encode(data: bytes) -> str:
return base64.b64encode(data).decode()
def _require_remote():
def _require_remote(operation: str = None):
from app.services.remote_provider import get_remote_provider
provider = get_remote_provider()
provider = get_remote_provider(operation)
if provider is None:
op_hint = f"AI_PROVIDER_{operation.upper()} or " if operation else ""
raise HTTPException(
status_code=503,
detail="No remote AI provider configured. Set AI_PROVIDER in .env (openai / invokeai / comfyui)."
detail=f"No remote AI provider configured for '{operation or 'default'}'. "
f"Set {op_hint}AI_PROVIDER in .env (openai / invokeai / comfyui)."
)
return provider
@@ -173,7 +175,7 @@ async def background_remove(req: BgRemoveRequest):
@router.post("/inpaint/remote")
async def inpaint_remote(req: InpaintRemoteRequest):
"""Inpaint via configured remote provider (InvokeAI / ComfyUI / OpenAI)."""
provider = _require_remote()
provider = _require_remote("inpaint")
try:
params = {
"negative_prompt": req.negative_prompt or "",
@@ -192,7 +194,7 @@ async def inpaint_remote(req: InpaintRemoteRequest):
@router.post("/generate/txt2img")
async def txt2img(req: Txt2ImgRequest):
"""Text-to-image via configured remote provider."""
provider = _require_remote()
provider = _require_remote("txt2img")
try:
params = {
"negative_prompt": req.negative_prompt or "",
@@ -212,7 +214,7 @@ async def txt2img(req: Txt2ImgRequest):
@router.post("/generate/img2img")
async def img2img(req: Img2ImgRequest):
"""Image-to-image via configured remote provider."""
provider = _require_remote()
provider = _require_remote("img2img")
try:
params = {
"negative_prompt": req.negative_prompt or "",
@@ -231,7 +233,7 @@ async def img2img(req: Img2ImgRequest):
@router.post("/generate/outpaint")
async def outpaint(req: OutpaintRequest):
"""Expand canvas in given direction via remote provider."""
provider = _require_remote()
provider = _require_remote("outpaint")
if req.direction not in ("left", "right", "top", "bottom"):
raise HTTPException(status_code=400, detail="direction must be left/right/top/bottom")
try:
@@ -246,6 +248,12 @@ async def outpaint(req: OutpaintRequest):
class ConfigUpdateRequest(BaseModel):
ai_provider: Optional[str] = None
# Per-operation overrides (blank = use default)
ai_provider_inpaint: Optional[str] = None
ai_provider_txt2img: Optional[str] = None
ai_provider_img2img: Optional[str] = None
ai_provider_outpaint: Optional[str] = None
# Credentials / URLs
openai_api_key: Optional[str] = None
openai_model: Optional[str] = None
invokeai_url: Optional[str] = None
@@ -265,49 +273,59 @@ async def update_config(req: ConfigUpdateRequest):
"""
from app.config import settings
if req.ai_provider is not None:
settings.ai_provider = req.ai_provider
if req.openai_api_key:
settings.openai_api_key = req.openai_api_key
if req.openai_model:
settings.openai_model = req.openai_model
if req.invokeai_url is not None:
settings.invokeai_url = req.invokeai_url
if req.invokeai_default_model:
settings.invokeai_default_model = req.invokeai_default_model
if req.comfyui_url is not None:
settings.comfyui_url = req.comfyui_url
if req.comfyui_default_model:
settings.comfyui_default_model = req.comfyui_default_model
if req.replicate_api_key:
settings.replicate_api_key = req.replicate_api_key
if req.stability_api_key:
settings.stability_api_key = req.stability_api_key
_str_fields = [
"ai_provider", "ai_provider_inpaint", "ai_provider_txt2img",
"ai_provider_img2img", "ai_provider_outpaint",
"openai_api_key", "openai_model",
"invokeai_url", "invokeai_default_model",
"comfyui_url", "comfyui_default_model",
"replicate_api_key", "stability_api_key",
]
for field in _str_fields:
val = getattr(req, field, None)
if val is not None:
setattr(settings, field, val)
return {"status": "ok", "ai_provider": settings.ai_provider}
return {
"status": "ok",
"ai_provider": settings.ai_provider,
"overrides": {
"inpaint": settings.ai_provider_inpaint or None,
"txt2img": settings.ai_provider_txt2img or None,
"img2img": settings.ai_provider_img2img or None,
"outpaint": settings.ai_provider_outpaint or None,
}
}
async def _check_provider(operation: str) -> dict:
"""Health-check the provider for a specific operation."""
from app.services.remote_provider import get_remote_provider
try:
p = get_remote_provider(operation)
if p is None:
return {"provider": None, "healthy": False}
healthy = await asyncio.wait_for(p.health(), timeout=5.0)
return {"provider": p.__class__.__name__.replace("Provider", "").lower(), "healthy": healthy}
except Exception:
return {"provider": None, "healthy": False}
@router.get("/config")
async def get_config():
"""
Return capability flags so the frontend can show/hide tools.
Frontend reads this on load.
Includes per-operation provider assignments and health status.
"""
from app.services.remote_provider import get_remote_provider
from app.config import settings
remote_caps: list[str] = []
remote_healthy = False
provider_name = (settings.ai_provider or "").lower()
# Run health checks for each operation concurrently
ops = ["inpaint", "txt2img", "img2img", "outpaint"]
results = await asyncio.gather(*[_check_provider(op) for op in ops])
op_status = dict(zip(ops, results))
if provider_name in ("openai", "invokeai", "comfyui"):
try:
provider = get_remote_provider()
if provider:
remote_caps = provider.capabilities()
remote_healthy = await asyncio.wait_for(provider.health(), timeout=5.0)
except Exception:
remote_healthy = False
# Default provider for display (used when no per-op override)
default_name = (settings.ai_provider or "").lower() or None
return {
"local": {
@@ -317,8 +335,16 @@ async def get_config():
"gpu_detected": gpu_available(),
},
"remote": {
"provider": provider_name or None,
"capabilities": remote_caps,
"healthy": remote_healthy,
"default_provider": default_name,
# Legacy field kept for backwards compat with badge/capabilities checks
"provider": default_name,
"healthy": any(v["healthy"] for v in op_status.values()),
"operations": op_status,
"overrides": {
"inpaint": settings.ai_provider_inpaint or None,
"txt2img": settings.ai_provider_txt2img or None,
"img2img": settings.ai_provider_img2img or None,
"outpaint": settings.ai_provider_outpaint or None,
},
}
}
+41 -6
View File
@@ -407,25 +407,60 @@ class ComfyUIProvider(RemoteAIProvider):
return ["inpaint", "txt2img", "img2img", "outpaint"]
def get_remote_provider() -> Optional[RemoteAIProvider]:
"""Return the configured remote provider, or None if not configured."""
def _build_provider(name: str) -> Optional[RemoteAIProvider]:
"""Instantiate a named provider from current settings."""
from app.config import settings
provider = (settings.ai_provider or "").lower()
name = (name or "").lower().strip()
if provider == "openai":
if name == "openai":
if not settings.openai_api_key:
return None
return OpenAIRemoteProvider(settings.openai_api_key, settings.openai_model)
if provider == "invokeai":
if name == "invokeai":
if not settings.invokeai_url:
return None
return InvokeAIProvider(settings.invokeai_url, settings.invokeai_default_model)
if provider == "comfyui":
if name == "comfyui":
if not settings.comfyui_url:
return None
return ComfyUIProvider(settings.comfyui_url, settings.comfyui_default_model)
return None
# Map operation names to the settings field that holds the override
_OP_FIELD = {
"inpaint": "ai_provider_inpaint",
"txt2img": "ai_provider_txt2img",
"img2img": "ai_provider_img2img",
"outpaint": "ai_provider_outpaint",
}
def get_remote_provider(operation: Optional[str] = None) -> Optional[RemoteAIProvider]:
"""
Return the provider for a given operation.
Resolution order:
1. Per-operation override (AI_PROVIDER_INPAINT, AI_PROVIDER_TXT2IMG, etc.)
2. Global default (AI_PROVIDER)
3. None (local-only mode)
Example .env for mixed setup:
AI_PROVIDER=invokeai # default for inpaint/img2img/outpaint
AI_PROVIDER_TXT2IMG=openai # use OpenAI only for text-to-image
"""
from app.config import settings
if operation and operation in _OP_FIELD:
override = getattr(settings, _OP_FIELD[operation], "")
if override:
provider = _build_provider(override)
if provider is not None:
return provider
# override configured but not usable (missing key/url) — fall through to default
return _build_provider(settings.ai_provider)