auravision-api / image_proxy.py
AuraVision Deployer
Deploying AuraVision backend engine to Hugging Face Spaces
c109c49
Raw History Blame Contribute Delete
4.19 kB
import httpx
import asyncio
import urllib.parse
import logging
from collections import OrderedDict
from fastapi import HTTPException, Request, Response
from security_utils import is_safe_url
logger = logging.getLogger(__name__)
# Bounded LRU Cache for proxied images to prevent memory exhaustion DoS
class AsyncLRUCache:
def __init__(self, capacity: int = 100):
self.capacity = capacity
self.cache = OrderedDict()
self.lock = asyncio.Lock()
async def get(self, key: str):
async with self.lock:
if key not in self.cache:
return None
self.cache.move_to_end(key)
return self.cache[key]
async def set(self, key: str, value: bytes):
async with self.lock:
self.cache[key] = value
self.cache.move_to_end(key)
if len(self.cache) > self.capacity:
self.cache.popitem(last=False)
IMAGE_CACHE = AsyncLRUCache(capacity=50)
IN_FLIGHT = {}
IN_FLIGHT_LOCK = asyncio.Lock()
async def proxy_image_logic(request: Request):
raw_query = request.url.query
if not raw_query.startswith("url="):
raise HTTPException(status_code=400, detail="Missing url query parameter.")
# Unquote twice to handle potential double-escaped query params
url = urllib.parse.unquote(urllib.parse.unquote(raw_query[4:]))
# SSRF Protection: Ensure target host is public and safe
if not is_safe_url(url):
logger.warning(f"SSRF attempt blocked. Insecure destination URL: {url}")
raise HTTPException(status_code=400, detail="Forbidden image destination URL host.")
cached_content = await IMAGE_CACHE.get(url)
if cached_content:
return Response(content=cached_content, media_type="image/jpeg")
# Guard in-flight duplicate requests
async with IN_FLIGHT_LOCK:
if url in IN_FLIGHT:
event = IN_FLIGHT[url]
else:
event = asyncio.Event()
IN_FLIGHT[url] = event
if event.is_set():
# Already fetched while we waited for lock
cached_content = await IMAGE_CACHE.get(url)
if cached_content:
return Response(content=cached_content, media_type="image/jpeg")
raise HTTPException(status_code=500, detail="Failed to fetch image.")
try:
# Check if URL belongs to pollination API prompts and trim to prevent HTTP header bloat
if "/prompt/" in url:
parts = url.split("?", 1)
base, params = parts[0], ("?" + parts[1] if len(parts) > 1 else "")
p_start = base.find("/prompt/") + 8
# Limit prompts to 600 chars to avoid command/argument length issues
base = base[:p_start] + base[p_start:][:600]
url = base + params
# Perform fetch with timeout
async with httpx.AsyncClient() as client:
for attempt in range(3):
try:
r = await client.get(url, timeout=15.0, follow_redirects=True)
if r.status_code == 200:
# Validate that returned content looks like an image
content_type = r.headers.get("Content-Type", "")
if not content_type.startswith("image/"):
logger.error(f"Target URL returned non-image content-type: {content_type}")
raise HTTPException(status_code=400, detail="Target URL is not an image.")
await IMAGE_CACHE.set(url, r.content)
return Response(content=r.content, media_type="image/jpeg")
if r.status_code == 429:
await asyncio.sleep(1)
except Exception as e:
logger.debug(f"Attempt {attempt} failed: {e}")
if attempt < 2:
await asyncio.sleep(1)
raise HTTPException(status_code=500, detail="Failed to fetch image.")
finally:
event.set()
async with IN_FLIGHT_LOCK:
if url in IN_FLIGHT:
del IN_FLIGHT[url]