import io
import time
import requests
import psycopg2
from psycopg2.extras import execute_values
import torch
import clip
from PIL import Image

# ================= 配置初始化 =================
API_BASE_URL = "http://192.168.1.67/PartVision/application/index.php"
SRC_NAME = "YIPARTS"

DB_CONFIG = {
    "host": "192.168.1.42",
    "database": "postgres",
    "user": "postgres",
    "password": "yiparts007"
}

# 💡 提示：如果你正在将 512 维全面升级为 1024 维，请将此开关设为 True。
# 当设置为 True 时，程序会强行重新下载并计算所有图片的 1024 维向量并覆盖数据库。
# 洗完一次数据后，可以将其改回 False，以恢复正常的高效增量时间戳校验。
FORCE_RECODE_ALL = True 

device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using device: {device}")

# 🚀 核心升级：加载旗舰级 RN50x64 模型（输出 1024 维向量）
model, preprocess = clip.load("RN50x64", device=device)

# 💡 规避兼容问题：如果是 CPU 运行，强制转换为 float32，防止部分 CPU 环境下半精度(float16)报错
if device == "cpu":
    model = model.float()

# ================= 核心处理函数 =================
def download_and_vectorize(img_url):
    """下载图片并转换为 RN50x64 (1024维) 向量"""
    try:
        response = requests.get(img_url, timeout=10)
        if response.status_code != 200:
            print(f" [Warning] Failed to download image: {img_url} (Status: {response.status_code})")
            return None
            
        image_bytes = io.BytesIO(response.content)
        image = preprocess(Image.open(image_bytes)).unsqueeze(0).to(device)
        
        # 如果是 CPU 运行，同步将图片张量转为 float32
        if device == "cpu":
            image = image.float()
        
        with torch.no_grad():
            vec = model.encode_image(image)
            
        # 向量归一化（以便后续直接进行余弦相似度点积计算）
        vec = vec / vec.norm(dim=-1, keepdim=True)
        
        # 转换为 Python List 结构，此时长度严格为 1024
        return vec.cpu().numpy()[0].tolist()
        
    except Exception as e:
        print(f" [Error] Process image error ({img_url}): {e}")
        return None

def main():
    conn = psycopg2.connect(**DB_CONFIG)
    cur = conn.cursor()
    
    last_id = None
    total_count = 0
    
    if FORCE_RECODE_ALL:
        print("🚀 WARNING: FORCE_RECODE_ALL is enabled. Will re-download and encode ALL images into 1024-dim vectors.")
    print("Start processing task with Two-Way Synchronization Strategy (1024-dim)...")
    
    while True:
        params = {"src": SRC_NAME}
        if last_id is not None:
            params["lastid"] = last_id
            
        try:
            response = requests.get(API_BASE_URL, params=params, timeout=15)
            if response.status_code != 200:
                print(f"API Request error: HTTP {response.status_code}. Retry in 5s...")
                time.sleep(5)
                continue
                
            data = response.json()
        except Exception as e:
            print(f"Network error or Invalid JSON: {e}. Retry in 5s...")
            time.sleep(5)
            continue

        if data.get("code") == 404:
            print(">>> API returned 404. No more data. Task finished successfully!")
            break

        if data.get("code") != 0:
            print(f"API returned unexpected code: {data}. stopping...")
            break

        src_id = data.get("src_id")
        brand = data.get("brand")
        number = data.get("number")
        partname_cn = data.get("partname_cn")
        partname_en = data.get("partname_en")
        product_url = data.get("url")
        images = data.get("images", [])

        print(f"Processing src_id: {src_id} | Number: {number} | Found {len(images)} images")

        try:
			# 1. 写入或更新商品主表（联合唯一键修复）
            cur.execute("""
                INSERT INTO part_products_1024 (src_id, src, brand, number, partname_cn, partname_en, url)
                VALUES (%s, %s, %s, %s, %s, %s, %s)
                ON CONFLICT (src_id, src) DO UPDATE SET 
                    brand = EXCLUDED.brand, 
                    number = EXCLUDED.number, 
                    partname_cn = EXCLUDED.partname_cn, 
                    partname_en = EXCLUDED.partname_en,
                    url = EXCLUDED.url
                RETURNING id;
            """, (src_id, SRC_NAME, brand, number, partname_cn, partname_en, product_url))
            
            product_db_id = cur.fetchone()[0]

            # =================== 【废弃图片清理】 ===================
            current_api_urls = [img.get("url") for img in images if img.get("url")]
            
            if current_api_urls:
                # 场景 A: 接口里还有图片。删除属于该商品，但不在当前接口列表里的本地图片记录
                cur.execute("""
                    DELETE FROM part_images 
                    WHERE product_id = %s AND image_url NOT IN %s;
                """, (product_db_id, tuple(current_api_urls)))
                if cur.rowcount > 0:
                    print(f"  [Clean] Deleted {cur.rowcount} obsolete image(s) from DB.")
            else:
                # 场景 B: 如果上游接口把该商品下的图片全部删空了
                cur.execute("""
                    DELETE FROM part_images WHERE product_id = %s;
                """, (product_db_id,))
                if cur.rowcount > 0:
                    print(f"  [Clean] API returned 0 images. Cleared all {cur.rowcount} image(s) for this product.")
            # ====================================================================

            # 2. 循环处理该商品下的所有图片（支持全量洗旧数据和增量更新）
            for img_info in images:
                img_url = img_info.get("url")
                img_type = img_info.get("type")
                img_last = img_info.get("last")
                
                if not img_url:
                    continue
                
                # 查询本地数据库中是否已有该产品的这张图片
                cur.execute("""
                    SELECT image_last_updated 
                    FROM part_images_1024 
                    WHERE product_id = %s AND image_url = %s
                    LIMIT 1;
                """, (product_db_id, img_url))
                
                db_record = cur.fetchone()
                
                if db_record is not None:
                    db_last_time = db_record[0]
                    db_last_str = db_last_time.strftime('%Y-%m-%d %H:%M:%S') if db_last_time else None
                    
                    # 💡 逻辑改造：如果开启了全量重刷，或者时间戳真的不一致
                    if FORCE_RECODE_ALL or (db_last_str != img_last):
                        if FORCE_RECODE_ALL:
                            print(f"  [Force Update] Upgrade to 1024-dim vector for: {img_url}")
                        else:
                            print(f"  [Update] Image expired (DB: {db_last_str} | API: {img_last}). Re-downloading...")
                            
                        embedding_vector = download_and_vectorize(img_url)
                        if embedding_vector is not None:
                            cur.execute("""
                                UPDATE part_images_1024 
                                SET image_last_updated = %s, embedding = %s 
                                WHERE product_id = %s AND image_url = %s
                            """, (img_last, embedding_vector, product_db_id, img_url))
                        continue
                    else:
                        # 没开启全量刷新且时间戳完全一致，跳过
                        continue

                # 新图入库
                embedding_vector = download_and_vectorize(img_url)
                if embedding_vector is not None:
                    cur.execute("""
                        INSERT INTO part_images_1024 (product_id, image_url, image_type, image_last_updated, embedding)
                        VALUES (%s, %s, %s, %s, %s)
                    """, (product_db_id, img_url, img_type, img_last, embedding_vector))
            
            # 提交事务
            conn.commit()
            total_count += 1
            last_id = src_id
            
        except Exception as db_err:
            conn.rollback()
            print(f" [Critical] Database operation failed for src_id {src_id}: {db_err}")
            break
            
        time.sleep(0.1)

    cur.close()
    conn.close()
    print(f"Done! Total synced records: {total_count}")

if __name__ == "__main__":
    main()