diff --git a/.env.example b/.env.example index 38d7804..aab2dce 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/backend/app/config.py b/backend/app/config.py index 1ac4fb4..0b80ad9 100644 --- a/backend/app/config.py +++ b/backend/app/config.py @@ -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" diff --git a/backend/app/routers/ai_tools.py b/backend/app/routers/ai_tools.py index 0f040aa..243c5ba 100644 --- a/backend/app/routers/ai_tools.py +++ b/backend/app/routers/ai_tools.py @@ -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, + }, } } diff --git a/backend/app/services/remote_provider.py b/backend/app/services/remote_provider.py index 6e899f7..6fb8d31 100644 --- a/backend/app/services/remote_provider.py +++ b/backend/app/services/remote_provider.py @@ -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) diff --git a/frontend/src/js/core/components/provider-badge.js b/frontend/src/js/core/components/provider-badge.js index 0af35f2..8724733 100644 --- a/frontend/src/js/core/components/provider-badge.js +++ b/frontend/src/js/core/components/provider-badge.js @@ -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'; diff --git a/frontend/src/js/modules/tools/ai_provider_settings.js b/frontend/src/js/modules/tools/ai_provider_settings.js index 088029a..4268346 100644 --- a/frontend/src/js/modules/tools/ai_provider_settings.js +++ b/frontend/src/js/modules/tools/ai_provider_settings.js @@ -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: '
Per-operation overrides — blank = use default above
', + }, + { + 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`, {