import io
import time
import requests
import psycopg2
from psycopg2.extras import execute_values
import torch
import torchvision.transforms as T
from PIL import Image

# ================= 配置初始化 =================
API_BASE_URL = "http://192.168.1.67/PartVision/application/index.php"
SRC_NAME = "MEVOTECH"

DB_CONFIG = {
    "host": "192.168.1.42",
    "database": "postgres",
    "user": "postgres",
    "password": "yiparts007"
}

device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using device: {device}")

# ================= 【核心修改 1：加载 DINOv2 模型】 =================
# 推荐使用 dinov2_vits14 (384维，极快) 或 dinov2_vitl14 (1024维，精细但稍慢)
# 这里选用 vitb14 (768维)，在体积和精度上对汽车零件非常平衡
MODEL_NAME = "dinov2_vitb14" 
print(f"Loading DINOv2 model [{MODEL_NAME}] from PyTorch Hub...")
model = torch.hub.load('facebookresearch/dinov2', MODEL_NAME).to(device)
model.eval()  # 必须切换到评估模式

# DINOv2 标准图像预处理流（替代 clip.load 返回的 preprocess）
preprocess = T.Compose([
    T.Resize(256, interpolation=T.InterpolationMode.BICUBIC),
    T.CenterCrop(224),
    T.ToTensor(),
    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])

# ================= 核心处理函数 =================
def download_and_vectorize(img_url):
    """下载图片并转换为 DINOv2 特征向量"""
    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)
        
        # 【核心修改 2：预处理与提取 DINOv2 向量】
        # DINOv2 需要保证图片为 RGB 格式
        img_pil = Image.open(image_bytes).convert('RGB')
        image = preprocess(img_pil).unsqueeze(0).to(device)
        
        with torch.no_grad():
            # DINOv2 直接前向传播即可输出类 Embedding 向量
            vec = model(image)
            
        # 依旧进行 L2 归一化，方便数据库进行余弦相似度计算
        vec = vec / vec.norm(dim=-1, keepdim=True)
        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
    
    print("Start processing task with DINOv2 & Two-Way Synchronization Strategy...")
    
    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_768 (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
                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:
                cur.execute("""
                    DELETE FROM part_images_768 
                    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:
                cur.execute("""
                    DELETE FROM part_images_768 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_768 
                    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 db_last_str == img_last:
                        continue
                    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_768 
                                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

                # 新图入库
                embedding_vector = download_and_vectorize(img_url)
                if embedding_vector is not None:
                    cur.execute("""
                        INSERT INTO part_images_768 (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()