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]