#!/usr/bin/env python3
"""
Image generation via CLIProxy (OpenAI-compatible endpoint).
Zero-dependency version: Uses standard library (urllib) only.
"""

import os
import sys
import json
import base64
import argparse
import urllib.request
import urllib.error
from pathlib import Path
from datetime import datetime
from concurrent.futures import ThreadPoolExecutor, as_completed

# Try to load dotenv if available (optional)
try:
    from dotenv import load_dotenv
    # Load .env from multiple locations (priority order)
    env_paths = [
        Path(__file__).parent / ".env",
        Path(__file__).parent.parent / ".env",
        Path.cwd() / ".env",
    ]
    for env_path in env_paths:
        if env_path.exists():
            load_dotenv(env_path)
except ImportError:
    pass

# Configuration with defaults
DEFAULT_BASE = os.getenv("CLIPROXY_BASE", "http://127.0.0.1:8317/v1")
DEFAULT_TOKEN = os.getenv("CLIPROXY_TOKEN", "ccs-internal-managed")
DEFAULT_MODEL = os.getenv("IMAGE_MODEL", "gemini-3-pro-image-preview")


def make_request(url, method="GET", headers=None, data=None, timeout=120):
    """Helper to make HTTP requests using standard urllib."""
    if headers is None:
        headers = {}
    
    if data is not None:
        json_data = json.dumps(data).encode("utf-8")
        headers["Content-Type"] = "application/json"
    else:
        json_data = None

    req = urllib.request.Request(url, data=json_data, headers=headers, method=method)
    
    try:
        with urllib.request.urlopen(req, timeout=timeout) as response:
            return {
                "status": response.status,
                "data": json.loads(response.read().decode("utf-8")),
                "error": None
            }
    except urllib.error.HTTPError as e:
        return {
            "status": e.code,
            "data": None,
            "error": f"HTTP {e.code}: {e.reason}"
        }
    except urllib.error.URLError as e:
        return {
            "status": 0,
            "data": None,
            "error": f"Connection failed: {e.reason}"
        }
    except Exception as e:
        return {
            "status": 0,
            "data": None,
            "error": str(e)
        }


def preflight_check(base_url: str = DEFAULT_BASE, token: str = DEFAULT_TOKEN) -> dict:
    """
    Check CLIProxy status and auth before generation.
    Returns: {"ready": bool, "error": str, "models": list}
    """
    result = {"ready": False, "error": "", "models": []}
    
    # Build headers with auth
    headers = {}
    if token:
        headers["Authorization"] = f"Bearer {token}"
    
    # Step 1: Check CLIProxy is running
    resp = make_request(f"{base_url}/models", headers=headers, timeout=5)
    
    if resp["error"]:
        if "Connection failed" in resp["error"]:
            result["error"] = (
                "❌ CLIProxy không chạy!\n"
                "   → Chạy: ccs config\n"
                "   → Hoặc: cli-proxy-api --login"
            )
        else:
            result["error"] = f"❌ Lỗi kết nối CLIProxy: {resp['error']}"
        return result

    data = resp["data"]
    models = [m.get("id") for m in data.get("data", [])]
    result["models"] = models
    
    # Step 2: Check if any models available (auth check)
    if not models:
        result["error"] = (
            "❌ Không có accounts nào được kết nối!\n"
            "   → Chạy: ccs agy --login\n"
            "   → Hoặc truy cập: http://localhost:3000/cliproxy"
        )
        return result
    
    # Step 3: Check image model specifically
    image_model = os.getenv("IMAGE_MODEL", DEFAULT_MODEL)
    if image_model not in models:
        available = ", ".join(models[:5])
        if len(models) > 5:
            available += f"... (+{len(models)-5} more)"
        result["error"] = (
            f"⚠️ Model '{image_model}' không khả dụng!\n"
            f"   → Models hiện có: {available}\n"
            f"   → Cần đăng nhập AGY account có quyền image generation."
        )
        return result
    
    result["ready"] = True
    return result


def encode_image(image_path: str) -> str:
    """Encode image to base64 data URL."""
    path = Path(image_path)
    if not path.exists():
        raise FileNotFoundError(f"Image not found: {image_path}")
    
    with open(path, "rb") as f:
        data = base64.b64encode(f.read()).decode("utf-8")
    
    ext = path.suffix.lower().replace(".", "")
    mime_map = {"jpg": "jpeg", "jpeg": "jpeg", "png": "png", "webp": "webp", "gif": "gif"}
    mime = mime_map.get(ext, "jpeg")
    
    return f"data:image/{mime};base64,{data}"


def generate_image(
    prompt: str,
    output_dir: Path,
    reference_images: list = None,
    model: str = DEFAULT_MODEL,
    base_url: str = DEFAULT_BASE,
    token: str = DEFAULT_TOKEN,
    prefix: str = "gen"
) -> dict:
    """
    Generate a single image via CLIProxy.
    Returns: {"success": bool, "path": str, "error": str}
    """
    # Build message content
    content = [{"type": "text", "text": prompt}]
    
    if reference_images:
        for img_path in reference_images:
            try:
                content.append({
                    "type": "image_url",
                    "image_url": {"url": encode_image(img_path)}
                })
            except FileNotFoundError as e:
                return {"success": False, "error": str(e)}
    
    payload = {
        "model": model,
        "messages": [{"role": "user", "content": content}],
        "max_tokens": 4096
    }
    
    headers = {}
    if token:
        headers["Authorization"] = f"Bearer {token}"
    
    resp = make_request(
        f"{base_url}/chat/completions",
        method="POST",
        headers=headers,
        data=payload,
        timeout=120
    )

    if resp["error"]:
        return {"success": False, "error": f"Request failed: {resp['error']}"}

    try:
        data = resp["data"]
        
        # Extract image (CLIProxy-specific path)
        choices = data.get("choices", [])
        if not choices:
             return {"success": False, "error": "No choices in response"}

        images = choices[0].get("message", {}).get("images", [])
        if not images:
            return {"success": False, "error": "No images in response"}
        
        img_url = images[0].get("image_url", {}).get("url", "")
        if not img_url.startswith("data:image"):
            return {"success": False, "error": "Invalid image data format"}
        
        # Decode and save
        header, b64_data = img_url.split(",", 1)
        ext = "jpg" if "jpeg" in header else "png"
        timestamp = int(datetime.now().timestamp() * 1000)
        filename = f"{prefix}_{timestamp}.{ext}"
        
        output_dir.mkdir(parents=True, exist_ok=True)
        output_path = output_dir / filename
        
        with open(output_path, "wb") as f:
            f.write(base64.b64decode(b64_data))
        
        return {"success": True, "path": str(output_path)}
    
    except Exception as e:
        return {"success": False, "error": str(e)}


def batch_generate(
    prompts: list,
    output_dir: Path,
    max_concurrent: int = 4,
    **kwargs
) -> list:
    """Generate multiple images in parallel."""
    results = []
    
    with ThreadPoolExecutor(max_workers=max_concurrent) as executor:
        futures = {}
        for i, prompt in enumerate(prompts):
            future = executor.submit(
                generate_image,
                prompt,
                output_dir,
                prefix=f"gen_{i+1:03d}",
                **kwargs
            )
            futures[future] = prompt
        
        for future in as_completed(futures):
            prompt = futures[future]
            result = future.result()
            result["prompt"] = prompt[:100]  # Truncate for display
            results.append(result)
            
            status = "✅" if result["success"] else "❌"
            print(f"{status} {prompt[:50]}...")
    
    return results


def main():
    parser = argparse.ArgumentParser(
        description="Generate images via CLIProxy",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog="""
Examples:
  %(prog)s --prompt "A cyberpunk hacker"
  %(prog)s --prompt "Same person, new outfit" --ref anchor.jpg
  %(prog)s --prompts prompts.txt --concurrent 4
  %(prog)s --check
        """
    )
    parser.add_argument("--prompt", help="Single prompt text")
    parser.add_argument("--prompts", help="File with prompts (one per line)")
    parser.add_argument("--ref", nargs="+", help="Reference image path(s)")
    parser.add_argument("--output", default="./output", help="Output directory")
    parser.add_argument("--model", default=DEFAULT_MODEL, help="Model name")
    parser.add_argument("--concurrent", type=int, default=4, help="Max parallel requests")
    parser.add_argument("--prefix", default="gen", help="Filename prefix")
    parser.add_argument("--check", action="store_true", help="Run preflight check only")
    
    args = parser.parse_args()
    
    # Preflight check
    print("🔍 Checking CLIProxy...")
    check = preflight_check()
    
    if not check["ready"]:
        print(check["error"])
        sys.exit(1)
    
    print(f"✅ Ready! Model: {args.model}")
    print(f"   Available models: {len(check['models'])}\n")
    
    if args.check:
        print("Models:", ", ".join(check["models"][:10]))
        if len(check["models"]) > 10:
            print(f"   ... and {len(check['models'])-10} more")
        sys.exit(0)
    
    # Determine prompts
    prompts = []
    if args.prompt:
        prompts = [args.prompt]
    elif args.prompts:
        with open(args.prompts, encoding="utf-8") as f:
            prompts = [line.strip() for line in f if line.strip()]
    else:
        parser.error("Either --prompt or --prompts is required")
    
    # Setup output with date subdirectory
    if Path(args.output).is_absolute():
        base_output = Path(args.output)
    else:
        # Force relative paths to resolve against current working directory (where user is)
        # NOT the script directory
        base_output = Path.cwd() / args.output
        
    output_dir = base_output / datetime.now().strftime("%Y%m%d")
    
    print(f"🎨 Generating {len(prompts)} image(s)...")
    print(f"📂 Output Dir: {output_dir.absolute()}\n")
    
    # Generate
    if len(prompts) == 1:
        result = generate_image(
            prompts[0],
            output_dir,
            reference_images=args.ref,
            model=args.model,
            prefix=args.prefix
        )
        if result["success"]:
            print(f"✅ Saved: {result['path']}")
        else:
            print(f"❌ Error: {result['error']}")
            sys.exit(1)
    else:
        results = batch_generate(
            prompts,
            output_dir,
            max_concurrent=args.concurrent,
            reference_images=args.ref,
            model=args.model
        )
        
        # Save metadata
        metadata_path = output_dir / "metadata.json"
        with open(metadata_path, "w", encoding="utf-8") as f:
            json.dump(results, f, indent=2, ensure_ascii=False)
        
        success_count = sum(1 for r in results if r["success"])
        print(f"\n✅ Generated {success_count}/{len(results)} images")
        print(f"📁 Output: {output_dir}")
        print(f"📋 Metadata: {metadata_path}")


if __name__ == "__main__":
    main()
