esotericelf/image_edit_creation
0
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 