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:
@@ -26,6 +26,13 @@
|
||||
|
||||
AI_PROVIDER=replicate
|
||||
|
||||
# Per-operation provider overrides (optional — blank means use AI_PROVIDER above)
|
||||
# Example: use OpenAI for text-to-image (best quality) but InvokeAI for everything else
|
||||
#AI_PROVIDER_TXT2IMG=openai
|
||||
#AI_PROVIDER_INPAINT=invokeai
|
||||
#AI_PROVIDER_IMG2IMG=invokeai
|
||||
#AI_PROVIDER_OUTPAINT=invokeai
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# STEP 2: Get Your API Key
|
||||
|
||||
+10
-1
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -34,8 +34,19 @@ export async function mountProviderBadge(container) {
|
||||
dot.style.background = '#44cc44';
|
||||
badge.style.background = '#1a2a1a';
|
||||
badge.style.color = '#aaffaa';
|
||||
label.textContent = remote.provider + (local.gpu_detected ? ' · GPU' : ' · CPU');
|
||||
badge.title = 'Remote provider: ' + remote.provider + '\nCapabilities: ' + (remote.capabilities || []).join(', ');
|
||||
|
||||
// Show override summary if any operations use different providers
|
||||
var overrides = remote.overrides || {};
|
||||
var overrideEntries = Object.entries(overrides).filter(([, v]) => v);
|
||||
var overrideStr = overrideEntries.length
|
||||
? ' · ' + overrideEntries.map(([k, v]) => `${k}→${v}`).join(', ')
|
||||
: '';
|
||||
label.textContent = remote.provider + overrideStr + (local.gpu_detected ? ' · GPU' : '');
|
||||
|
||||
var opLines = Object.entries(remote.operations || {})
|
||||
.map(([op, s]) => `${op}: ${s.provider || remote.provider} ${s.healthy ? '✓' : '✗'}`)
|
||||
.join('\n');
|
||||
badge.title = opLines || ('Provider: ' + remote.provider);
|
||||
} else if (remote.provider && !remote.healthy) {
|
||||
dot.style.background = '#ffaa00';
|
||||
badge.style.background = '#2a2000';
|
||||
|
||||
@@ -47,11 +47,44 @@ class Tools_ai_provider_settings_class {
|
||||
},
|
||||
{
|
||||
name: 'provider',
|
||||
title: 'Remote provider:',
|
||||
title: 'Default provider (used unless overridden below):',
|
||||
value: ls_get('provider', remote.provider || ''),
|
||||
values: ['', 'openai', 'invokeai', 'comfyui', 'replicate'],
|
||||
type: 'select',
|
||||
},
|
||||
// ── Per-operation overrides ───────────────────────────────
|
||||
{
|
||||
title: '',
|
||||
html: '<div style="font-size:11px;color:#888;margin:2px 0 6px;">Per-operation overrides — blank = use default above</div>',
|
||||
},
|
||||
{
|
||||
name: 'provider_inpaint',
|
||||
title: 'Inpaint / Replace Selection:',
|
||||
value: ls_get('provider_inpaint', remote.overrides?.inpaint || ''),
|
||||
values: ['', 'openai', 'invokeai', 'comfyui', 'replicate'],
|
||||
type: 'select',
|
||||
},
|
||||
{
|
||||
name: 'provider_txt2img',
|
||||
title: 'Text → Image:',
|
||||
value: ls_get('provider_txt2img', remote.overrides?.txt2img || ''),
|
||||
values: ['', 'openai', 'invokeai', 'comfyui', 'replicate'],
|
||||
type: 'select',
|
||||
},
|
||||
{
|
||||
name: 'provider_img2img',
|
||||
title: 'Image → Image:',
|
||||
value: ls_get('provider_img2img', remote.overrides?.img2img || ''),
|
||||
values: ['', 'openai', 'invokeai', 'comfyui', 'replicate'],
|
||||
type: 'select',
|
||||
},
|
||||
{
|
||||
name: 'provider_outpaint',
|
||||
title: 'Expand Canvas (Outpaint):',
|
||||
value: ls_get('provider_outpaint', remote.overrides?.outpaint || ''),
|
||||
values: ['', 'openai', 'invokeai', 'comfyui', 'replicate'],
|
||||
type: 'select',
|
||||
},
|
||||
// ── OpenAI ────────────────────────────────────────────────
|
||||
{
|
||||
name: 'openai_key',
|
||||
@@ -108,7 +141,11 @@ class Tools_ai_provider_settings_class {
|
||||
|
||||
async _save(params) {
|
||||
// Persist to localStorage
|
||||
ls_set('provider', params.provider || '');
|
||||
ls_set('provider', params.provider || '');
|
||||
ls_set('provider_inpaint', params.provider_inpaint || '');
|
||||
ls_set('provider_txt2img', params.provider_txt2img || '');
|
||||
ls_set('provider_img2img', params.provider_img2img || '');
|
||||
ls_set('provider_outpaint', params.provider_outpaint || '');
|
||||
ls_set('openai_key', params.openai_key || '');
|
||||
ls_set('openai_model', params.openai_model || 'dall-e-3');
|
||||
ls_set('invokeai_url', params.invokeai_url || '');
|
||||
@@ -117,17 +154,21 @@ class Tools_ai_provider_settings_class {
|
||||
ls_set('comfyui_model', params.comfyui_model || 'v1-5-pruned-emaonly.ckpt');
|
||||
ls_set('replicate_key', params.replicate_key || '');
|
||||
|
||||
// Push to backend (requires a running server that accepts runtime config)
|
||||
// Push to backend
|
||||
try {
|
||||
var payload = {
|
||||
ai_provider: params.provider || '',
|
||||
openai_api_key: params.openai_key || '',
|
||||
openai_model: params.openai_model || 'dall-e-3',
|
||||
invokeai_url: params.invokeai_url || '',
|
||||
invokeai_default_model: params.invokeai_model || 'flux-dev',
|
||||
comfyui_url: params.comfyui_url || '',
|
||||
comfyui_default_model: params.comfyui_model || '',
|
||||
replicate_api_key: params.replicate_key || '',
|
||||
ai_provider: params.provider || '',
|
||||
ai_provider_inpaint: params.provider_inpaint || '',
|
||||
ai_provider_txt2img: params.provider_txt2img || '',
|
||||
ai_provider_img2img: params.provider_img2img || '',
|
||||
ai_provider_outpaint: params.provider_outpaint || '',
|
||||
openai_api_key: params.openai_key || '',
|
||||
openai_model: params.openai_model || 'dall-e-3',
|
||||
invokeai_url: params.invokeai_url || '',
|
||||
invokeai_default_model: params.invokeai_model || 'flux-dev',
|
||||
comfyui_url: params.comfyui_url || '',
|
||||
comfyui_default_model: params.comfyui_model || '',
|
||||
replicate_api_key: params.replicate_key || '',
|
||||
};
|
||||
var base = window.API_BASE_URL || '';
|
||||
var r = await fetch(`${base}/api/config`, {
|
||||
|
||||
Reference in New Issue
Block a user