CoolFace
Apppublic

esotericelf/image_edit_creation

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
main.py466 linesDownload Raw Back to root
1import asyncio2import logging3import os4from contextlib import asynccontextmanager5from typing import Any6 7import fal_client8from fastapi import Depends, FastAPI, Header, HTTPException, Request9from fastapi.exceptions import RequestValidationError10from fastapi.middleware.cors import CORSMiddleware11from fastapi.responses import JSONResponse12from sqlalchemy import select13from sqlalchemy.ext.asyncio import AsyncSession14 15from auth import (16    clear_admin_session_cookie,17    create_admin_session_token,18    is_admin_session_configured,19    read_admin_session,20    require_admin_session,21    set_admin_session_cookie,22    verify_master_password,23)24from database import GiftToken, get_db_session, init_db25from gift_tokens import (26    GiftTokenError,27    consume_gift_token,28    create_gift_token,29    get_gift_token,30    validate_gift_token,31)32from schemas import (33    PREMIUM_RESOLUTION_MULTIPLIERS,34    PREMIUM_RESOLUTIONS,35    AdminLoginRequest,36    AdminLoginResponse,37    CostSafetyViolation,38    GenerateRequest,39    GenerateResponse,40    GenerationMode,41    GiftTokenCreateResponse,42    GiftTokenStatusResponse,43    MAX_NUM_IMAGES,44    Resolution,45    SAFE_RESOLUTION,46    ThinkingLevel,47)48from test_connection import get_health_snapshot, resolve_health_status_code49 50logging.basicConfig(level=logging.INFO)51logger = logging.getLogger(__name__)52 53NANO_BANANA_2_TEXT_ENDPOINT = "fal-ai/nano-banana-2"54NANO_BANANA_2_EDIT_ENDPOINT = "fal-ai/nano-banana-2/edit"55REFERENCE_ASSET_FIELDS = ("image_urls", "video_url", "audio_url", "pdf_url")56PARENT_IMAGE_URL_FIELD = "parent_image_url"57GIFT_TOKEN_HEADER = "x-gift-token"58 59 60@asynccontextmanager61async def lifespan(_app: FastAPI):62    await init_db()63    yield64 65 66app = FastAPI(67    title="Nano Banana 2 Image Generation API",68    description="FastAPI proxy for fal.ai Nano Banana 2 with gift-link access control.",69    version="1.3.0",70    lifespan=lifespan,71)72 73NETLIFY_FRONTEND_ORIGIN = "https://image-edit-creation.netlify.app"74DEFAULT_CORS_ORIGINS = f"{NETLIFY_FRONTEND_ORIGIN},http://localhost:5173"75 76cors_origins = [77    origin.strip()78    for origin in os.getenv("CORS_ORIGINS", DEFAULT_CORS_ORIGINS).split(",")79    if origin.strip() and origin.strip() != "*"80]81if NETLIFY_FRONTEND_ORIGIN not in cors_origins:82    cors_origins.insert(0, NETLIFY_FRONTEND_ORIGIN)83app.add_middleware(84    CORSMiddleware,85    allow_origins=cors_origins,86    allow_credentials=True,87    allow_methods=["*"],88    allow_headers=["*"],89)90 91 92def _is_truthy_env(name: str) -> bool:93    return os.getenv(name, "").strip().lower() in {"1", "true", "yes", "on"}94 95 96def _premium_resolution_allowed() -> bool:97    return _is_truthy_env("ALLOW_PREMIUM_RESOLUTION")98 99 100def _premium_features_allowed() -> bool:101    return _is_truthy_env("ALLOW_PREMIUM_FEATURES")102 103 104def _validate_premium_resolution(request: GenerateRequest) -> None:105    if request.resolution not in PREMIUM_RESOLUTIONS:106        return107 108    if _premium_resolution_allowed():109        logger.info(110            "Premium resolution '%s' allowed via ALLOW_PREMIUM_RESOLUTION.",111            request.resolution.value,112        )113        return114 115    multiplier = PREMIUM_RESOLUTION_MULTIPLIERS[request.resolution]116    raise CostSafetyViolation(117        f"Resolution '{request.resolution.value}' is blocked because it carries a "118        f"{multiplier} pricing premium. Allowed values without authorization: "119        f"'{SAFE_RESOLUTION.value}' or '{Resolution.HALF_K.value}'. "120        f"Set ALLOW_PREMIUM_RESOLUTION=true on the server to opt in.",121        field="resolution",122    )123 124 125def _apply_cost_guardrails(request: GenerateRequest) -> dict[str, Any]:126    _validate_premium_resolution(request)127 128    arguments = request.model_dump(mode="json", exclude_unset=True, exclude_none=True)129    arguments.pop("gift_token", None)130    arguments.pop("is_admin", None)131 132    requested_num_images = arguments.get("num_images", MAX_NUM_IMAGES)133    if requested_num_images > MAX_NUM_IMAGES:134        raise CostSafetyViolation(135            "num_images exceeds the cost-safety cap of 1 per request. "136            "Submit separate requests if you need additional images.",137            field="num_images",138        )139    if requested_num_images != MAX_NUM_IMAGES:140        logger.warning(141            "Cost safety: overriding num_images=%s to %s before upstream dispatch.",142            requested_num_images,143            MAX_NUM_IMAGES,144        )145    arguments["num_images"] = MAX_NUM_IMAGES146 147    if _premium_features_allowed() and request.enable_web_search:148        arguments["enable_web_search"] = True149    else:150        if request.enable_web_search and not _premium_features_allowed():151            logger.warning(152                "Cost safety: blocking enable_web_search=true (premium surcharge). "153                "Set ALLOW_PREMIUM_FEATURES=true to authorize."154            )155        arguments["enable_web_search"] = False156 157    thinking_level = arguments.get("thinking_level")158    if thinking_level == ThinkingLevel.HIGH.value:159        if not _premium_features_allowed():160            raise CostSafetyViolation(161                "thinking_level='high' is blocked because it incurs additional per-request "162                "surcharges. Omit thinking_level, use 'minimal', or set "163                "ALLOW_PREMIUM_FEATURES=true on the server.",164                field="thinking_level",165            )166    elif thinking_level is None and "thinking_level" in request.model_fields_set:167        arguments.pop("thinking_level", None)168    elif not _premium_features_allowed() and thinking_level == ThinkingLevel.MINIMAL.value:169        logger.info("Cost safety: allowing thinking_level='minimal' (low-cost mode).")170 171    arguments["limit_generations"] = True172 173    logger.info(174        "Cost guardrails applied: num_images=1, limit_generations=True, "175        "enable_web_search=%s, resolution=%s",176        arguments.get("enable_web_search", False),177        arguments.get("resolution", SAFE_RESOLUTION.value),178    )179    return arguments180 181 182def _prepare_upstream_call(183    request: GenerateRequest,184    arguments: dict[str, Any],185) -> tuple[str, dict[str, Any]]:186    upstream_arguments = dict(arguments)187    upstream_arguments.pop("mode", None)188    parent_image_url = upstream_arguments.pop(PARENT_IMAGE_URL_FIELD, None)189 190    if parent_image_url:191        existing_urls = list(upstream_arguments.get("image_urls") or [])192        image_urls = [parent_image_url]193        for url in existing_urls:194            if url != parent_image_url:195                image_urls.append(url)196        upstream_arguments["image_urls"] = image_urls197        logger.info(198            "Routing refinement request to %s with parent_image_url as primary asset.",199            NANO_BANANA_2_EDIT_ENDPOINT,200        )201        return NANO_BANANA_2_EDIT_ENDPOINT, upstream_arguments202 203    if request.mode == GenerationMode.CREATE:204        for field in REFERENCE_ASSET_FIELDS:205            upstream_arguments.pop(field, None)206        return NANO_BANANA_2_TEXT_ENDPOINT, upstream_arguments207 208    if not request.image_urls:209        raise HTTPException(210            status_code=400,211            detail="Edit mode requires at least one image URL in image_urls.",212        )213 214    return NANO_BANANA_2_EDIT_ENDPOINT, upstream_arguments215 216 217def _resolve_gift_token(218    request: GenerateRequest,219    gift_token_header: str | None,220) -> str | None:221    return (gift_token_header or request.gift_token or "").strip() or None222 223 224@app.exception_handler(CostSafetyViolation)225async def cost_safety_exception_handler(226    _request: Request,227    exc: CostSafetyViolation,228) -> JSONResponse:229    logger.warning("Request blocked by cost guardrails: %s", exc.message)230    detail: dict[str, Any] = {231        "error": "cost_safety_violation",232        "message": exc.message,233    }234    if exc.field:235        detail["field"] = exc.field236    return JSONResponse(status_code=400, content={"detail": detail})237 238 239@app.exception_handler(GiftTokenError)240async def gift_token_exception_handler(241    _request: Request,242    exc: GiftTokenError,243) -> JSONResponse:244    return JSONResponse(245        status_code=exc.status_code,246        content={"detail": {"error": "gift_token_invalid", "message": exc.message}},247    )248 249 250@app.exception_handler(RequestValidationError)251async def validation_exception_handler(252    _request: Request,253    exc: RequestValidationError,254) -> JSONResponse:255    errors = exc.errors()256    for error in errors:257        message = error.get("msg", "")258        if "Cost safety policy" in message:259            logger.warning("Request blocked by schema cost guardrails: %s", message)260            return JSONResponse(261                status_code=400,262                content={263                    "detail": {264                        "error": "cost_safety_violation",265                        "message": message,266                        "field": ".".join(str(part) for part in error.get("loc", [])),267                    }268                },269            )270 271    return JSONResponse(status_code=422, content={"detail": errors})272 273 274@app.get("/api/v1/health", response_class=JSONResponse)275async def health() -> JSONResponse:276    snapshot = get_health_snapshot()277    status_code = resolve_health_status_code(snapshot)278    return JSONResponse(status_code=status_code, content=snapshot)279 280 281@app.post("/api/v1/admin/login", response_model=AdminLoginResponse)282async def admin_login(payload: AdminLoginRequest) -> JSONResponse:283    if len((os.getenv("MASTER_PASSWORD") or "").strip()) == 0:284        raise HTTPException(285            status_code=500,286            detail="MASTER_PASSWORD is not configured on the server.",287        )288    if not is_admin_session_configured():289        raise HTTPException(290            status_code=500,291            detail="ADMIN_SESSION_SECRET is not configured on the server.",292        )293    if not verify_master_password(payload.password):294        raise HTTPException(status_code=401, detail="Invalid master password.")295 296    token = create_admin_session_token()297    response = JSONResponse(content={"message": "Admin session established."})298    set_admin_session_cookie(response, token)299    return response300 301 302@app.post("/api/v1/admin/logout")303async def admin_logout() -> JSONResponse:304    response = JSONResponse(content={"message": "Logged out."})305    clear_admin_session_cookie(response)306    return response307 308 309@app.get("/api/v1/admin/session")310async def admin_session(request: Request) -> dict[str, bool]:311    return {"authenticated": read_admin_session(request)}312 313 314@app.post("/api/v1/admin/generate-token", response_model=GiftTokenCreateResponse)315async def admin_generate_token(316    request: Request,317    session: AsyncSession = Depends(get_db_session),318) -> GiftTokenCreateResponse:319    require_admin_session(request)320    gift = await create_gift_token(session)321    invite_path = f"/invite/{gift.token}"322    return GiftTokenCreateResponse(323        token=gift.token,324        invite_path=invite_path,325        expires_at=gift.expires_at,326        created_at=gift.created_at,327    )328 329 330@app.get("/api/v1/admin/tokens", response_model=list[GiftTokenCreateResponse])331async def admin_list_tokens(332    request: Request,333    session: AsyncSession = Depends(get_db_session),334) -> list[GiftTokenCreateResponse]:335    require_admin_session(request)336    result = await session.execute(337        select(GiftToken).order_by(GiftToken.created_at.desc()).limit(25)338    )339    tokens = result.scalars().all()340    return [341        GiftTokenCreateResponse(342            token=row.token,343            invite_path=f"/invite/{row.token}",344            expires_at=row.expires_at,345            created_at=row.created_at,346        )347        for row in tokens348    ]349 350 351@app.get("/api/v1/gift-tokens/{token}", response_model=GiftTokenStatusResponse)352async def gift_token_status(353    token: str,354    session: AsyncSession = Depends(get_db_session),355) -> GiftTokenStatusResponse:356    from datetime import datetime, timezone357 358    gift = await get_gift_token(session, token)359    if gift is None:360        return GiftTokenStatusResponse(361            token=token,362            valid=False,363            is_used=False,364            expires_at=datetime.now(timezone.utc),365            message="Invalid or unknown gift token.",366        )367 368    try:369        validated = await validate_gift_token(session, token)370        return GiftTokenStatusResponse(371            token=validated.token,372            valid=True,373            is_used=False,374            expires_at=validated.expires_at,375            message="Gift token is valid and ready to use.",376        )377    except GiftTokenError as exc:378        return GiftTokenStatusResponse(379            token=token,380            valid=False,381            is_used=gift.is_used,382            expires_at=gift.expires_at,383            message=exc.message,384        )385 386 387@app.post("/api/v1/generate", response_model=GenerateResponse)388@app.post("/api/v1/edit", response_model=GenerateResponse)389async def generate(390    request: GenerateRequest,391    http_request: Request,392    gift_token_header: str | None = Header(default=None, alias="X-Gift-Token"),393    session: AsyncSession = Depends(get_db_session),394) -> GenerateResponse:395    if not os.getenv("FAL_KEY", "").strip():396        raise HTTPException(397            status_code=500,398            detail="FAL_KEY environment variable is not configured.",399        )400 401    is_admin = read_admin_session(http_request)402    gift_token_value = _resolve_gift_token(request, gift_token_header)403 404    if is_admin:405        logger.info(406            "Master admin session authenticated; bypassing gift token validation and burn tracking."407        )408    elif not gift_token_value:409        raise HTTPException(410            status_code=403,411            detail="A valid gift token is required for generation.",412        )413    else:414        await validate_gift_token(session, gift_token_value)415 416    try:417        arguments = _apply_cost_guardrails(request)418        endpoint, upstream_arguments = _prepare_upstream_call(request, arguments)419    except CostSafetyViolation:420        raise421 422    logger.info(423        "Submitting guarded request to %s (mode=%s, refinement=%s, admin=%s, gift=%s)",424        endpoint,425        request.mode.value,426        bool(request.parent_image_url),427        is_admin,428        bool(gift_token_value) and not is_admin,429    )430 431    try:432        result = await asyncio.to_thread(433            fal_client.subscribe,434            endpoint,435            arguments=upstream_arguments,436            with_logs=True,437        )438    except Exception as exc:439        logger.exception("Upstream fal.ai request failed")440        raise HTTPException(441            status_code=502,442            detail=f"Upstream fal.ai request failed: {exc}",443        ) from exc444 445    try:446        response = GenerateResponse.model_validate(result)447    except Exception as exc:448        logger.exception("Failed to validate upstream response")449        raise HTTPException(450            status_code=502,451            detail=f"Upstream fal.ai returned an invalid response: {exc}",452        ) from exc453 454    if gift_token_value and not is_admin:455        await consume_gift_token(session, gift_token_value)456        logger.info("Gift token %s marked as used after successful generation.", gift_token_value)457 458    return response459 460 461if __name__ == "__main__":462    import uvicorn463 464    port = int(os.environ.get("PORT", 7860))465    uvicorn.run("main:app", host="0.0.0.0", port=port, reload=False)466